Skip to content

Commit 8625ada

Browse files
authored
Merge pull request #2234 from Jazzcort/wire-shield-into-response
LCORE-3202: Wire shield into response
2 parents 9e5afb5 + 9f2ffe5 commit 8625ada

11 files changed

Lines changed: 355 additions & 110 deletions

File tree

src/app/endpoints/a2a.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -60,7 +60,7 @@
6060
from log import get_logger
6161
from models.api.requests import QueryRequest
6262
from models.config import Action
63-
from utils.agents.query import map_agent_inference_error
63+
from utils.agents.error_handler import map_agent_inference_error
6464
from utils.conversation_compaction import apply_compaction_blocking
6565
from utils.mcp_headers import McpHeaders, mcp_headers_dependency
6666
from utils.pydantic_ai_helpers import build_agent

src/app/endpoints/responses.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,7 @@
105105
select_model_for_responses,
106106
)
107107
from utils.rh_identity import get_rh_identity_context
108-
from utils.shields import run_shield_moderation
108+
from utils.shields import run_shield_moderation_v2
109109
from utils.suid import (
110110
normalize_conversation_id,
111111
)
@@ -424,11 +424,11 @@ async def responses_endpoint_handler(
424424
attachments_text = extract_attachments_text(original_request.input)
425425

426426
endpoint_path = ENDPOINT_PATH_RESPONSES
427-
moderation_result = await run_shield_moderation(
428-
client,
427+
428+
moderation_result = await run_shield_moderation_v2(
429429
input_text + "\n\n" + attachments_text,
430-
endpoint_path,
431-
original_request.shield_ids,
430+
configuration.configuration.shields,
431+
responses_request.shield_ids,
432432
)
433433

434434
filter_server_tools = (

src/utils/agents/error_handler.py

Lines changed: 103 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,103 @@
1+
"""Error mapping for agent inference failures to structured API error responses."""
2+
3+
from typing import TypeAlias
4+
5+
from ogx_client import APIConnectionError, APIStatusError
6+
from pydantic_ai.exceptions import (
7+
AgentRunError,
8+
ContentFilterError,
9+
IncompleteToolCall,
10+
ModelAPIError,
11+
ModelHTTPError,
12+
UnexpectedModelBehavior,
13+
UsageLimitExceeded,
14+
)
15+
16+
from log import get_logger
17+
from models.api.responses.error import (
18+
AbstractErrorResponse,
19+
InternalServerErrorResponse,
20+
PromptTooLongResponse,
21+
QuotaExceededResponse,
22+
ServiceUnavailableResponse,
23+
)
24+
from utils.query import (
25+
handle_known_apistatus_errors,
26+
is_context_length_error,
27+
)
28+
29+
AgentInferenceError: TypeAlias = (
30+
AgentRunError | APIStatusError | APIConnectionError | RuntimeError
31+
)
32+
33+
logger = get_logger(__name__)
34+
35+
36+
def map_agent_inference_error(
37+
exc: AgentInferenceError,
38+
model_id: str,
39+
) -> AbstractErrorResponse:
40+
"""Map agent run failures from pydantic-ai or Llama Stack to an LCS error response.
41+
42+
Args:
43+
exc: Agent, HTTP status, connection, or context-length runtime error.
44+
model_id: Model identifier in provider/model format.
45+
46+
Returns:
47+
Structured error response for HTTP or SSE error events.
48+
49+
Raises:
50+
RuntimeError: Re-raised when ``exc`` is a non-agent ``RuntimeError`` that is
51+
not a recognized context-length failure.
52+
"""
53+
match exc:
54+
case AgentRunError() as agent_exc:
55+
return map_pydantic_agent_run_error(agent_exc, model_id)
56+
case APIStatusError() as status_exc:
57+
return handle_known_apistatus_errors(status_exc, model_id)
58+
case APIConnectionError() as connection_exc:
59+
return ServiceUnavailableResponse(
60+
backend_name="OGX",
61+
cause=str(connection_exc),
62+
)
63+
case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)):
64+
return PromptTooLongResponse(model=model_id)
65+
case _:
66+
return InternalServerErrorResponse.generic()
67+
68+
69+
def map_pydantic_agent_run_error(
70+
exc: AgentRunError, model_id: str
71+
) -> AbstractErrorResponse:
72+
"""Map pydantic-ai ``AgentRunError`` subclasses to LCS error responses.
73+
74+
Args:
75+
exc: Agent exception to map.
76+
model_id: Model identifier in provider/model format.
77+
78+
Returns:
79+
Structured error response for HTTP or SSE error events.
80+
"""
81+
match exc:
82+
case ContentFilterError() as filter_exc:
83+
return InternalServerErrorResponse.query_failed(str(filter_exc))
84+
case IncompleteToolCall():
85+
return PromptTooLongResponse(model=model_id)
86+
case UnexpectedModelBehavior():
87+
logger.error("Unexpected model behavior: %s", exc, exc_info=True)
88+
return InternalServerErrorResponse.generic()
89+
case UsageLimitExceeded():
90+
return QuotaExceededResponse.model(model_id)
91+
case ModelHTTPError() as http_exc if is_context_length_error(str(http_exc)):
92+
return PromptTooLongResponse(model=model_id)
93+
case ModelHTTPError(status_code=429):
94+
return QuotaExceededResponse.model(model_id)
95+
case ModelHTTPError():
96+
return InternalServerErrorResponse.generic()
97+
case ModelAPIError() as api_exc:
98+
return ServiceUnavailableResponse(
99+
backend_name="OGX",
100+
cause=str(api_exc),
101+
)
102+
case _:
103+
return InternalServerErrorResponse.query_failed(str(exc))

src/utils/agents/query.py

Lines changed: 1 addition & 80 deletions
Original file line numberDiff line numberDiff line change
@@ -9,12 +9,6 @@
99
from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient
1010
from pydantic_ai.exceptions import (
1111
AgentRunError,
12-
ContentFilterError,
13-
IncompleteToolCall,
14-
ModelAPIError,
15-
ModelHTTPError,
16-
UnexpectedModelBehavior,
17-
UsageLimitExceeded,
1812
)
1913
from pydantic_ai.messages import ModelRequest, ModelResponse, ToolReturnPart
2014
from pydantic_ai.run import AgentRunResult
@@ -27,14 +21,13 @@
2721
AbstractErrorResponse,
2822
InternalServerErrorResponse,
2923
PromptTooLongResponse,
30-
QuotaExceededResponse,
31-
ServiceUnavailableResponse,
3224
)
3325
from models.common.agents import AgentTurnAccumulator
3426
from models.common.moderation import ShieldModerationResult
3527
from models.common.responses.responses_api_params import ResponsesApiParams
3628
from models.common.responses.types import ResponseInput
3729
from models.common.turn_summary import TurnSummary
30+
from utils.agents.error_handler import map_agent_inference_error
3831
from utils.agents.tool_processor import (
3932
process_function_tool_call,
4033
process_function_tool_result,
@@ -45,8 +38,6 @@
4538
from utils.pydantic_ai_helpers import build_agent
4639
from utils.query import (
4740
extract_provider_and_model_from_model_id,
48-
handle_known_apistatus_errors,
49-
is_context_length_error,
5041
)
5142
from utils.responses import extract_vector_store_ids_from_tools
5243
from utils.token_counter import TokenCounter
@@ -68,76 +59,6 @@ class AgentFinishReason(str, Enum):
6859
ERROR = "error"
6960

7061

71-
def map_agent_inference_error(
72-
exc: AgentInferenceError,
73-
model_id: str,
74-
) -> AbstractErrorResponse:
75-
"""Map agent run failures from pydantic-ai or Llama Stack to an LCS error response.
76-
77-
Args:
78-
exc: Agent, HTTP status, connection, or context-length runtime error.
79-
model_id: Model identifier in provider/model format.
80-
81-
Returns:
82-
Structured error response for HTTP or SSE error events.
83-
84-
Raises:
85-
RuntimeError: Re-raised when ``exc`` is a non-agent ``RuntimeError`` that is
86-
not a recognized context-length failure.
87-
"""
88-
match exc:
89-
case AgentRunError() as agent_exc:
90-
return map_pydantic_agent_run_error(agent_exc, model_id)
91-
case APIStatusError() as status_exc:
92-
return handle_known_apistatus_errors(status_exc, model_id)
93-
case APIConnectionError() as connection_exc:
94-
return ServiceUnavailableResponse(
95-
backend_name="OGX",
96-
cause=str(connection_exc),
97-
)
98-
case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)):
99-
return PromptTooLongResponse(model=model_id)
100-
case _:
101-
return InternalServerErrorResponse.generic()
102-
103-
104-
def map_pydantic_agent_run_error(
105-
exc: AgentRunError, model_id: str
106-
) -> AbstractErrorResponse:
107-
"""Map pydantic-ai ``AgentRunError`` subclasses to LCS error responses.
108-
109-
Args:
110-
exc: Agent exception to map.
111-
model_id: Model identifier in provider/model format.
112-
113-
Returns:
114-
Structured error response for HTTP or SSE error events.
115-
"""
116-
match exc:
117-
case ContentFilterError() as filter_exc:
118-
return InternalServerErrorResponse.query_failed(str(filter_exc))
119-
case IncompleteToolCall():
120-
return PromptTooLongResponse(model=model_id)
121-
case UnexpectedModelBehavior():
122-
logger.error("Unexpected model behavior: %s", exc, exc_info=True)
123-
return InternalServerErrorResponse.generic()
124-
case UsageLimitExceeded():
125-
return QuotaExceededResponse.model(model_id)
126-
case ModelHTTPError() as http_exc if is_context_length_error(str(http_exc)):
127-
return PromptTooLongResponse(model=model_id)
128-
case ModelHTTPError(status_code=429):
129-
return QuotaExceededResponse.model(model_id)
130-
case ModelHTTPError():
131-
return InternalServerErrorResponse.generic()
132-
case ModelAPIError() as api_exc:
133-
return ServiceUnavailableResponse(
134-
backend_name="OGX",
135-
cause=str(api_exc),
136-
)
137-
case _:
138-
return InternalServerErrorResponse.query_failed(str(exc))
139-
140-
14162
def get_agent_finish_reason(response: ModelResponse) -> AgentFinishReason:
14263
"""Get the finish reason from a completed agent model response.
14364

src/utils/agents/streaming.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -43,12 +43,12 @@
4343
from models.common.responses.contexts import ResponseGeneratorContext
4444
from models.common.responses.responses_api_params import ResponsesApiParams
4545
from models.common.turn_summary import TurnSummary
46+
from utils.agents.error_handler import map_agent_inference_error
4647
from utils.agents.query import (
4748
AgentFinishReason,
4849
extract_agent_token_usage,
4950
get_agent_finish_reason,
5051
get_finish_reason_error,
51-
map_agent_inference_error,
5252
)
5353
from utils.agents.tool_processor import (
5454
process_function_tool_call,

src/utils/shields.py

Lines changed: 76 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -3,24 +3,33 @@
33
from typing import Optional
44

55
from fastapi import HTTPException
6-
from ogx_client import (
7-
APIConnectionError,
8-
AsyncOgxClient,
9-
)
10-
from ogx_client import (
11-
APIStatusError as LLSApiStatusError,
12-
)
6+
from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient
7+
from pydantic_ai.exceptions import AgentRunError
138

149
from configuration import AppConfig
10+
from log import get_logger
1511
from models.api.requests import QueryRequest
1612
from models.api.responses.error import (
1713
InternalServerErrorResponse,
1814
NotFoundResponse,
1915
ServiceUnavailableResponse,
2016
UnprocessableEntityResponse,
2117
)
22-
from models.common import ShieldModerationPassed, ShieldModerationResult
23-
from models.config import ShieldConfiguration
18+
from models.common.moderation import (
19+
ShieldModerationPassed,
20+
ShieldModerationResult,
21+
)
22+
from models.config import QuestionValidityConfig, RedactionConfig, ShieldConfiguration
23+
from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability
24+
from pydantic_ai_lightspeed.capabilities.question_validity._capability import (
25+
QuestionValidity,
26+
)
27+
from pydantic_ai_lightspeed.capabilities.redaction._capability import (
28+
PiiRedactionCapability,
29+
)
30+
from utils.agents.error_handler import map_agent_inference_error
31+
32+
logger = get_logger(__name__)
2433

2534

2635
def validate_shield_ids_override(
@@ -59,6 +68,63 @@ def validate_shield_ids_override(
5968
raise HTTPException(**response.model_dump())
6069

6170

71+
async def run_shield_moderation_v2(
72+
input_text: str,
73+
shield_configs: list[ShieldConfiguration],
74+
selected_shield_ids: Optional[list[str]] = None,
75+
) -> ShieldModerationResult:
76+
"""Run v2 shield moderation on input text.
77+
78+
Iterates through configured shields and runs moderation checks.
79+
80+
Parameters:
81+
input_text: The text to moderate.
82+
shield_configs: List of shield configurations to evaluate.
83+
selected_shield_ids: Optional list of shield names to filter by.
84+
85+
Returns:
86+
Result indicating if content was blocked or passed.
87+
"""
88+
selected_shield_configs = get_shields_for_request(
89+
shield_configs, selected_shield_ids
90+
)
91+
92+
for shield_config in selected_shield_configs:
93+
shield = build_shield(shield_config)
94+
95+
try:
96+
shield_result = await shield.run(input_text)
97+
# APIConnectionError and APIStatusError from ogx should not be raised from model_request,
98+
# because they will be caught inside AsyncOpenAI and transferred into openai's
99+
# APIConnectionError. The openai's exceptions will further transferred into ModelHTTPError
100+
# or ModelAPIError by _map_api_errors in OpenAIResponseModel.
101+
except (AgentRunError, RuntimeError) as exc:
102+
model_id = getattr(shield_config.config, "model_id", "unknown-shield-model")
103+
response = map_agent_inference_error(exc, model_id)
104+
raise HTTPException(**response.model_dump()) from exc
105+
106+
if shield_result.decision == "blocked":
107+
return shield_result
108+
109+
return ShieldModerationPassed()
110+
111+
112+
def build_shield(shield_config: ShieldConfiguration) -> AbstractSafetyCapability:
113+
"""Build a safety capability instance from a shield configuration.
114+
115+
Parameters:
116+
shield_config: The shield configuration to build from.
117+
118+
Returns:
119+
The constructed safety capability.
120+
"""
121+
match shield_config.config:
122+
case QuestionValidityConfig():
123+
return QuestionValidity(shield_config.config)
124+
case RedactionConfig():
125+
return PiiRedactionCapability(shield_config.config)
126+
127+
62128
async def run_shield_moderation(
63129
_client: AsyncOgxClient,
64130
_input_text: str,
@@ -124,7 +190,7 @@ async def append_turn_to_conversation(
124190
cause=str(e),
125191
)
126192
raise HTTPException(**error_response.model_dump()) from e
127-
except LLSApiStatusError as e:
193+
except APIStatusError as e:
128194
error_response = InternalServerErrorResponse.generic()
129195
raise HTTPException(**error_response.model_dump()) from e
130196

tests/integration/endpoints/test_responses_integration.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -173,7 +173,7 @@ def _configure_shield_blocked(
173173
moderation_id=moderation_id,
174174
)
175175
mocker.patch(
176-
"app.endpoints.responses.run_shield_moderation",
176+
"app.endpoints.responses.run_shield_moderation_v2",
177177
return_value=blocked,
178178
)
179179

0 commit comments

Comments
 (0)