Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/app/endpoints/a2a.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@
from log import get_logger
from models.api.requests import QueryRequest
from models.config import Action
from utils.agents.query import map_agent_inference_error
from utils.agents.error_handler import map_agent_inference_error
from utils.conversation_compaction import apply_compaction_blocking
from utils.mcp_headers import McpHeaders, mcp_headers_dependency
from utils.pydantic_ai_helpers import build_agent
Expand Down
10 changes: 5 additions & 5 deletions src/app/endpoints/responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,7 +105,7 @@
select_model_for_responses,
)
from utils.rh_identity import get_rh_identity_context
from utils.shields import run_shield_moderation
from utils.shields import run_shield_moderation_v2
from utils.suid import (
normalize_conversation_id,
)
Expand Down Expand Up @@ -424,11 +424,11 @@ async def responses_endpoint_handler(
attachments_text = extract_attachments_text(original_request.input)

endpoint_path = ENDPOINT_PATH_RESPONSES
moderation_result = await run_shield_moderation(
client,

moderation_result = await run_shield_moderation_v2(
input_text + "\n\n" + attachments_text,
endpoint_path,
original_request.shield_ids,
configuration.configuration.shields,
responses_request.shield_ids,
)

filter_server_tools = (
Expand Down
103 changes: 103 additions & 0 deletions src/utils/agents/error_handler.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
"""Error mapping for agent inference failures to structured API error responses."""

from typing import TypeAlias

from ogx_client import APIConnectionError, APIStatusError
from pydantic_ai.exceptions import (
AgentRunError,
ContentFilterError,
IncompleteToolCall,
ModelAPIError,
ModelHTTPError,
UnexpectedModelBehavior,
UsageLimitExceeded,
)

from log import get_logger
from models.api.responses.error import (
AbstractErrorResponse,
InternalServerErrorResponse,
PromptTooLongResponse,
QuotaExceededResponse,
ServiceUnavailableResponse,
)
from utils.query import (
handle_known_apistatus_errors,
is_context_length_error,
)

AgentInferenceError: TypeAlias = (
AgentRunError | APIStatusError | APIConnectionError | RuntimeError
)

logger = get_logger(__name__)


def map_agent_inference_error(
exc: AgentInferenceError,
model_id: str,
) -> AbstractErrorResponse:
"""Map agent run failures from pydantic-ai or Llama Stack to an LCS error response.

Args:
exc: Agent, HTTP status, connection, or context-length runtime error.
model_id: Model identifier in provider/model format.

Returns:
Structured error response for HTTP or SSE error events.

Raises:
RuntimeError: Re-raised when ``exc`` is a non-agent ``RuntimeError`` that is
not a recognized context-length failure.
"""
match exc:
case AgentRunError() as agent_exc:
return map_pydantic_agent_run_error(agent_exc, model_id)
case APIStatusError() as status_exc:
return handle_known_apistatus_errors(status_exc, model_id)
case APIConnectionError() as connection_exc:
return ServiceUnavailableResponse(
backend_name="OGX",
cause=str(connection_exc),
)
case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)):
return PromptTooLongResponse(model=model_id)
case _:
return InternalServerErrorResponse.generic()


def map_pydantic_agent_run_error(
exc: AgentRunError, model_id: str
) -> AbstractErrorResponse:
"""Map pydantic-ai ``AgentRunError`` subclasses to LCS error responses.

Args:
exc: Agent exception to map.
model_id: Model identifier in provider/model format.

Returns:
Structured error response for HTTP or SSE error events.
"""
match exc:
case ContentFilterError() as filter_exc:
return InternalServerErrorResponse.query_failed(str(filter_exc))
case IncompleteToolCall():
return PromptTooLongResponse(model=model_id)
case UnexpectedModelBehavior():
logger.error("Unexpected model behavior: %s", exc, exc_info=True)
return InternalServerErrorResponse.generic()
case UsageLimitExceeded():
return QuotaExceededResponse.model(model_id)
case ModelHTTPError() as http_exc if is_context_length_error(str(http_exc)):
return PromptTooLongResponse(model=model_id)
case ModelHTTPError(status_code=429):
return QuotaExceededResponse.model(model_id)
case ModelHTTPError():
return InternalServerErrorResponse.generic()
case ModelAPIError() as api_exc:
return ServiceUnavailableResponse(
backend_name="OGX",
cause=str(api_exc),
)
case _:
return InternalServerErrorResponse.query_failed(str(exc))
81 changes: 1 addition & 80 deletions src/utils/agents/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,6 @@
from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient
from pydantic_ai.exceptions import (
AgentRunError,
ContentFilterError,
IncompleteToolCall,
ModelAPIError,
ModelHTTPError,
UnexpectedModelBehavior,
UsageLimitExceeded,
)
from pydantic_ai.messages import ModelRequest, ModelResponse, ToolReturnPart
from pydantic_ai.run import AgentRunResult
Expand All @@ -27,14 +21,13 @@
AbstractErrorResponse,
InternalServerErrorResponse,
PromptTooLongResponse,
QuotaExceededResponse,
ServiceUnavailableResponse,
)
from models.common.agents import AgentTurnAccumulator
from models.common.moderation import ShieldModerationResult
from models.common.responses.responses_api_params import ResponsesApiParams
from models.common.responses.types import ResponseInput
from models.common.turn_summary import TurnSummary
from utils.agents.error_handler import map_agent_inference_error
from utils.agents.tool_processor import (
process_function_tool_call,
process_function_tool_result,
Expand All @@ -45,8 +38,6 @@
from utils.pydantic_ai_helpers import build_agent
from utils.query import (
extract_provider_and_model_from_model_id,
handle_known_apistatus_errors,
is_context_length_error,
)
from utils.responses import extract_vector_store_ids_from_tools
from utils.token_counter import TokenCounter
Expand All @@ -68,76 +59,6 @@ class AgentFinishReason(str, Enum):
ERROR = "error"


def map_agent_inference_error(
exc: AgentInferenceError,
model_id: str,
) -> AbstractErrorResponse:
"""Map agent run failures from pydantic-ai or Llama Stack to an LCS error response.

Args:
exc: Agent, HTTP status, connection, or context-length runtime error.
model_id: Model identifier in provider/model format.

Returns:
Structured error response for HTTP or SSE error events.

Raises:
RuntimeError: Re-raised when ``exc`` is a non-agent ``RuntimeError`` that is
not a recognized context-length failure.
"""
match exc:
case AgentRunError() as agent_exc:
return map_pydantic_agent_run_error(agent_exc, model_id)
case APIStatusError() as status_exc:
return handle_known_apistatus_errors(status_exc, model_id)
case APIConnectionError() as connection_exc:
return ServiceUnavailableResponse(
backend_name="OGX",
cause=str(connection_exc),
)
case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)):
return PromptTooLongResponse(model=model_id)
case _:
return InternalServerErrorResponse.generic()


def map_pydantic_agent_run_error(
exc: AgentRunError, model_id: str
) -> AbstractErrorResponse:
"""Map pydantic-ai ``AgentRunError`` subclasses to LCS error responses.

Args:
exc: Agent exception to map.
model_id: Model identifier in provider/model format.

Returns:
Structured error response for HTTP or SSE error events.
"""
match exc:
case ContentFilterError() as filter_exc:
return InternalServerErrorResponse.query_failed(str(filter_exc))
case IncompleteToolCall():
return PromptTooLongResponse(model=model_id)
case UnexpectedModelBehavior():
logger.error("Unexpected model behavior: %s", exc, exc_info=True)
return InternalServerErrorResponse.generic()
case UsageLimitExceeded():
return QuotaExceededResponse.model(model_id)
case ModelHTTPError() as http_exc if is_context_length_error(str(http_exc)):
return PromptTooLongResponse(model=model_id)
case ModelHTTPError(status_code=429):
return QuotaExceededResponse.model(model_id)
case ModelHTTPError():
return InternalServerErrorResponse.generic()
case ModelAPIError() as api_exc:
return ServiceUnavailableResponse(
backend_name="OGX",
cause=str(api_exc),
)
case _:
return InternalServerErrorResponse.query_failed(str(exc))


def get_agent_finish_reason(response: ModelResponse) -> AgentFinishReason:
"""Get the finish reason from a completed agent model response.

Expand Down
2 changes: 1 addition & 1 deletion src/utils/agents/streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,12 +43,12 @@
from models.common.responses.contexts import ResponseGeneratorContext
from models.common.responses.responses_api_params import ResponsesApiParams
from models.common.turn_summary import TurnSummary
from utils.agents.error_handler import map_agent_inference_error
from utils.agents.query import (
AgentFinishReason,
extract_agent_token_usage,
get_agent_finish_reason,
get_finish_reason_error,
map_agent_inference_error,
)
from utils.agents.tool_processor import (
process_function_tool_call,
Expand Down
86 changes: 76 additions & 10 deletions src/utils/shields.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,24 +3,33 @@
from typing import Optional

from fastapi import HTTPException
from ogx_client import (
APIConnectionError,
AsyncOgxClient,
)
from ogx_client import (
APIStatusError as LLSApiStatusError,
)
from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient
from pydantic_ai.exceptions import AgentRunError

from configuration import AppConfig
from log import get_logger
from models.api.requests import QueryRequest
from models.api.responses.error import (
InternalServerErrorResponse,
NotFoundResponse,
ServiceUnavailableResponse,
UnprocessableEntityResponse,
)
from models.common import ShieldModerationPassed, ShieldModerationResult
from models.config import ShieldConfiguration
from models.common.moderation import (
ShieldModerationPassed,
ShieldModerationResult,
)
from models.config import QuestionValidityConfig, RedactionConfig, ShieldConfiguration
from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability
from pydantic_ai_lightspeed.capabilities.question_validity._capability import (
QuestionValidity,
)
from pydantic_ai_lightspeed.capabilities.redaction._capability import (
PiiRedactionCapability,
)
from utils.agents.error_handler import map_agent_inference_error

logger = get_logger(__name__)


def validate_shield_ids_override(
Expand Down Expand Up @@ -59,6 +68,63 @@ def validate_shield_ids_override(
raise HTTPException(**response.model_dump())


async def run_shield_moderation_v2(
input_text: str,
shield_configs: list[ShieldConfiguration],
selected_shield_ids: Optional[list[str]] = None,
) -> ShieldModerationResult:
"""Run v2 shield moderation on input text.

Iterates through configured shields and runs moderation checks.

Parameters:
input_text: The text to moderate.
shield_configs: List of shield configurations to evaluate.
selected_shield_ids: Optional list of shield names to filter by.

Returns:
Result indicating if content was blocked or passed.
"""
selected_shield_configs = get_shields_for_request(
shield_configs, selected_shield_ids
)

for shield_config in selected_shield_configs:
shield = build_shield(shield_config)

try:
shield_result = await shield.run(input_text)
# APIConnectionError and APIStatusError from ogx should not be raised from model_request,
# because they will be caught inside AsyncOpenAI and transferred into openai's
# APIConnectionError. The openai's exceptions will further transferred into ModelHTTPError
# or ModelAPIError by _map_api_errors in OpenAIResponseModel.
except (AgentRunError, RuntimeError) as exc:
model_id = getattr(shield_config.config, "model_id", "unknown-shield-model")
response = map_agent_inference_error(exc, model_id)
raise HTTPException(**response.model_dump()) from exc

if shield_result.decision == "blocked":
return shield_result

return ShieldModerationPassed()


def build_shield(shield_config: ShieldConfiguration) -> AbstractSafetyCapability:
"""Build a safety capability instance from a shield configuration.

Parameters:
shield_config: The shield configuration to build from.

Returns:
The constructed safety capability.
"""
match shield_config.config:
case QuestionValidityConfig():
return QuestionValidity(shield_config.config)
case RedactionConfig():
return PiiRedactionCapability(shield_config.config)


async def run_shield_moderation(
_client: AsyncOgxClient,
_input_text: str,
Expand Down Expand Up @@ -124,7 +190,7 @@ async def append_turn_to_conversation(
cause=str(e),
)
raise HTTPException(**error_response.model_dump()) from e
except LLSApiStatusError as e:
except APIStatusError as e:
error_response = InternalServerErrorResponse.generic()
raise HTTPException(**error_response.model_dump()) from e

Expand Down
2 changes: 1 addition & 1 deletion tests/integration/endpoints/test_responses_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -173,7 +173,7 @@ def _configure_shield_blocked(
moderation_id=moderation_id,
)
mocker.patch(
"app.endpoints.responses.run_shield_moderation",
"app.endpoints.responses.run_shield_moderation_v2",
return_value=blocked,
)

Expand Down
Loading
Loading