Skip to content

Commit 9f2ffe5

Browse files
committed
Wire run_shield_moderation_v2 into Responses API endpoint
Add run_shield_moderation_v2 and build_shield to utils/shields.py to run shield moderation through AbstractSafetyCapability instances instead of the Llama Stack client. Update the responses endpoint to call the new function with shield configs directly. Include unit tests covering pass, block, filtering, and error handling paths.
1 parent eb34920 commit 9f2ffe5

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)