diff --git a/Makefile b/Makefile index 223acf7a0..8e8db60db 100644 --- a/Makefile +++ b/Makefile @@ -144,7 +144,7 @@ clean-llama-stack: remove-llama-stack-container ## Remove container and image run-llama-stack: ## Start Llama Stack with enriched config (for local service mode) uv run src/llama_stack_configuration.py -c $(CONFIG) -i $(LLAMA_STACK_CONFIG) -o $(LLAMA_STACK_CONFIG) && \ - uv run ogx stack run $(LLAMA_STACK_CONFIG) + uv run ogx stack run --insecure $(LLAMA_STACK_CONFIG) test-unit: ## Run the unit tests @echo "Running unit tests..." diff --git a/pyproject.toml b/pyproject.toml index b9e4ad2ed..4db826587 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -80,7 +80,7 @@ dependencies = [ # Used for token estimation before LLM calls (LCORE-1569 / conversation compaction) "tiktoken>=0.8.0", # Used for Pydantic AI - "pydantic-ai==2.16.0", + "pydantic-ai>=2.23.0", "pydantic-ai-skills>=0.11.0", # Used for OpenTelemetry instrumentation "opentelemetry-distro>=0.49b0", diff --git a/run.yaml b/run.yaml index e4d5aad47..0812c093b 100644 --- a/run.yaml +++ b/run.yaml @@ -59,6 +59,8 @@ providers: provider_type: inline::reference server: port: 8321 + # OGX 1.3+ requires TLS by default; e2e/local use plain HTTP. + insecure: true storage: backends: kv_default: # Define the storage backend type for RAG, in this case registry and RAG are unified i.e. information on registered resources (e.g. models, vector_stores) are saved together with the RAG chunks diff --git a/scripts/llama-stack-entrypoint.sh b/scripts/llama-stack-entrypoint.sh index 2ddcfd2e8..ce3a89c88 100755 --- a/scripts/llama-stack-entrypoint.sh +++ b/scripts/llama-stack-entrypoint.sh @@ -19,9 +19,10 @@ if [ -f "$LIGHTSPEED_CONFIG" ]; then if [ -f "$ENRICHED_CONFIG" ] && [ "$ENRICHMENT_FAILED" -eq 0 ]; then echo "Using enriched config: $ENRICHED_CONFIG" - exec ogx stack run "$ENRICHED_CONFIG" + # OGX 1.3+ requires TLS unless --insecure is set (e2e/local HTTP). + exec ogx stack run --insecure "$ENRICHED_CONFIG" fi fi echo "Using original config: $INPUT_CONFIG" -exec ogx stack run "$INPUT_CONFIG" +exec ogx stack run --insecure "$INPUT_CONFIG" diff --git a/src/app/endpoints/a2a.py b/src/app/endpoints/a2a.py index 1ff7a31d7..2425497ca 100644 --- a/src/app/endpoints/a2a.py +++ b/src/app/endpoints/a2a.py @@ -33,7 +33,7 @@ ) from a2a.utils import new_agent_text_message, new_task from fastapi import APIRouter, Depends, HTTPException, Request, status -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError from pydantic_ai.messages import ( @@ -378,7 +378,7 @@ async def _process_task_streaming( # pylint: disable=too-many-locals configuration, shields=query_request.shield_ids, ) - except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as e: + except (AgentRunError, ApiException, RuntimeError) as e: error_response = map_agent_inference_error(e, query_request.model or "") logger.error("Error preparing A2A agent: %s", str(e), exc_info=True) await task_updater.update_status( @@ -437,7 +437,7 @@ async def _process_task_streaming( # pylint: disable=too-many-locals ): aggregator.process_event(a2a_event) await event_queue.enqueue_event(a2a_event) - except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as e: + except (AgentRunError, ApiException, RuntimeError) as e: error_response = map_agent_inference_error(e, responses_params.model) logger.error("Error during A2A agent run: %s", str(e), exc_info=True) await task_updater.update_status( diff --git a/src/app/endpoints/conversations_v1.py b/src/app/endpoints/conversations_v1.py index 6ab693658..932e0f94e 100644 --- a/src/app/endpoints/conversations_v1.py +++ b/src/app/endpoints/conversations_v1.py @@ -4,10 +4,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request from ogx_api import ConversationNotFoundError, InvalidParameterError -from ogx_client import ( - APIConnectionError, - APIStatusError, -) +from ogx_client import ApiException from sqlalchemy.exc import SQLAlchemyError from app.database import get_session @@ -273,15 +270,21 @@ async def get_conversation_endpoint_handler( # pylint: disable=too-many-locals, chat_history=chat_history, ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse( - backend_name="OGX", cause=str(e) + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse( + backend_name="OGX", cause=str(e) + ).model_dump() + raise HTTPException(**response) from e + # In library mode, ConversationNotFoundError is raised instead of ApiException + logger.error("Conversation not found: %s", e) + response = NotFoundResponse( + resource="conversation", resource_id=normalized_conv_id ).model_dump() raise HTTPException(**response) from e - - except (APIStatusError, ConversationNotFoundError) as e: - # In library mode, ConversationNotFoundError is raised instead of APIStatusError + except ConversationNotFoundError as e: + # In library mode, ConversationNotFoundError is raised instead of ApiException logger.error("Conversation not found: %s", e) response = NotFoundResponse( resource="conversation", resource_id=normalized_conv_id @@ -384,12 +387,17 @@ async def delete_conversation_endpoint_handler( delete_response.deleted, ) - except APIConnectionError as e: - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e - - except (APIStatusError, ConversationNotFoundError, InvalidParameterError): - # In library mode, ConversationNotFoundError is raised instead of APIStatusError + except ApiException as e: + if not e.status: + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + # In library mode, ConversationNotFoundError is raised instead of ApiException + logger.warning( + "Conversation %s in LlamaStack not found. Treating as already deleted.", + normalized_conv_id, + ) + except (ConversationNotFoundError, InvalidParameterError): + # In library mode, ConversationNotFoundError is raised instead of ApiException logger.warning( "Conversation %s in LlamaStack not found. Treating as already deleted.", normalized_conv_id, @@ -460,8 +468,7 @@ async def update_conversation_endpoint_handler( # If reached this, user is authorized to update this conversation try: - conversation = retrieve_conversation(normalized_conv_id) - if conversation is None: + if retrieve_conversation(normalized_conv_id) is None: response = NotFoundResponse( resource="conversation", resource_id=normalized_conv_id ).model_dump() @@ -520,14 +527,20 @@ async def update_conversation_endpoint_handler( message="Topic summary updated successfully", ) - except APIConnectionError as e: - response = ServiceUnavailableResponse( - backend_name="OGX", cause=str(e) + except ApiException as e: + if not e.status: + response = ServiceUnavailableResponse( + backend_name="OGX", cause=str(e) + ).model_dump() + raise HTTPException(**response) from e + # In library mode, ConversationNotFoundError is raised instead of ApiException + logger.error("Conversation not found: %s", e) + response = NotFoundResponse( + resource="conversation", resource_id=normalized_conv_id ).model_dump() raise HTTPException(**response) from e - - except (APIStatusError, ConversationNotFoundError) as e: - # In library mode, ConversationNotFoundError is raised instead of APIStatusError + except ConversationNotFoundError as e: + # In library mode, ConversationNotFoundError is raised instead of ApiException logger.error("Conversation not found: %s", e) response = NotFoundResponse( resource="conversation", resource_id=normalized_conv_id diff --git a/src/app/endpoints/health.py b/src/app/endpoints/health.py index 0294df8f4..c929cd150 100644 --- a/src/app/endpoints/health.py +++ b/src/app/endpoints/health.py @@ -8,7 +8,7 @@ from typing import Annotated, Any from fastapi import APIRouter, Depends, Response, status -from ogx_client import APIConnectionError +from ogx_client import ApiException from authentication import get_auth_dependency from authentication.interface import AuthTuple @@ -78,7 +78,7 @@ async def get_providers_health_statuses() -> list[ProviderHealthStatus]: for provider in providers ] - except APIConnectionError as e: + except ApiException as e: logger.error("Failed to check providers health: %s", e) return [ ProviderHealthStatus( diff --git a/src/app/endpoints/info.py b/src/app/endpoints/info.py index 569966af7..0df993c70 100644 --- a/src/app/endpoints/info.py +++ b/src/app/endpoints/info.py @@ -3,7 +3,7 @@ from typing import Annotated, Any from fastapi import APIRouter, Depends, HTTPException, Request -from ogx_client import APIConnectionError +from ogx_client import ApiException from authentication import get_auth_dependency from authentication.interface import AuthTuple @@ -83,7 +83,7 @@ async def info_endpoint_handler( llama_stack_version=llama_stack_version, ) # connection to Llama Stack server - except APIConnectionError as e: + except ApiException as e: logger.error("Unable to connect to Llama Stack: %s", e) response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e diff --git a/src/app/endpoints/models.py b/src/app/endpoints/models.py index 6a9dd6262..cf7e6529b 100644 --- a/src/app/endpoints/models.py +++ b/src/app/endpoints/models.py @@ -4,7 +4,7 @@ from fastapi import APIRouter, HTTPException, Query, Request from fastapi.params import Depends -from ogx_client import APIConnectionError +from ogx_client import ApiException from authentication import get_auth_dependency from authentication.interface import AuthTuple @@ -106,7 +106,7 @@ async def models_endpoint_handler( return ModelsResponse(models=parsed_models) # Connection to Llama Stack server failed - except APIConnectionError as e: + except ApiException as e: logger.error("Unable to connect to Llama Stack: %s", e) response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e diff --git a/src/app/endpoints/prompts.py b/src/app/endpoints/prompts.py index 651abf1c1..d6f8db344 100644 --- a/src/app/endpoints/prompts.py +++ b/src/app/endpoints/prompts.py @@ -3,9 +3,7 @@ from typing import Annotated, Any, Optional from fastapi import APIRouter, Depends, HTTPException, Request -from ogx_client import APIConnectionError, BadRequestError -from ogx_client import APIStatusError as LLSApiStatusError -from openai._exceptions import APIStatusError as OpenAIAPIStatusError +from ogx_client import ApiException, BadRequestError from authentication import get_auth_dependency from authentication.interface import AuthTuple @@ -139,11 +137,12 @@ async def create_prompt_handler( payload = body.model_dump(exclude_none=True) created = await client.prompts.create(**payload) return PromptResourceResponse.model_validate(created.model_dump()) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + logger.error("API status error while creating prompt: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -188,11 +187,12 @@ async def list_prompts_handler( items = await client.prompts.list() data = [PromptResourceResponse.model_validate(p.model_dump()) for p in items] return PromptsListResponse(data=data) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + logger.error("API status error while listing prompts: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -247,15 +247,16 @@ async def get_prompt_handler( client = AsyncOgxClientHolder().get_client() retrieved = await client.prompts.retrieve(prompt_id, version=version) return PromptResourceResponse.model_validate(retrieved.model_dump()) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except (BadRequestError, ValueError) as e: logger.error("Prompt not found: %s", e) response = NotFoundResponse(resource="prompt", resource_id=prompt_id) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + logger.error("API status error while retrieving prompt: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -316,15 +317,16 @@ async def update_prompt_handler( payload = body.model_dump(exclude_none=True, exclude_unset=True) updated = await client.prompts.update(prompt_id, **payload) return PromptResourceResponse.model_validate(updated.model_dump()) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except (BadRequestError, ValueError) as e: logger.error("Prompt update failed: %s", e) response = NotFoundResponse(resource="prompt", resource_id=prompt_id) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + logger.error("API status error while updating prompt: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -381,14 +383,15 @@ async def delete_prompt_handler( client = AsyncOgxClientHolder().get_client() await client.prompts.delete(prompt_id) return PromptDeleteResponse(deleted=True, prompt_id=prompt_id) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except (BadRequestError, ValueError) as e: logger.error("Prompt delete failed: %s", e) return PromptDeleteResponse(deleted=False, prompt_id=prompt_id) - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + logger.error("API status error while deleting prompt: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e diff --git a/src/app/endpoints/providers.py b/src/app/endpoints/providers.py index abc876473..b0de35473 100644 --- a/src/app/endpoints/providers.py +++ b/src/app/endpoints/providers.py @@ -4,7 +4,7 @@ from fastapi import APIRouter, HTTPException, Request from fastapi.params import Depends -from ogx_client import APIConnectionError, BadRequestError +from ogx_client import ApiException, BadRequestError from ogx_client.models.list_providers_response import ListProvidersResponse from authentication import get_auth_dependency @@ -92,7 +92,7 @@ async def providers_endpoint_handler( try: client = AsyncOgxClientHolder().get_client() providers: ListProvidersResponse = await client.providers.list() - except APIConnectionError as e: + except ApiException as e: logger.error("Unable to connect to Llama Stack: %s", e) response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e @@ -100,7 +100,9 @@ async def providers_endpoint_handler( return ProvidersListResponse(providers=group_providers(providers)) -def group_providers(providers: ListProvidersResponse) -> dict[str, list[dict[str, Any]]]: +def group_providers( + providers: ListProvidersResponse, +) -> dict[str, list[dict[str, Any]]]: """Group a list of ProviderInfo objects by their API type. Args: @@ -164,11 +166,15 @@ async def get_provider_endpoint_handler( provider = await client.providers.retrieve(provider_id) return ProviderResponse(**provider.model_dump()) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + except (BadRequestError, ValueError) as e: + # Server mode raises BadRequestError; library mode raises ValueError. + logger.error("Provider not found: %s", e) + response = NotFoundResponse(resource="provider", resource_id=provider_id) raise HTTPException(**response.model_dump()) from e - except BadRequestError as e: - response = NotFoundResponse(resource="provider", resource_id=provider_id) + except ApiException as e: + if e.status: + raise + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e diff --git a/src/app/endpoints/rags.py b/src/app/endpoints/rags.py index 8ba2f1918..5b31a074e 100644 --- a/src/app/endpoints/rags.py +++ b/src/app/endpoints/rags.py @@ -4,7 +4,7 @@ from fastapi import APIRouter, HTTPException, Request from fastapi.params import Depends -from ogx_client import APIConnectionError, BadRequestError +from ogx_client import ApiException, BadRequestError from authentication import get_auth_dependency from authentication.interface import AuthTuple @@ -99,14 +99,13 @@ async def rags_endpoint_handler( # Map llama-stack vector store IDs to user-facing rag_ids from config rag_id_mapping = configuration.rag_id_mapping rag_ids = [ - configuration.resolve_index_name(rag.id, rag_id_mapping) - for rag in rags + configuration.resolve_index_name(rag.id, rag_id_mapping) for rag in rags ] return RAGListResponse(rags=rag_ids) # connection to Llama Stack server - except APIConnectionError as e: + except ApiException as e: logger.error("Unable to connect to Llama Stack: %s", e) response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e @@ -202,11 +201,13 @@ async def get_rag_endpoint_handler( status=rag_info.status or "unknown", usage_bytes=rag_info.usage_bytes or 0, ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("RAG not found: %s", e) response = NotFoundResponse(resource="rag", resource_id=rag_id) raise HTTPException(**response.model_dump()) from e + except ApiException as e: + if e.status: + raise + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index 47035eb9b..c20657c6a 100644 --- a/src/app/endpoints/responses.py +++ b/src/app/endpoints/responses.py @@ -21,15 +21,8 @@ from ogx_api import ( OpenAIResponseObjectStreamResponseOutputItemDone as OutputItemDoneChunk, ) -from ogx_client import ( - APIConnectionError, -) -from ogx_client import ( - APIStatusError as LLSApiStatusError, -) -from openai._exceptions import ( - APIStatusError as OpenAIAPIStatusError, -) +from ogx_client import ApiException +from openai._exceptions import APIStatusError as OpenAIAPIStatusError from app.endpoints.responses_telemetry import ( queue_blocked_response_event, @@ -183,12 +176,12 @@ def _http_exception_for_response_api_error( if not is_context_length_error(str(error)): return None error_response = PromptTooLongResponse(model=api_params.model) - elif isinstance(error, APIConnectionError): + elif isinstance(error, ApiException) and not error.status: error_response = ServiceUnavailableResponse( backend_name="OGX", cause=str(error), ) - elif isinstance(error, (LLSApiStatusError, OpenAIAPIStatusError)): + elif isinstance(error, (ApiException, OpenAIAPIStatusError)): error_response = handle_known_apistatus_errors(error, api_params.model) else: return None @@ -596,8 +589,7 @@ async def handle_streaming_response( ) except ( RuntimeError, - APIConnectionError, - LLSApiStatusError, + ApiException, OpenAIAPIStatusError, ) as e: _record_response_inference_result( @@ -1118,8 +1110,7 @@ async def handle_non_streaming_response( except ( RuntimeError, - APIConnectionError, - LLSApiStatusError, + ApiException, OpenAIAPIStatusError, ) as e: if not inference_metric_recorded: diff --git a/src/app/endpoints/rlsapi_v1.py b/src/app/endpoints/rlsapi_v1.py index 39bdfcfb4..0b092ae3d 100644 --- a/src/app/endpoints/rlsapi_v1.py +++ b/src/app/endpoints/rlsapi_v1.py @@ -13,7 +13,7 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request from jinja2.sandbox import SandboxedEnvironment from ogx_api.openai_responses import OpenAIResponseObject -from ogx_client import APIConnectionError, APIStatusError, RateLimitError +from ogx_client import ApiException, RateLimitError from openai._exceptions import APIStatusError as OpenAIAPIStatusError import constants @@ -79,9 +79,8 @@ class TemplateRenderError(Exception): _INFER_HANDLED_EXCEPTIONS = ( TemplateRenderError, RuntimeError, - APIConnectionError, + ApiException, RateLimitError, - APIStatusError, OpenAIAPIStatusError, ) @@ -189,13 +188,14 @@ async def _get_default_model_id() -> str: client = AsyncOgxClientHolder().get_client() try: models = parse_model_list_response(await client.openai.list()) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e @@ -259,7 +259,7 @@ async def _call_llm( The full OpenAIResponseObject from the LLM. Raises: - APIConnectionError: If the Llama Stack service is unreachable. + ApiException: If the Llama Stack service is unreachable. HTTPException: 503 if no default model is configured. """ client = AsyncOgxClientHolder().get_client() @@ -616,7 +616,7 @@ def _map_inference_error_to_http_exception( # pylint: disable=too-many-return-s ) return None - if isinstance(error, APIConnectionError): + if isinstance(error, ApiException) and not error.status: logger.error( "Unable to connect to OGX for request %s: %s", request_id, @@ -640,7 +640,7 @@ def _map_inference_error_to_http_exception( # pylint: disable=too-many-return-s ) return HTTPException(**error_response.model_dump()) - if isinstance(error, (APIStatusError, OpenAIAPIStatusError)): + if isinstance(error, (ApiException, OpenAIAPIStatusError)): logger.error("API error for request %s: %s", request_id, type(error).__name__) error_response = handle_known_apistatus_errors(error, model_id) return HTTPException(**error_response.model_dump()) diff --git a/src/app/endpoints/streaming_query.py b/src/app/endpoints/streaming_query.py index b83ae47af..fc08beb3f 100644 --- a/src/app/endpoints/streaming_query.py +++ b/src/app/endpoints/streaming_query.py @@ -7,12 +7,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse -from ogx_client import ( - APIConnectionError, -) -from ogx_client import ( - APIStatusError as LLSApiStatusError, -) +from ogx_client import ApiException from openai._exceptions import APIStatusError as OpenAIAPIStatusError from authentication import get_auth_dependency @@ -404,13 +399,19 @@ async def generate_response_with_compaction( ) yield stream_http_error_event(error_response, media_type) return - except APIConnectionError as e: + except ApiException as e: + if not e.status: + yield stream_http_error_event( + ServiceUnavailableResponse(backend_name="OGX", cause=str(e)), + media_type, + ) + return + yield stream_http_error_event( - ServiceUnavailableResponse(backend_name="OGX", cause=str(e)), - media_type, + handle_known_apistatus_errors(e, responses_params.model), media_type ) return - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except OpenAIAPIStatusError as e: yield stream_http_error_event( handle_known_apistatus_errors(e, responses_params.model), media_type ) diff --git a/src/app/endpoints/vector_stores.py b/src/app/endpoints/vector_stores.py index 38d3ccbe2..b72b1e3c7 100644 --- a/src/app/endpoints/vector_stores.py +++ b/src/app/endpoints/vector_stores.py @@ -6,13 +6,7 @@ from typing import Annotated, Any, Optional from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status -from ogx_client import ( - APIConnectionError, - BadRequestError, -) -from ogx_client import ( - APIStatusError as LLSApiStatusError, -) +from ogx_client import ApiException, BadRequestError from openai._exceptions import APIStatusError as OpenAIAPIStatusError from authentication import get_auth_dependency @@ -192,11 +186,16 @@ async def create_vector_store( usage_bytes=vector_store.usage_bytes or 0, metadata=vector_store.metadata, ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while creating vector store: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while creating vector store: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -248,11 +247,16 @@ async def list_vector_stores( ] return VectorStoresListResponse(data=data) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while listing vector stores: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while listing vector stores: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -302,17 +306,22 @@ async def get_vector_store( usage_bytes=vector_store.usage_bytes or 0, metadata=vector_store.metadata, ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("Vector store not found: %s", e) response = NotFoundResponse( resource="vector store", resource_id=vector_store_id ) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while getting vector store: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while getting vector store: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -366,17 +375,22 @@ async def update_vector_store( usage_bytes=vector_store.usage_bytes or 0, metadata=vector_store.metadata or None, ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("Vector store not found: %s", e) response = NotFoundResponse( resource="vector store", resource_id=vector_store_id ) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while updating vector store: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while updating vector store: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -418,14 +432,19 @@ async def delete_vector_store( client = AsyncOgxClientHolder().get_client() await client.vector_stores.delete(vector_store_id) return VectorStoreDeleteResponse(deleted=True, vector_store_id=vector_store_id) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except (BadRequestError, ValueError) as e: logger.error("Vector store delete failed: %s", e) return VectorStoreDeleteResponse(deleted=False, vector_store_id=vector_store_id) - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while deleting vector store: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while deleting vector store: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -527,10 +546,6 @@ async def create_file( # pylint: disable=too-many-branches,too-many-statements purpose=file_obj.purpose or "assistants", object=file_obj.object or "file", ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("Bad request for file upload: %s", e) # Check if backend rejected due to file size @@ -545,7 +560,16 @@ async def create_file( # pylint: disable=too-many-branches,too-many-statements response.status_code = status.HTTP_400_BAD_REQUEST response.detail.response = "Invalid file upload" raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while uploading file: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while uploading file: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -645,16 +669,10 @@ async def add_file_to_vector_store( # pylint: disable=too-many-locals,too-many- status=vs_file.status or "unknown", attributes=vs_file.attributes, last_error=( - vs_file.last_error.message - if vs_file.last_error is not None - else None + vs_file.last_error.message if vs_file.last_error is not None else None ), object=vs_file.object or "vector_store.file", ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("Vector store file operation failed: %s", e) # Don't assume which resource is missing - could be vector_store_id OR file_id @@ -663,7 +681,16 @@ async def add_file_to_vector_store( # pylint: disable=too-many-locals,too-many- resource_id=f"vector_store={vector_store_id}, file={body.file_id}", ) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while adding file to vector store: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while adding file to vector store: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -712,27 +739,28 @@ async def list_vector_store_files( vector_store_id=f.vector_store_id or vector_store_id, status=f.status or "unknown", attributes=f.attributes, - last_error=( - f.last_error.message - if f.last_error is not None - else None - ), + last_error=(f.last_error.message if f.last_error is not None else None), object=f.object or "vector_store.file", ) for f in files ] return VectorStoreFilesListResponse(data=data) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("Vector store not found: %s", e) response = NotFoundResponse( resource="vector_store", resource_id=vector_store_id ) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while listing vector store files: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while listing vector store files: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -786,21 +814,24 @@ async def get_vector_store_file( status=vs_file.status or "unknown", attributes=vs_file.attributes, last_error=( - vs_file.last_error.message - if vs_file.last_error is not None - else None + vs_file.last_error.message if vs_file.last_error is not None else None ), object=vs_file.object or "vector_store.file", ) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except BadRequestError as e: logger.error("Vector store file not found: %s", e) response = NotFoundResponse(resource="file", resource_id=file_id) raise HTTPException(**response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while getting vector store file: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while getting vector store file: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e @@ -847,14 +878,19 @@ async def delete_vector_store_file( file_id=file_id, ) return VectorStoreFileDeleteResponse(deleted=True, file_id=file_id) - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) - raise HTTPException(**response.model_dump()) from e except (BadRequestError, ValueError) as e: logger.error("Vector store file delete failed: %s", e) return VectorStoreFileDeleteResponse(deleted=False, file_id=file_id) - except (LLSApiStatusError, OpenAIAPIStatusError) as e: + except ApiException as e: + if not e.status: + logger.error("Unable to connect to Llama Stack: %s", e) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) + raise HTTPException(**response.model_dump()) from e + + logger.error("API status error while deleting vector store file: %s", e) + error_response = handle_known_apistatus_errors(e, "llama-stack") + raise HTTPException(**error_response.model_dump()) from e + except OpenAIAPIStatusError as e: logger.error("API status error while deleting vector store file: %s", e) error_response = handle_known_apistatus_errors(e, "llama-stack") raise HTTPException(**error_response.model_dump()) from e diff --git a/src/app/main.py b/src/app/main.py index 330eb9789..57fd70180 100644 --- a/src/app/main.py +++ b/src/app/main.py @@ -9,7 +9,7 @@ from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse -from ogx_client import APIConnectionError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from starlette.routing import Mount, Route, WebSocketRoute from starlette.types import ASGIApp, Message, Receive, Scope, Send @@ -99,7 +99,7 @@ async def lifespan(_app: FastAPI) -> AsyncIterator[None]: else: logger.debug("Llama Stack version: %s", llama_stack_version) degraded_tracker.set_healthy() - except APIConnectionError as e: + except ApiException as e: # if degraded mode is allowed, simply ignore the exception llama_stack_url = llama_stack_config.url logger.error( @@ -126,7 +126,7 @@ async def lifespan(_app: FastAPI) -> AsyncIterator[None]: if not degraded_tracker.is_degraded(): try: await setup_model_metrics() - except APIConnectionError as e: + except ApiException as e: logger.warning("Failed to set up model metrics: %s", e, exc_info=True) logger.info("App startup complete") diff --git a/src/client.py b/src/client.py index 1e6be32bb..2b4d2b48c 100644 --- a/src/client.py +++ b/src/client.py @@ -8,7 +8,7 @@ import yaml from fastapi import HTTPException from ogx.core.library_client import AsyncOGXAsLibraryClient -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient import constants from authorization.azure_token_manager import AzureEntraIDManager @@ -221,7 +221,7 @@ async def reload_library_client(self) -> AsyncOgxClient: try: client = AsyncOGXAsLibraryClient(self._config_path) await client.initialize() - except APIConnectionError as e: + except ApiException as e: error_response = ServiceUnavailableResponse( backend_name="OGX", cause=str(e), @@ -264,7 +264,7 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: except RuntimeError as e: logger.warning("Client not initialized, skipping model check: %s", e) return False, f"Client not initialized: {e!s}" - except (APIConnectionError, APIStatusError) as e: + except ApiException as e: logger.error("Error checking model availability: %s", e) return False, f"Error checking model availability: {e!s}" @@ -292,8 +292,7 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: except ( RuntimeError, HTTPException, - APIConnectionError, - APIStatusError, + ApiException, ) as err: logger.error("Client reload failed: %s", err) @@ -357,7 +356,7 @@ async def get_azure_base_url(self) -> Optional[str]: try: providers = await self._lsc.providers.list() - except (APIConnectionError, APIStatusError) as err: + except ApiException as err: logger.warning("Failed to list providers for Azure base_url: %s", err) return None diff --git a/src/constants.py b/src/constants.py index 253f79c71..25ec6dca9 100644 --- a/src/constants.py +++ b/src/constants.py @@ -7,7 +7,7 @@ # Minimal and maximal supported Llama Stack version MINIMAL_SUPPORTED_LLAMA_STACK_VERSION: Final[str] = "0.2.17" -MAXIMAL_SUPPORTED_LLAMA_STACK_VERSION: Final[str] = "1.2.2" +MAXIMAL_SUPPORTED_LLAMA_STACK_VERSION: Final[str] = "1.3.0" # Path to the lightspeed-stack.yaml, exported so uvicorn workers (separate # processes) can reload the configuration that the parent process selected. diff --git a/src/pydantic_ai_lightspeed/llamastack/_model.py b/src/pydantic_ai_lightspeed/llamastack/_model.py index 9a10856ff..1782a15c3 100644 --- a/src/pydantic_ai_lightspeed/llamastack/_model.py +++ b/src/pydantic_ai_lightspeed/llamastack/_model.py @@ -386,6 +386,9 @@ async def request_stream( # pylint: disable=unused-argument f"Expected ResponseCreatedEvent, got {type(first_chunk).__name__}" ) + tool_call_ids_are_response_scoped = self.profile.get( # type: ignore[attr-defined] + "openai_responses_tool_call_ids_are_response_scoped", False + ) yield OpenAIResponsesStreamedResponse( model_request_parameters=model_request_parameters, _model_name=first_chunk.response.model, @@ -398,6 +401,7 @@ async def request_stream( # pylint: disable=unused-argument if first_chunk.response.created_at else None ), + _tool_call_ids_are_response_scoped=tool_call_ids_are_response_scoped, ) @staticmethod diff --git a/src/utils/agents/error_handler.py b/src/utils/agents/error_handler.py index aeeddec0c..dd0d8de99 100644 --- a/src/utils/agents/error_handler.py +++ b/src/utils/agents/error_handler.py @@ -2,7 +2,7 @@ from typing import TypeAlias -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pydantic_ai.exceptions import ( AgentRunError, ContentFilterError, @@ -26,9 +26,7 @@ is_context_length_error, ) -AgentInferenceError: TypeAlias = ( - AgentRunError | APIStatusError | APIConnectionError | RuntimeError -) +AgentInferenceError: TypeAlias = AgentRunError | ApiException | RuntimeError logger = get_logger(__name__) @@ -53,13 +51,13 @@ def map_agent_inference_error( 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: + case ApiException() as connection_exc if not connection_exc.status: return ServiceUnavailableResponse( backend_name="OGX", cause=str(connection_exc), ) + case ApiException() as status_exc: + return handle_known_apistatus_errors(status_exc, model_id) case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)): return PromptTooLongResponse(model=model_id) case _: diff --git a/src/utils/agents/query.py b/src/utils/agents/query.py index 044cf67cc..e1fe8d0c7 100644 --- a/src/utils/agents/query.py +++ b/src/utils/agents/query.py @@ -6,7 +6,7 @@ from typing import Optional, TypeAlias, cast from fastapi import HTTPException -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from pydantic_ai.exceptions import ( AgentRunError, ) @@ -46,9 +46,7 @@ logger = get_logger(__name__) -AgentInferenceError: TypeAlias = ( - AgentRunError | APIStatusError | APIConnectionError | RuntimeError -) +AgentInferenceError: TypeAlias = AgentRunError | ApiException | RuntimeError class AgentFinishReason(str, Enum): @@ -259,7 +257,7 @@ async def retrieve_agent_response( else: prompt = cast(str, responses_params.input) run_result = await agent.run(prompt) - except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as exc: + except (AgentRunError, ApiException, RuntimeError) as exc: response = map_agent_inference_error(exc, responses_params.model) raise HTTPException(**response.model_dump()) from exc diff --git a/src/utils/agents/streaming.py b/src/utils/agents/streaming.py index e03b2326b..90fd0ecb4 100644 --- a/src/utils/agents/streaming.py +++ b/src/utils/agents/streaming.py @@ -11,7 +11,7 @@ from typing import Any, Final, Optional, TypeAlias, cast from fastapi import HTTPException -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pydantic_ai import Agent, AgentRunError, AgentRunResultEvent, ToolReturnPart from pydantic_ai.messages import ( AgentStreamEvent, @@ -148,7 +148,7 @@ async def retrieve_agent_response_generator( ), turn_summary, ) - except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as exc: + except (AgentRunError, ApiException, RuntimeError) as exc: response = map_agent_inference_error(exc, responses_params.model) raise HTTPException(**response.model_dump()) from exc @@ -205,7 +205,7 @@ async def generate_agent_response( stream_completed = True - except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as exc: + except (AgentRunError, ApiException, RuntimeError) as exc: error_response = map_agent_inference_error(exc, responses_params.model) yield serialize_event( ErrorStreamPayload.from_error_response(error_response), diff --git a/src/utils/builtin_tools.py b/src/utils/builtin_tools.py index 53055d8a3..e08bfa7df 100644 --- a/src/utils/builtin_tools.py +++ b/src/utils/builtin_tools.py @@ -5,7 +5,7 @@ from typing import Final from fastapi import HTTPException -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from ogx_client.models.provider_info import ProviderInfo from log import get_logger @@ -88,10 +88,10 @@ async def get_file_search_tools( """ try: providers = await client.providers.list() - except APIStatusError as exc: - logger.warning("Unable to list providers for file-search tools: %s", exc) - return [] - except APIConnectionError as e: + except ApiException as e: + if e.status: + logger.warning("Unable to list providers for file-search tools: %s", e) + return [] logger.error("Unable to connect to OGX: %s", e) response = ServiceUnavailableResponse( backend_name="OGX", cause=str(e) diff --git a/src/utils/conversations.py b/src/utils/conversations.py index 87eaa15e7..44357584c 100644 --- a/src/utils/conversations.py +++ b/src/utils/conversations.py @@ -7,7 +7,7 @@ from fastapi import HTTPException from ogx_api import OpenAIResponseOutput -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from ogx_client.models.add_items_request import AddItemsRequest from ogx_client.models.open_ai_response_input_function_tool_call_output import ( OpenAIResponseInputFunctionToolCallOutput as FunctionCallOutput, @@ -551,13 +551,14 @@ async def append_turn_items_to_conversation( conversation_id, add_items_request=build_add_items_request(items), ) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e @@ -589,13 +590,14 @@ async def get_all_conversation_items( has_more = page.has_more after = page.last_id return items - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e @@ -633,12 +635,13 @@ async def append_turn_to_conversation( ] ), ) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e diff --git a/src/utils/llama_stack_version.py b/src/utils/llama_stack_version.py index fb8598178..8d89040a3 100644 --- a/src/utils/llama_stack_version.py +++ b/src/utils/llama_stack_version.py @@ -4,7 +4,7 @@ import re from typing import Optional -from ogx_client import APIConnectionError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from semver import Version from constants import ( @@ -42,7 +42,7 @@ async def check_llama_stack_version( retry_delay: Delay in seconds between retry attempts. Raises: - APIConnectionError: If Llama Stack is unreachable after all retries. + ApiException: If Llama Stack is unreachable after all retries. InvalidLlamaStackVersionException: If the detected version is outside the supported range or cannot be parsed. """ @@ -58,7 +58,7 @@ async def check_llama_stack_version( MAXIMAL_SUPPORTED_LLAMA_STACK_VERSION, ) return version_info.version - except APIConnectionError: + except ApiException: if attempt == max_retries - 1: raise logger.warning( diff --git a/src/utils/query.py b/src/utils/query.py index c7cbab6ad..f928b5c90 100644 --- a/src/utils/query.py +++ b/src/utils/query.py @@ -6,9 +6,7 @@ import psycopg2 from fastapi import HTTPException -from ogx_client import ( - APIStatusError as LLSApiStatusError, -) +from ogx_client import ApiException from openai._exceptions import APIStatusError as OpenAIAPIStatusError from pydantic_ai.messages import ImageUrl, UserContent from sqlalchemy import func @@ -540,7 +538,7 @@ def normalize_vertex_ai_model_id(model_id: str) -> str: def handle_known_apistatus_errors( - error: LLSApiStatusError | OpenAIAPIStatusError, model_id: str + error: ApiException | OpenAIAPIStatusError, model_id: str ) -> AbstractErrorResponse: """Handle known API status errors from both Llama Stack and OpenAI. @@ -554,6 +552,7 @@ def handle_known_apistatus_errors( error_message = getattr(error, "message", str(error)) if is_context_length_error(error_message): return PromptTooLongResponse(model=model_id) - if error.status_code == 429: + status = error.status if isinstance(error, ApiException) else error.status_code + if status == 429: return QuotaExceededResponse.model(model_id) return InternalServerErrorResponse.generic() diff --git a/src/utils/responses.py b/src/utils/responses.py index c6c95d88f..2cca45c8c 100644 --- a/src/utils/responses.py +++ b/src/utils/responses.py @@ -78,7 +78,7 @@ from ogx_api.openai_responses import ( OpenAIResponseUsageOutputTokensDetails as UsageOutputTokensDetails, ) -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient import constants from configuration import configuration @@ -154,13 +154,14 @@ async def get_vector_store_ids( try: vector_stores = await client.vector_stores.list() return [vector_store.id for vector_store in vector_stores] - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e @@ -192,13 +193,14 @@ async def get_topic_summary( # pylint: disable=too-many-nested-blocks store=False, # Don't store topic summary requests ), ) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = handle_known_apistatus_errors(e, model_id) raise HTTPException(**error_response.model_dump()) from e @@ -389,13 +391,13 @@ async def prepare_responses_params( # pylint: disable=too-many-arguments,too-ma logger.debug("No conversation_id provided, creating new conversation") try: conversation = await client.conversations.create(metadata={}) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e @@ -1340,15 +1342,15 @@ async def check_model_configured( ) and model.identifier == model_id.removeprefix("watsonx/"): return True return False - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e response = InternalServerErrorResponse.generic() raise HTTPException(**response.model_dump()) from e - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e async def select_model_for_responses( @@ -1396,13 +1398,14 @@ async def select_model_for_responses( # 3. Fetch models list and select the first LLM model (model_type="llm") try: models = parse_model_list_response(await client.openai.list()) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e @@ -1612,13 +1615,14 @@ async def create_new_conversation( try: conversation = await client.conversations.create(metadata={}) return conversation.id - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="OGX", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except APIStatusError as e: + except ApiException as e: + if not e.status: + error_response = ServiceUnavailableResponse( + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e diff --git a/src/utils/shields.py b/src/utils/shields.py index dd72da72c..ecce8f49d 100644 --- a/src/utils/shields.py +++ b/src/utils/shields.py @@ -92,9 +92,9 @@ async def run_shield_moderation_v2( try: shield_result = await shield.run(input_text) - # APIConnectionError and APIStatusError from ogx should not be raised from model_request, + # ApiException 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 + # APIStatusError. 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") diff --git a/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-prow.yaml b/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-prow.yaml index 271304cbd..d5b4b4622 100644 --- a/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-prow.yaml +++ b/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-prow.yaml @@ -157,12 +157,12 @@ spec: if [[ -f "$ENRICHED_CONFIG" ]] && [[ "$ENRICHMENT_FAILED" -eq 0 ]]; then echo "Using enriched config: $ENRICHED_CONFIG" restore_rag_seed - exec ogx stack run "$ENRICHED_CONFIG" + exec ogx stack run --insecure "$ENRICHED_CONFIG" fi fi echo "Using original config: $INPUT_CONFIG" restore_rag_seed - exec ogx stack run "$INPUT_CONFIG" + exec ogx stack run --insecure "$INPUT_CONFIG" ports: - containerPort: 8321 readinessProbe: diff --git a/tests/e2e/utils/llama_stack_utils.py b/tests/e2e/utils/llama_stack_utils.py index bfb7d4fe6..d86207b01 100644 --- a/tests/e2e/utils/llama_stack_utils.py +++ b/tests/e2e/utils/llama_stack_utils.py @@ -13,8 +13,7 @@ from typing import Optional from ogx_client import ( - APIConnectionError, - APIStatusError, + ApiException, AsyncOgxClient, ) @@ -62,11 +61,11 @@ async def _unregister_shield_async(identifier: str) -> Optional[tuple[str, str]] return None try: await client.shields.delete(identifier) - except APIConnectionError: - raise - except APIStatusError as e: + except ApiException as e: + if not e.status: + raise # 400 "not found": shield already absent, scenario can proceed - if e.status_code == 400 and "not found" in str(e).lower(): + if e.status == 400 and "not found" in str(e).lower(): return None raise if provider_id is not None and provider_shield_id is not None: diff --git a/tests/integration/endpoints/test_conversations_v1_integration.py b/tests/integration/endpoints/test_conversations_v1_integration.py index b0e94f8f0..2a2f05747 100644 --- a/tests/integration/endpoints/test_conversations_v1_integration.py +++ b/tests/integration/endpoints/test_conversations_v1_integration.py @@ -8,7 +8,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pytest_mock import AsyncMockType, MockerFixture from sqlalchemy.orm import Session @@ -320,7 +320,6 @@ async def test_conversation_error_handling( # pylint: disable=too-many-locals non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, - mocker: MockerFixture, ) -> None: """Data-driven test for conversation endpoint error handling. @@ -336,7 +335,6 @@ async def test_conversation_error_handling( # pylint: disable=too-many-locals non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session - mocker: pytest-mock fixture """ _ = test_config @@ -365,13 +363,9 @@ async def test_conversation_error_handling( # pylint: disable=too-many-locals mock_method = getattr(mock_method, attr) if error_type == "connection": - mock_method.side_effect = APIConnectionError(request=mocker.Mock()) + mock_method.side_effect = ApiException(status=None) elif error_type == "api_status": - mock_method.side_effect = APIStatusError( - message="Server error", - response=mocker.Mock(status_code=500), - body=None, - ) + mock_method.side_effect = ApiException(status=500, reason="Server error") # Call the appropriate endpoint and expect error with pytest.raises(HTTPException) as exc_info: @@ -654,7 +648,6 @@ async def test_delete_conversation_handles_not_found_in_llama_stack( non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, - mocker: MockerFixture, ) -> None: """Test that delete conversation handles not found in Llama Stack gracefully. @@ -669,7 +662,6 @@ async def test_delete_conversation_handles_not_found_in_llama_stack( non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session - mocker: pytest-mock fixture """ _ = test_config @@ -688,10 +680,8 @@ async def test_delete_conversation_handles_not_found_in_llama_stack( patch_db_session.commit() # Configure mock to raise not found error - mock_ogx_client.conversations.delete.side_effect = APIStatusError( - message="Not found", - response=mocker.Mock(status_code=404), - body=None, + mock_ogx_client.conversations.delete.side_effect = ApiException( + status=404, reason="Not found" ) response = await delete_conversation_endpoint_handler( diff --git a/tests/integration/endpoints/test_health_integration.py b/tests/integration/endpoints/test_health_integration.py index 4a0dafea3..9f2c3ce94 100644 --- a/tests/integration/endpoints/test_health_integration.py +++ b/tests/integration/endpoints/test_health_integration.py @@ -139,7 +139,7 @@ async def test_health_readiness_client_error( This integration test verifies: - RuntimeError from uninitialized client is NOT caught by the endpoint - Error propagates from the endpoint handler (desired behavior) - - The endpoint does not catch RuntimeError, only APIConnectionError + - The endpoint does not catch RuntimeError, only ApiException Parameters: ---------- diff --git a/tests/integration/endpoints/test_info_integration.py b/tests/integration/endpoints/test_info_integration.py index 44a7681d4..20b503dc7 100644 --- a/tests/integration/endpoints/test_info_integration.py +++ b/tests/integration/endpoints/test_info_integration.py @@ -5,7 +5,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError +from ogx_client import ApiException from ogx_client.models.version_info import VersionInfo from pytest_mock import AsyncMockType, MockerFixture @@ -92,7 +92,6 @@ async def test_info_endpoint_handles_connection_error( mock_ogx_client: AsyncMockType, test_request: Request, test_auth: AuthTuple, - mocker: MockerFixture, ) -> None: """Test that info endpoint properly handles Llama Stack connection errors. @@ -107,14 +106,11 @@ async def test_info_endpoint_handles_connection_error( mock_ogx_client: Mocked Llama Stack client test_request: FastAPI request test_auth: noop authentication tuple - mocker: pytest-mock fixture for creating mocks """ # test_config fixture loads configuration, which is required for the endpoint _ = test_config # Configure mock to raise connection error - mock_ogx_client.inspect.version.side_effect = APIConnectionError( - request=mocker.Mock() - ) + mock_ogx_client.inspect.version.side_effect = ApiException(status=None) # Verify that HTTPException is raised with pytest.raises(HTTPException) as exc_info: diff --git a/tests/integration/endpoints/test_model_list.py b/tests/integration/endpoints/test_model_list.py index ef1f95a51..47df10782 100644 --- a/tests/integration/endpoints/test_model_list.py +++ b/tests/integration/endpoints/test_model_list.py @@ -6,7 +6,7 @@ import pytest from fastapi import Request from fastapi.exceptions import HTTPException -from ogx_client import APIConnectionError +from ogx_client import ApiException from pytest_mock import AsyncMockType, MockerFixture from app.endpoints.models import models_endpoint_handler @@ -78,7 +78,7 @@ def mock_ogx_client_failing_fixture( mock_client = mocker.AsyncMock() - mock_client.openai.list.side_effect = APIConnectionError(request=mocker.Mock()) + mock_client.openai.list.side_effect = ApiException(status=None) # Create a mock holder instance mock_holder_instance = mock_holder_class.return_value @@ -187,14 +187,14 @@ async def test_models_list_on_api_connection_error( Parameters: ---------- test_config: Test configuration - mock_ogx_client_failing: Mocked Llama Stack client that raises APIConnectionError + mock_ogx_client_failing: Mocked Llama Stack client that raises ApiException test_request: FastAPI request test_auth: noop authentication tuple """ _ = test_config _ = mock_ogx_client_failing - # we should catch HTTPException, not APIConnectionError! + # we should catch HTTPException, not ApiException! with pytest.raises(HTTPException) as exc_info: await models_endpoint_handler( request=test_request, diff --git a/tests/integration/endpoints/test_query_integration.py b/tests/integration/endpoints/test_query_integration.py index 34c29e7a2..e2c4fe46c 100644 --- a/tests/integration/endpoints/test_query_integration.py +++ b/tests/integration/endpoints/test_query_integration.py @@ -6,7 +6,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError +from ogx_client import ApiException from pytest_mock import AsyncMockType, MockerFixture from sqlalchemy.orm import Session @@ -97,7 +97,6 @@ async def test_query_v2_endpoint_handles_connection_error( mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, - mocker: MockerFixture, ) -> None: """Test that query v2 endpoint properly handles Llama Stack connection errors. @@ -113,7 +112,6 @@ async def test_query_v2_endpoint_handles_connection_error( mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple - mocker: pytest-mock fixture Returns: ------- @@ -123,7 +121,7 @@ async def test_query_v2_endpoint_handles_connection_error( _ = mock_ogx_client _ = mock_query_agent - mock_query_agent.run.side_effect = APIConnectionError(request=mocker.Mock()) + mock_query_agent.run.side_effect = ApiException(status=None) query_request = QueryRequest(query="What is Ansible?") diff --git a/tests/integration/endpoints/test_rlsapi_v1_integration.py b/tests/integration/endpoints/test_rlsapi_v1_integration.py index 3fdf37e86..a427090a3 100644 --- a/tests/integration/endpoints/test_rlsapi_v1_integration.py +++ b/tests/integration/endpoints/test_rlsapi_v1_integration.py @@ -14,7 +14,7 @@ import pytest from fastapi import HTTPException, status from fastapi.testclient import TestClient -from ogx_client import APIConnectionError +from ogx_client import ApiException from pytest_mock import MockerFixture import constants @@ -266,9 +266,7 @@ async def test_rlsapi_v1_infer_connection_error_returns_503( _ = rlsapi_config mock_responses = mocker.Mock() - mock_responses.create = mocker.AsyncMock( - side_effect=APIConnectionError(request=mocker.Mock()) - ) + mock_responses.create = mocker.AsyncMock(side_effect=ApiException(status=None)) mock_client = mocker.Mock() mock_client.responses = mock_responses diff --git a/tests/unit/app/endpoints/test_a2a.py b/tests/unit/app/endpoints/test_a2a.py index b7bf30bb6..ddefd5b1f 100644 --- a/tests/unit/app/endpoints/test_a2a.py +++ b/tests/unit/app/endpoints/test_a2a.py @@ -6,7 +6,6 @@ from typing import Any -import httpx import pytest from a2a.server.agent_execution import RequestContext from a2a.server.events import EventQueue @@ -22,7 +21,7 @@ ) from a2a.utils import new_agent_text_message from fastapi import HTTPException, Request -from ogx_client import APIConnectionError +from ogx_client import ApiException from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError from pydantic_ai.messages import ( @@ -37,8 +36,6 @@ from pydantic_ai.messages import TextPart as PydanticTextPart from pytest_mock import MockerFixture -from tests.unit.conftest import make_openai_model, make_openai_models_list_response - from app.endpoints.a2a import ( A2AAgentExecutor, TaskResultAggregator, @@ -51,6 +48,7 @@ ) from configuration import AppConfig from models.config import Action +from tests.unit.conftest import make_openai_model, make_openai_models_list_response # User ID must be proper UUID MOCK_AUTH = ( @@ -687,7 +685,7 @@ async def test_process_task_streaming_handles_api_connection_error_on_models_lis mocker: MockerFixture, setup_configuration: AppConfig, # pylint: disable=unused-argument ) -> None: - """Test _process_task_streaming handles APIConnectionError from openai.list().""" + """Test _process_task_streaming handles ApiException from openai.list().""" executor = A2AAgentExecutor(auth_token="test-token") # Mock the context with valid input @@ -717,19 +715,16 @@ async def test_process_task_streaming_handles_api_connection_error_on_models_lis "app.endpoints.a2a._get_context_store", return_value=mock_context_store ) - # Mock the client to raise APIConnectionError on openai.list() + # Mock the client to raise ApiException on openai.list() mock_client = mocker.AsyncMock() - # Create a mock httpx.Request for APIConnectionError - mock_request = httpx.Request("GET", "http://test-llama-stack/models") - mock_client.openai.list.side_effect = APIConnectionError( - message="Connection refused: unable to reach Llama Stack", - request=mock_request, + mock_client.openai.list.side_effect = ApiException( + status=None, reason="Connection refused: unable to reach Llama Stack" ) mocker.patch( "app.endpoints.a2a.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client - # prepare_responses_params raises HTTPException when APIConnectionError occurs + # prepare_responses_params raises HTTPException when ApiException occurs with pytest.raises(HTTPException) as exc_info: await executor._process_task_streaming( context, task_updater, context.task_id, context.context_id @@ -745,7 +740,7 @@ async def test_process_task_streaming_handles_api_connection_error( # pylint: d mocker: MockerFixture, setup_configuration: AppConfig, # pylint: disable=unused-argument ) -> None: - """Test _process_task_streaming handles APIConnectionError during agent run.""" + """Test _process_task_streaming handles ApiException during agent run.""" executor = A2AAgentExecutor(auth_token="test-token") # Mock the context with valid input @@ -807,13 +802,11 @@ async def test_process_task_streaming_handles_api_connection_error( # pylint: d ) # Mock build_agent to return an agent whose run_stream_events raises - mock_request = httpx.Request("POST", "http://test-llama-stack/responses") mock_agent = mocker.MagicMock() mock_stream_ctx = mocker.AsyncMock() mock_stream_ctx.__aenter__ = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection timeout during streaming", - request=mock_request, + side_effect=ApiException( + status=None, reason="Connection timeout during streaming" ) ) mock_agent.run_stream_events.return_value = mock_stream_ctx diff --git a/tests/unit/app/endpoints/test_conversations.py b/tests/unit/app/endpoints/test_conversations.py index 07dbd7e25..ddb2333c6 100644 --- a/tests/unit/app/endpoints/test_conversations.py +++ b/tests/unit/app/endpoints/test_conversations.py @@ -8,7 +8,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError, APIStatusError, NotFoundError +from ogx_client import ApiException, NotFoundError from pytest_mock import MockerFixture, MockType from sqlalchemy.exc import SQLAlchemyError @@ -529,9 +529,7 @@ async def test_llama_stack_connection_error( mock_database_session(mocker, query_result=[mock_conversation], db_turns=[]) mock_client = mocker.AsyncMock() - mock_client.items.list.side_effect = APIConnectionError( - request=None # type: ignore[arg-type] - ) + mock_client.items.list.side_effect = ApiException(status=None) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) @@ -579,9 +577,7 @@ async def test_llama_stack_not_found_error( mock_client = mocker.AsyncMock() mock_client.items.list.side_effect = NotFoundError( - message="Conversation not found", - response=mocker.Mock(request=None), - body=None, + status=404, reason="Conversation not found" ) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" @@ -686,9 +682,7 @@ async def test_get_others_conversations_allowed_for_authorized_user( mock_item2.content = "Hi there!" mock_items_response.data = [mock_item1, mock_item2] mock_items_response.has_more = False - mock_client.items.list = mocker.AsyncMock( - return_value=mock_items_response - ) + mock_client.items.list = mocker.AsyncMock(return_value=mock_items_response) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" @@ -818,9 +812,7 @@ async def test_no_items_found_in_get_conversation( mock_items_response = mocker.Mock() mock_items_response.data = [] mock_items_response.has_more = False - mock_client.items.list = mocker.AsyncMock( - return_value=mock_items_response - ) + mock_client.items.list = mocker.AsyncMock(return_value=mock_items_response) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) @@ -847,9 +839,9 @@ async def test_api_status_error_in_get_conversation( dummy_request: Request, mock_conversation: MockType, ) -> None: - """Test when APIStatusError is raised during conversation retrieval. + """Test when ApiException is raised during conversation retrieval. - get_all_conversation_items maps APIStatusError to HTTP 500. + get_all_conversation_items maps ApiException to HTTP 500. """ mock_authorization_resolvers(mocker) mocker.patch( @@ -864,10 +856,8 @@ async def test_api_status_error_in_get_conversation( mock_database_session(mocker, db_turns=[]) mock_client = mocker.AsyncMock() - mock_client.items.list.side_effect = APIStatusError( - message="Conversation not found", - response=mocker.Mock(status_code=404, request=None), - body=None, + mock_client.items.list.side_effect = ApiException( + status=404, reason="Conversation not found" ) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" @@ -1109,9 +1099,7 @@ async def test_llama_stack_connection_error( ) mock_client = mocker.AsyncMock() - mock_client.conversations.delete.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.conversations.delete.side_effect = ApiException(status=None) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) @@ -1151,10 +1139,8 @@ async def test_llama_stack_not_found_error( ) mock_client = mocker.AsyncMock() - mock_client.conversations.delete.side_effect = APIStatusError( - message="Conversation not found", - response=mocker.Mock(status_code=404, request=None), - body=None, + mock_client.conversations.delete.side_effect = ApiException( + status=404, reason="Conversation not found" ) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" @@ -2031,11 +2017,9 @@ async def test_llama_stack_connection_error_in_update( return_value=mock_conversation, ) - # Mock AsyncOgxClientHolder to raise APIConnectionError + # Mock AsyncOgxClientHolder to raise ApiException mock_client = mocker.AsyncMock() - mock_client.conversations.update.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.conversations.update.side_effect = ApiException(status=None) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) @@ -2077,12 +2061,10 @@ async def test_llama_stack_not_found_error_in_update( return_value=mock_conversation, ) - # Mock AsyncOgxClientHolder to raise APIStatusError + # Mock AsyncOgxClientHolder to raise ApiException mock_client = mocker.AsyncMock() - mock_client.conversations.update.side_effect = APIStatusError( - message="Conversation not found", - response=mocker.Mock(status_code=404, request=None), - body=None, + mock_client.conversations.update.side_effect = ApiException( + status=404, reason="Conversation not found" ) mock_client_holder = mocker.patch( "app.endpoints.conversations_v1.AsyncOgxClientHolder" diff --git a/tests/unit/app/endpoints/test_health.py b/tests/unit/app/endpoints/test_health.py index d32561153..e8d2271cd 100644 --- a/tests/unit/app/endpoints/test_health.py +++ b/tests/unit/app/endpoints/test_health.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from ogx_client import APIConnectionError +from ogx_client import ApiException from pytest_mock import MockerFixture from app.endpoints.health import ( @@ -267,16 +267,15 @@ async def test_get_providers_health_statuses_connection_error( mock_lsc = mocker.patch("client.AsyncOgxClientHolder.get_client") # Mock get_ogx_client to raise an exception - mock_lsc.side_effect = APIConnectionError(request=mocker.Mock()) + mock_lsc.side_effect = ApiException(status=None, reason="Connection error.") result = await get_providers_health_statuses() assert len(result) == 1 assert result[0].provider_id == "unknown" assert result[0].status == HealthStatus.ERROR.value - assert ( - result[0].message == "Failed to initialize health check: Connection error." - ) + assert result[0].message.startswith("Failed to initialize health check:") + assert "Connection error." in (result[0].message or "") class TestCheckDefaultModelAvailable: diff --git a/tests/unit/app/endpoints/test_info.py b/tests/unit/app/endpoints/test_info.py index 8ef603252..ca8674fb4 100644 --- a/tests/unit/app/endpoints/test_info.py +++ b/tests/unit/app/endpoints/test_info.py @@ -4,7 +4,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError +from ogx_client import ApiException from ogx_client.models.version_info import VersionInfo from pytest_mock import MockerFixture @@ -84,7 +84,7 @@ async def test_info_endpoint_connection_error(mocker: MockerFixture) -> None: Sets up application configuration and patches the LlamaStack client so that calling its version inspection raises an - APIConnectionError, then asserts the raised HTTPException has + ApiException, then asserts the raised HTTPException has status code 503 and a detail payload containing a "response" of "Service unavailable" and a "cause" that includes "Unable to connect to Llama Stack". @@ -119,7 +119,7 @@ async def test_info_endpoint_connection_error(mocker: MockerFixture) -> None: # Mock the LlamaStack client mock_client = mocker.AsyncMock() - mock_client.inspect.version.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.inspect.version.side_effect = ApiException(status=None) # type: ignore mock_lsc = mocker.patch("client.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() diff --git a/tests/unit/app/endpoints/test_models.py b/tests/unit/app/endpoints/test_models.py index 231ee7fe0..923ff54a0 100644 --- a/tests/unit/app/endpoints/test_models.py +++ b/tests/unit/app/endpoints/test_models.py @@ -4,7 +4,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError +from ogx_client import ApiException from pytest_mock import MockerFixture from pytest_subtests import SubTests @@ -56,7 +56,7 @@ async def test_models_endpoint_handler_configuration_loaded( Loads an AppConfig from a test dictionary, patches the endpoint's configuration and AsyncOgxClientHolder so that get_client raises - APIConnectionError, issues a request with an authorization header, and + ApiException, issues a request with an authorization header, and asserts that calling the handler raises an HTTPException with status 503 and a detail response of "Unable to connect to OGX". """ @@ -90,9 +90,7 @@ async def test_models_endpoint_handler_configuration_loaded( mocker.patch("app.endpoints.models.configuration", cfg) mock_client_holder = mocker.patch("app.endpoints.models.AsyncOgxClientHolder") - mock_client_holder.return_value.get_client.side_effect = APIConnectionError( - request=mocker.Mock() - ) + mock_client_holder.return_value.get_client.side_effect = ApiException(status=None) request = Request( scope={ @@ -429,10 +427,10 @@ async def test_models_endpoint_llama_stack_connection_error( "authentication": {"module": "noop"}, } - # mock AsyncOgxClientHolder to raise APIConnectionError + # mock AsyncOgxClientHolder to raise ApiException # when openai.list() method is called mock_client = mocker.AsyncMock() - mock_client.openai.list.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.openai.list.side_effect = ApiException(status=None) # type: ignore mock_client_holder = mocker.patch("app.endpoints.models.AsyncOgxClientHolder") mock_client_holder.return_value.get_client.return_value = mock_client diff --git a/tests/unit/app/endpoints/test_prompts.py b/tests/unit/app/endpoints/test_prompts.py index 1c82c756f..c751b85e7 100644 --- a/tests/unit/app/endpoints/test_prompts.py +++ b/tests/unit/app/endpoints/test_prompts.py @@ -4,7 +4,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError, BadRequestError +from ogx_client import ApiException, BadRequestError from ogx_client.models.prompt import Prompt from pytest_mock import MockerFixture @@ -199,9 +199,7 @@ async def test_delete_prompt_not_found_returns_body( _, mock_prompts = prompts_client_mocks mock_response = mocker.Mock() mock_response.request = mocker.Mock() - mock_prompts.delete.side_effect = BadRequestError( - message="not found", response=mock_response, body=None - ) + mock_prompts.delete.side_effect = BadRequestError(status=400, reason="not found") result = await delete_prompt_handler( request=prompts_http_request, @@ -237,9 +235,9 @@ async def test_get_prompt_api_connection_error( prompts_client_mocks: tuple[Any, Any], prompts_http_request: Request, ) -> None: - """get_prompt maps APIConnectionError to 503.""" + """get_prompt maps ApiException to 503.""" _, mock_prompts = prompts_client_mocks - mock_prompts.retrieve.side_effect = APIConnectionError(request=None) # type: ignore + mock_prompts.retrieve.side_effect = ApiException(status=None) # type: ignore with pytest.raises(HTTPException) as exc_info: await get_prompt_handler( @@ -261,9 +259,7 @@ async def test_get_prompt_bad_request_maps_to_404( _, mock_prompts = prompts_client_mocks mock_response = mocker.Mock() mock_response.request = mocker.Mock() - mock_prompts.retrieve.side_effect = BadRequestError( - message="not found", response=mock_response, body=None - ) + mock_prompts.retrieve.side_effect = BadRequestError(status=400, reason="not found") with pytest.raises(HTTPException) as exc_info: await get_prompt_handler( @@ -292,7 +288,7 @@ async def test_update_prompt_bad_request_maps_to_404( mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_prompts.update.side_effect = BadRequestError( - message="invalid version", response=mock_response, body=None + status=400, reason="invalid version" ) body = PromptUpdateRequest( diff --git a/tests/unit/app/endpoints/test_providers.py b/tests/unit/app/endpoints/test_providers.py index 404e5bc1f..f58daccf6 100644 --- a/tests/unit/app/endpoints/test_providers.py +++ b/tests/unit/app/endpoints/test_providers.py @@ -2,7 +2,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError, BadRequestError +from ogx_client import ApiException, BadRequestError from ogx_client.models.provider_info import ProviderInfo from pytest_mock import MockerFixture @@ -44,7 +44,7 @@ async def test_providers_endpoint_connection_error( mocker.patch( "app.endpoints.providers.AsyncOgxClientHolder" - ).return_value.get_client.side_effect = APIConnectionError(request=mocker.Mock()) + ).return_value.get_client.side_effect = ApiException(status=None) request = Request(scope={"type": "http"}) @@ -117,11 +117,7 @@ async def test_get_provider_not_found( mock_client_holder = mocker.patch("app.endpoints.providers.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() mock_client.providers.retrieve = mocker.AsyncMock( - side_effect=BadRequestError( - message="Provider not found", - response=mocker.Mock(request=None), - body=None, - ) + side_effect=BadRequestError(status=400, reason="Provider not found") ) # type: ignore mock_client_holder.return_value.get_client.return_value = mock_client @@ -141,6 +137,34 @@ async def test_get_provider_not_found( assert "Provider with ID openai does not exist" in detail["cause"] # type: ignore +@pytest.mark.asyncio +async def test_get_provider_not_found_library_mode( + mocker: MockerFixture, minimal_config: AppConfig +) -> None: + """Test that /providers/{provider_id} maps library-mode ValueError to HTTP 404.""" + mocker.patch("app.endpoints.providers.configuration", minimal_config) + + mock_client_holder = mocker.patch("app.endpoints.providers.AsyncOgxClientHolder") + mock_client = mocker.AsyncMock() + mock_client.providers.retrieve = mocker.AsyncMock( + side_effect=ValueError("Provider faisss not found") + ) + mock_client_holder.return_value.get_client.return_value = mock_client + + request = Request(scope={"type": "http"}) + auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") + + with pytest.raises(HTTPException) as e: + await get_provider_endpoint_handler( + request=request, provider_id="faisss", auth=auth + ) + assert e.value.status_code == status.HTTP_404_NOT_FOUND + detail = e.value.detail + assert isinstance(detail, dict) + assert "not found" in detail["response"] # type: ignore + assert "Provider with ID faisss does not exist" in detail["cause"] # type: ignore + + @pytest.mark.asyncio async def test_get_provider_success( mocker: MockerFixture, minimal_config: AppConfig @@ -183,7 +207,7 @@ async def test_get_provider_connection_error( mocker.patch( "app.endpoints.providers.AsyncOgxClientHolder" - ).return_value.get_client.side_effect = APIConnectionError(request=mocker.Mock()) + ).return_value.get_client.side_effect = ApiException(status=None) request = Request(scope={"type": "http"}) diff --git a/tests/unit/app/endpoints/test_rags.py b/tests/unit/app/endpoints/test_rags.py index 15c2d8350..604223aa6 100644 --- a/tests/unit/app/endpoints/test_rags.py +++ b/tests/unit/app/endpoints/test_rags.py @@ -6,7 +6,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError, BadRequestError +from ogx_client import ApiException, BadRequestError from pytest_mock import MockerFixture from app.endpoints.rags import ( @@ -45,7 +45,7 @@ async def test_rags_endpoint_connection_error( """Test that /rags endpoint raises HTTP 503 if Llama Stack connection fails.""" mocker.patch("app.endpoints.rags.configuration", minimal_config) mock_client = mocker.AsyncMock() - mock_client.vector_stores.list.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.vector_stores.list.side_effect = ApiException(status=None) # type: ignore mocker.patch( "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client @@ -151,11 +151,7 @@ async def test_rag_info_endpoint_rag_not_found( mocker.patch("app.endpoints.rags.configuration", minimal_config) mock_client = mocker.AsyncMock() mock_client.vector_stores.retrieve = mocker.AsyncMock( - side_effect=BadRequestError( - message="RAG not found", - response=mocker.Mock(request=None), - body=None, - ) + side_effect=BadRequestError(status=400, reason="RAG not found") ) # type: ignore mocker.patch( "app.endpoints.rags.AsyncOgxClientHolder" @@ -182,9 +178,7 @@ async def test_rag_info_endpoint_connection_error( """Test that /rags/{rag_id} endpoint raises HTTP 503 if Llama Stack connection fails.""" mocker.patch("app.endpoints.rags.configuration", minimal_config) mock_client = mocker.AsyncMock() - mock_client.vector_stores.retrieve.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.vector_stores.retrieve.side_effect = ApiException(status=None) mocker.patch( "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index 4997427d7..ef7f5882a 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -16,7 +16,7 @@ from ogx_api.openai_responses import ( OpenAIResponseMessage, ) -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from pytest_mock import MockerFixture from app.endpoints.responses import ( @@ -131,6 +131,8 @@ def _patch_base(mocker: MockerFixture, config: AppConfig) -> None: def _patch_client(mocker: MockerFixture) -> Any: """Patch AsyncOgxClientHolder; return (mock_client, mock_holder).""" mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") mock_vector_stores = mocker.Mock() mock_vector_stores.list = mocker.AsyncMock(return_value=[]) mock_client.vector_stores = mock_vector_stores @@ -491,6 +493,10 @@ async def test_responses_azure_token_refresh( mock_azure.refresh_token.return_value = True mocker.patch(f"{MODULE}.AzureEntraIDManager", return_value=mock_azure) updated_client = mocker.AsyncMock(spec=AsyncOgxClient) + updated_client.attach_mock(mocker.AsyncMock(), "responses") + updated_client.attach_mock(mocker.AsyncMock(), "items") + updated_client.attach_mock(mocker.AsyncMock(), "openai") + updated_client.attach_mock(mocker.AsyncMock(), "conversations") mock_holder.update_azure_token = mocker.AsyncMock(return_value=updated_client) _patch_rag(mocker) _patch_moderation(mocker, decision="passed") @@ -747,6 +753,10 @@ async def test_handle_non_streaming_blocked_returns_refusal( """Test that blocked moderation returns response with refusal message.""" request = _request_with_model_and_conv("Bad input") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" mock_moderation.message = "Content blocked" @@ -757,6 +767,10 @@ async def test_handle_non_streaming_blocked_returns_refusal( mock_moderation.refusal_response = mock_refusal _patch_handle_non_streaming_common(mocker, minimal_config) + mocker.patch( + f"{MODULE}.append_turn_items_to_conversation", + new=mocker.AsyncMock(), + ) mock_client.items.create = mocker.AsyncMock() mock_api_response = mocker.Mock() mock_api_response.output = [mock_refusal] @@ -809,6 +823,10 @@ async def test_handle_non_streaming_success_returns_response( """Test successful handle_non_streaming_response returns ResponsesResponse.""" request = _request_with_model_and_conv("Hello") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -889,6 +907,10 @@ async def test_handle_non_streaming_with_previous_response_id_appends_turn( """Test append_turn_items_to_conversation triggers with store and previous_response_id.""" request = _request_with_previous_response_id("Hi", previous_response_id="r1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -971,6 +993,10 @@ async def test_handle_non_streaming_context_length_raises_413( """Test that RuntimeError with context_length raises 413.""" request = _request_with_model_and_conv("Long input") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( side_effect=RuntimeError("context_length exceeded") ) @@ -1007,14 +1033,15 @@ async def test_handle_non_streaming_connection_error_raises_503( minimal_config: AppConfig, mocker: MockerFixture, ) -> None: - """Test that APIConnectionError raises 503.""" + """Test that ApiException raises 503.""" request = _request_with_model_and_conv("Hi") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", - request=mocker.Mock(), - ) + side_effect=ApiException(status=None, reason="Connection failed") ) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1053,15 +1080,15 @@ async def test_handle_non_streaming_api_status_error_raises_http( minimal_config: AppConfig, mocker: MockerFixture, ) -> None: - """Test that APIStatusError is handled and re-raised as HTTPException.""" + """Test that ApiException is handled and re-raised as HTTPException.""" request = _request_with_model_and_conv("Hi") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( - side_effect=APIStatusError( - message="API error", - response=mocker.Mock(request=None), - body=None, - ) + side_effect=ApiException(status=500, reason="API error") ) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1108,6 +1135,10 @@ async def test_handle_non_streaming_runtime_error_without_context_reraises( """Test that RuntimeError without context_length is re-raised.""" request = _request_with_model_and_conv("Hi") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( side_effect=RuntimeError("Some other error") ) @@ -1149,6 +1180,10 @@ async def test_handle_streaming_blocked_returns_sse_consumes_shield_generator( """Test streaming with blocked moderation yields SSE from shield_violation_generator.""" request = _request_with_model_and_conv("Bad", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" mock_moderation.message = "Blocked" @@ -1213,6 +1248,10 @@ async def test_handle_streaming_success_returns_sse_consumes_response_generator( """Test streaming with passed moderation yields SSE from response_generator.""" request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1292,6 +1331,10 @@ async def test_handle_streaming_in_progress_chunk_sets_quotas_and_output_text( """Test in_progress chunk includes available_quotas and output_text.""" request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1379,6 +1422,10 @@ async def test_handle_streaming_builds_tool_call_summary_from_output( """Test that response output items are passed to build_tool_call_summary.""" request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1466,6 +1513,10 @@ async def test_handle_streaming_with_previous_response_id_appends_turn( "Hi", previous_response_id="r_prev" ) mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1549,6 +1600,10 @@ async def test_handle_streaming_context_length_raises_413( """Test streaming raises 413 when create raises RuntimeError context_length.""" request = _request_with_model_and_conv("Long", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( side_effect=RuntimeError("context_length exceeded") ) @@ -1583,14 +1638,15 @@ async def test_handle_streaming_connection_error_raises_503( minimal_config: AppConfig, mocker: MockerFixture, ) -> None: - """Test streaming raises 503 when create raises APIConnectionError.""" + """Test streaming raises 503 when create raises ApiException.""" request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", - request=mocker.Mock(), - ) + side_effect=ApiException(status=None, reason="Connection failed") ) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2273,6 +2329,10 @@ async def test_non_streaming_sanitizes_mcp_output_and_model( ) mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2402,7 +2462,7 @@ def _make_streaming_completed_chunk(self, mocker: MockerFixture) -> Any: return completed_chunk @pytest.mark.asyncio - async def test_streaming_sanitizes_mcp_output_model_and_instructions( + async def test_streaming_sanitizes_mcp_output_model_and_instructions( # pylint: disable=too-many-statements self, minimal_config: AppConfig, mocker: MockerFixture, @@ -2429,6 +2489,10 @@ async def test_streaming_sanitizes_mcp_output_model_and_instructions( conversation=VALID_CONV_ID_NORMALIZED, ) mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2512,7 +2576,7 @@ class TestMcpEventsFilteredUnconditionally: """Integration test: MCP events are filtered regardless of X-LCS-Merge-Server-Tools.""" @pytest.mark.asyncio - async def test_mcp_events_filtered_without_merge_server_tools_header( + async def test_mcp_events_filtered_without_merge_server_tools_header( # pylint: disable=too-many-statements self, minimal_config: AppConfig, mocker: MockerFixture, @@ -2530,6 +2594,10 @@ async def test_mcp_events_filtered_without_merge_server_tools_header( request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2627,6 +2695,10 @@ async def test_mcp_events_filtered_with_no_mcp_servers_configured( """ request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2721,6 +2793,8 @@ async def test_response_generator_records_failure_when_stream_iteration_raises( """Test that response_generator records a failure metric when the stream raises.""" request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" diff --git a/tests/unit/app/endpoints/test_responses_splunk.py b/tests/unit/app/endpoints/test_responses_splunk.py index 8d74c2495..83e824d77 100644 --- a/tests/unit/app/endpoints/test_responses_splunk.py +++ b/tests/unit/app/endpoints/test_responses_splunk.py @@ -9,8 +9,7 @@ from fastapi.responses import StreamingResponse from ogx_api import OpenAIResponseObject from ogx_api.openai_responses import OpenAIResponseMessage -from ogx_client import APIConnectionError, AsyncOgxClient -from ogx_client import APIStatusError as LLSApiStatusError +from ogx_client import ApiException, AsyncOgxClient from openai._exceptions import APIStatusError as OpenAIAPIStatusError from pytest_mock import MockerFixture @@ -231,6 +230,10 @@ async def test_non_streaming_shield_blocked( """Blocked moderation fires responses_shield_blocked telemetry.""" request = _request_with_model_and_conv("Bad input") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" @@ -294,18 +297,16 @@ async def test_non_streaming_shield_blocked( "exc_factory", [ pytest.param( - lambda m: APIConnectionError(request=m.Mock()), - id="APIConnectionError", + lambda m: ApiException(status=None), + id="ApiException-connection", ), pytest.param( lambda m: RuntimeError("context_length exceeded"), id="RuntimeError-context-length", ), pytest.param( - lambda m: LLSApiStatusError( - message="API error", response=m.Mock(request=None), body=None - ), - id="LLSApiStatusError", + lambda m: ApiException(status=500, reason="API error"), + id="ApiException-status", ), pytest.param( lambda m: OpenAIAPIStatusError( @@ -325,6 +326,10 @@ async def test_non_streaming_error_fires_telemetry( """Each error branch fires responses_error telemetry with fire_and_forget.""" request = _request_with_model_and_conv("Hello") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -378,6 +383,10 @@ async def test_non_streaming_success( """Successful non-streaming response fires responses_completed with token counts.""" request = _request_with_model_and_conv("Hello") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -468,6 +477,10 @@ async def test_streaming_shield_blocked( """Blocked moderation in streaming fires responses_shield_blocked telemetry.""" request = _request_with_model_and_conv("Bad", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" @@ -522,18 +535,16 @@ async def test_streaming_shield_blocked( "exc_factory", [ pytest.param( - lambda m: APIConnectionError(request=m.Mock()), - id="APIConnectionError", + lambda m: ApiException(status=None), + id="ApiException-connection", ), pytest.param( lambda m: RuntimeError("context_length exceeded"), id="RuntimeError-context-length", ), pytest.param( - lambda m: LLSApiStatusError( - message="API error", response=m.Mock(request=None), body=None - ), - id="LLSApiStatusError", + lambda m: ApiException(status=500, reason="API error"), + id="ApiException-status", ), pytest.param( lambda m: OpenAIAPIStatusError( @@ -553,6 +564,10 @@ async def test_streaming_error_fires_telemetry( """Each streaming error branch fires responses_error telemetry with fire_and_forget.""" request = _request_with_model_and_conv("Hello") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -606,6 +621,10 @@ async def test_streaming_success( """Successful streaming fires responses_completed after consuming the stream.""" request = _request_with_model_and_conv("Hi", model="provider/model1") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -702,6 +721,10 @@ async def test_splunk_disabled_no_background_tasks( """When background_tasks is None, queue_responses_splunk_event is called but is a no-op.""" request = _request_with_model_and_conv("Bad input") mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" diff --git a/tests/unit/app/endpoints/test_rlsapi_v1.py b/tests/unit/app/endpoints/test_rlsapi_v1.py index 68c955a71..2a8b6baba 100644 --- a/tests/unit/app/endpoints/test_rlsapi_v1.py +++ b/tests/unit/app/endpoints/test_rlsapi_v1.py @@ -13,12 +13,10 @@ import pytest from fastapi import HTTPException, status -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pydantic import ValidationError from pytest_mock import MockerFixture -from tests.unit.conftest import make_openai_model, make_openai_models_list_response - import constants from app.endpoints.rlsapi_v1 import ( AUTH_DISABLED, @@ -51,6 +49,7 @@ RedactionRule, RedactionShieldConfiguration, ) +from tests.unit.conftest import make_openai_model, make_openai_models_list_response from tests.unit.utils.auth_helpers import mock_authorization_resolvers from utils.rh_identity import get_rh_identity_context from utils.suid import check_suid @@ -180,10 +179,10 @@ def mock_model_configured_fixture(mocker: MockerFixture) -> None: @pytest.fixture(name="mock_api_connection_error") def mock_api_connection_error_fixture(mocker: MockerFixture) -> None: - """Mock responses.create() to raise APIConnectionError.""" + """Mock responses.create() to raise ApiException.""" _setup_responses_mock( mocker, - mocker.AsyncMock(side_effect=APIConnectionError(request=mocker.Mock())), + mocker.AsyncMock(side_effect=ApiException(status=None)), ) @@ -198,16 +197,14 @@ def mock_generic_runtime_error_fixture(mocker: MockerFixture) -> None: @pytest.fixture(name="mock_api_status_error_with_private_text") def mock_api_status_error_with_private_text_fixture(mocker: MockerFixture) -> None: - """Mock responses.create() to raise APIStatusError with private text.""" + """Mock responses.create() to raise ApiException with private text.""" mock_response = mocker.Mock(request=None) mock_response.status_code = 500 _setup_responses_mock( mocker, mocker.AsyncMock( - side_effect=APIStatusError( - message="Backend echoed PRIVATE prompt sk-backend-secret", - response=mock_response, - body=None, + side_effect=ApiException( + status=500, reason="Backend echoed PRIVATE prompt sk-backend-secret" ) ), ) @@ -377,7 +374,7 @@ async def test_get_default_model_id_errors( ) else: mock_client.openai.list = mocker.AsyncMock( - side_effect=APIConnectionError(request=mocker.Mock()) + side_effect=ApiException(status=None) ) mock_client_holder = mocker.Mock() @@ -402,7 +399,7 @@ async def test_config_error_503_matches_llm_error_503_shape( ) -> None: """Test that auto-discovery 503s have the same shape as LLM error 503s. - Both _get_default_model_id() no-LLM auto-discovery errors and APIConnectionError + Both _get_default_model_id() no-LLM auto-discovery errors and ApiException handlers use ServiceUnavailableResponse, producing identical detail shapes with 'response' and 'cause' keys. """ @@ -771,7 +768,7 @@ async def test_infer_api_status_error_logs_class_without_private_text( ) assert exc_info.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR - assert "APIStatusError" in caplog.text + assert "ApiException" in caplog.text assert "sk-backend-secret" not in caplog.text assert "PRIVATE prompt" not in caplog.text diff --git a/tests/unit/app/endpoints/test_vector_stores.py b/tests/unit/app/endpoints/test_vector_stores.py index f49102f37..8ed8285b1 100644 --- a/tests/unit/app/endpoints/test_vector_stores.py +++ b/tests/unit/app/endpoints/test_vector_stores.py @@ -7,7 +7,7 @@ import pytest from fastapi import HTTPException, Request, status -from ogx_client import APIConnectionError, BadRequestError +from ogx_client import ApiException, BadRequestError from pytest_mock import MockerFixture from app.endpoints.vector_stores import ( @@ -232,7 +232,7 @@ async def test_create_vector_store_connection_error(mocker: MockerFixture) -> No cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores.create.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.vector_stores.create.side_effect = ApiException(status=None) # type: ignore mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -323,7 +323,7 @@ async def test_get_vector_store_not_found(mocker: MockerFixture) -> None: mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores.retrieve.side_effect = BadRequestError( - message="Not found", response=mock_response, body=None + status=400, reason="Not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -693,7 +693,7 @@ async def test_list_vector_stores_connection_error(mocker: MockerFixture) -> Non cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores.list.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.vector_stores.list.side_effect = ApiException(status=None) # type: ignore mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -718,7 +718,7 @@ async def test_update_vector_store_connection_error(mocker: MockerFixture) -> No cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores.update.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.vector_stores.update.side_effect = ApiException(status=None) # type: ignore mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -749,7 +749,7 @@ async def test_update_vector_store_not_found(mocker: MockerFixture) -> None: mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores.update.side_effect = BadRequestError( - message="Not found", response=mock_response, body=None + status=400, reason="Not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -778,7 +778,7 @@ async def test_delete_vector_store_connection_error(mocker: MockerFixture) -> No cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores.delete.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.vector_stores.delete.side_effect = ApiException(status=None) # type: ignore mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -806,7 +806,7 @@ async def test_delete_vector_store_not_found(mocker: MockerFixture) -> None: mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores.delete.side_effect = BadRequestError( - message="Not found", response=mock_response, body=None + status=400, reason="Not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -834,7 +834,7 @@ async def test_create_file_connection_error(mocker: MockerFixture) -> None: cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.files.create.side_effect = APIConnectionError(request=None) # type: ignore + mock_client.files.create.side_effect = ApiException(status=None) # type: ignore mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -867,7 +867,7 @@ async def test_create_file_bad_request(mocker: MockerFixture) -> None: mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.files.create.side_effect = BadRequestError( - message="File too large", response=mock_response, body=None + status=400, reason="File too large" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -963,9 +963,7 @@ async def test_add_file_to_vector_store_connection_error( cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores_files.create.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.vector_stores_files.create.side_effect = ApiException(status=None) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -996,7 +994,7 @@ async def test_add_file_to_vector_store_not_found(mocker: MockerFixture) -> None mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores_files.create.side_effect = BadRequestError( - message="File not found", response=mock_response, body=None + status=400, reason="File not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -1027,9 +1025,7 @@ async def test_list_vector_store_files_connection_error( cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores_files.list.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.vector_stores_files.list.side_effect = ApiException(status=None) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -1059,7 +1055,7 @@ async def test_list_vector_store_files_not_found(mocker: MockerFixture) -> None: mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores_files.list.side_effect = BadRequestError( - message="Vector store not found", response=mock_response, body=None + status=400, reason="Vector store not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -1088,9 +1084,7 @@ async def test_get_vector_store_file_connection_error(mocker: MockerFixture) -> cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores_files.retrieve.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.vector_stores_files.retrieve.side_effect = ApiException(status=None) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -1120,7 +1114,7 @@ async def test_get_vector_store_file_not_found(mocker: MockerFixture) -> None: mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores_files.retrieve.side_effect = BadRequestError( - message="File not found", response=mock_response, body=None + status=400, reason="File not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -1150,9 +1144,7 @@ async def test_delete_vector_store_file_connection_error( cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores_files.delete.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.vector_stores_files.delete.side_effect = ApiException(status=None) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -1182,7 +1174,7 @@ async def test_delete_vector_store_file_not_found(mocker: MockerFixture) -> None mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.vector_stores_files.delete.side_effect = BadRequestError( - message="File not found", response=mock_response, body=None + status=400, reason="File not found" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" @@ -1210,9 +1202,7 @@ async def test_get_vector_store_connection_error(mocker: MockerFixture) -> None: cfg.init_from_dict(config_dict) mock_client = mocker.AsyncMock() - mock_client.vector_stores.retrieve.side_effect = APIConnectionError( - request=None # type: ignore - ) + mock_client.vector_stores.retrieve.side_effect = ApiException(status=None) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) @@ -1277,7 +1267,7 @@ async def test_create_file_non_size_bad_request_returns_400( mock_response = mocker.Mock() mock_response.request = mocker.Mock() mock_client.files.create.side_effect = BadRequestError( - message="Invalid file format", response=mock_response, body=None + status=400, reason="Invalid file format" ) mock_lsc = mocker.patch( "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" diff --git a/tests/unit/authorization/test_azure_token_manager.py b/tests/unit/authorization/test_azure_token_manager.py index 89565bf95..b252a6aa7 100644 --- a/tests/unit/authorization/test_azure_token_manager.py +++ b/tests/unit/authorization/test_azure_token_manager.py @@ -18,6 +18,7 @@ ) from configuration import AzureEntraIdConfiguration from constants import DEFAULT_LOGGER_NAME +from utils.types import Singleton @pytest.fixture(name="dummy_config") @@ -33,9 +34,10 @@ def dummy_config_fixture() -> AzureEntraIdConfiguration: @pytest.fixture(autouse=True) def reset_singleton() -> Generator[None, None, None]: - """Reset the singleton instance before each test.""" - AzureEntraIDManager._instances = {} # type: ignore[attr-defined] + """Reset the AzureEntraIDManager singleton before each test.""" + Singleton._instances.pop(AzureEntraIDManager, None) yield + Singleton._instances.pop(AzureEntraIDManager, None) @pytest.fixture(name="token_manager") diff --git a/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py b/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py index 962f7625a..9a059cd72 100644 --- a/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py +++ b/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py @@ -477,6 +477,12 @@ def model_fixture(self, mocker: MockerFixture) -> OgxResponsesModel: mocker.patch( "pydantic_ai_lightspeed.llamastack._model.check_allow_model_requests" ) + mocker.patch.object( + type(model), + "profile", + new_callable=mocker.PropertyMock, + return_value={}, + ) return model @pytest.mark.asyncio diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py index 73c47afb7..1d0a239da 100644 --- a/tests/unit/test_client.py +++ b/tests/unit/test_client.py @@ -9,22 +9,22 @@ import pytest from fastapi import HTTPException -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pydantic import AnyHttpUrl, SecretStr from pytest_mock import MockerFixture from authorization.azure_token_manager import AzureEntraIDManager from client import AsyncOgxClientHolder from configuration import AzureEntraIdConfiguration -from tests.unit.conftest import make_openai_model, make_openai_models_list_response from models.config import LlamaStackConfiguration +from tests.unit.conftest import make_openai_model, make_openai_models_list_response from utils.types import Singleton @pytest.fixture(autouse=True) def reset_singleton() -> None: """Reset singleton state between tests.""" - Singleton._instances = {} + Singleton._instances.clear() def test_async_client_get_client_method() -> None: @@ -57,9 +57,7 @@ async def test_get_async_llama_stack_library_client() -> None: async with client.get_client() as ls_client: assert ls_client is not None - assert not ls_client.is_closed() await ls_client.close() - assert ls_client.is_closed() @pytest.mark.asyncio @@ -112,7 +110,7 @@ async def test_get_async_llama_stack_wrong_configuration( @pytest.mark.asyncio async def test_update_azure_token_service_client() -> None: """Test update_azure_token replaces the service client with new provider headers.""" - AzureEntraIDManager._instances = {} # type: ignore[attr-defined] + Singleton._instances.pop(AzureEntraIDManager, None) manager = AzureEntraIDManager() manager.set_config( AzureEntraIdConfiguration( @@ -151,7 +149,7 @@ async def test_update_azure_token_service_client() -> None: @pytest.mark.asyncio async def test_load_service_client_defers_azure_provider_data() -> None: """Test service client load does not set Azure headers until update_azure_token.""" - AzureEntraIDManager._instances = {} # type: ignore[attr-defined] + Singleton._instances.pop(AzureEntraIDManager, None) manager = AzureEntraIDManager() manager.set_config( AzureEntraIdConfiguration( @@ -230,7 +228,6 @@ def holder_with_mock_client( @pytest.mark.asyncio async def test_model_available( self, - mocker: MockerFixture, holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns True when the model is found in the registry.""" @@ -247,7 +244,6 @@ async def test_model_available( @pytest.mark.asyncio async def test_model_not_found_service_client( self, - mocker: MockerFixture, holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns False and skips reload for non-library (service) clients.""" @@ -276,15 +272,11 @@ async def test_client_not_initialized(self) -> None: "exception_factory", [ pytest.param( - lambda m: APIConnectionError(request=m.Mock()), + lambda m: ApiException(status=None), id="connection_error", ), pytest.param( - lambda m: APIStatusError( - message="Internal error", - response=m.Mock(status_code=500, headers={}), - body=None, - ), + lambda m: ApiException(status=500, reason="Internal error"), id="api_status_error", ), ], diff --git a/tests/unit/utils/agents/test_query.py b/tests/unit/utils/agents/test_query.py index 97aa40649..d266f20d0 100644 --- a/tests/unit/utils/agents/test_query.py +++ b/tests/unit/utils/agents/test_query.py @@ -6,7 +6,7 @@ import pytest from fastapi import HTTPException -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from pydantic_ai.messages import ( FinishReason, ImageUrl, @@ -479,9 +479,7 @@ async def test_agent_connection_error_raises_http_exception( ) -> None: """Test Llama Stack connection errors are mapped to HTTPException.""" mock_agent = mocker.AsyncMock() - mock_agent.run = mocker.AsyncMock( - side_effect=APIConnectionError(request=mocker.Mock()) - ) + mock_agent.run = mocker.AsyncMock(side_effect=ApiException(status=None)) mocker.patch( "utils.agents.query.build_agent", return_value=mock_agent, @@ -507,11 +505,7 @@ async def test_api_status_error_raises_http_exception( """Test API status errors from the agent run are mapped to HTTPException.""" mock_agent = mocker.AsyncMock() mock_agent.run = mocker.AsyncMock( - side_effect=APIStatusError( - message="quota exceeded", - response=mocker.Mock(), - body=None, - ) + side_effect=ApiException(status=500, reason="quota exceeded") ) mocker.patch( "utils.agents.query.build_agent", diff --git a/tests/unit/utils/agents/test_streaming.py b/tests/unit/utils/agents/test_streaming.py index f452f8106..c1ed200a0 100644 --- a/tests/unit/utils/agents/test_streaming.py +++ b/tests/unit/utils/agents/test_streaming.py @@ -10,7 +10,7 @@ import pytest from fastapi import HTTPException -from ogx_client import APIStatusError +from ogx_client import ApiException from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError from pydantic_ai.messages import ( @@ -735,11 +735,7 @@ async def inner() -> AsyncIterator[str]: TokenStreamPayload.create(chunk_id=0, token="partial"), MEDIA_TYPE_JSON, ) - raise APIStatusError( - message="quota exceeded", - response=mocker.Mock(), - body=None, - ) + raise ApiException(status=500, reason="quota exceeded") mock_error = mocker.Mock() mock_error.status_code = 429 diff --git a/tests/unit/utils/test_builtin_tools.py b/tests/unit/utils/test_builtin_tools.py index c7ba9949c..a0cc794a8 100644 --- a/tests/unit/utils/test_builtin_tools.py +++ b/tests/unit/utils/test_builtin_tools.py @@ -4,7 +4,7 @@ import pytest from fastapi import HTTPException -from ogx_client import APIConnectionError +from ogx_client import ApiException from ogx_client.models.provider_info import ProviderInfo from pytest_mock import MockerFixture @@ -99,7 +99,7 @@ async def test_get_file_search_tools_raises_503_on_provider_connection_error( """Raise HTTP 503 when Llama Stack is unreachable during provider discovery.""" client = mocker.AsyncMock() client.providers.list = mocker.AsyncMock( - side_effect=APIConnectionError(message="down", request=mocker.Mock()) + side_effect=ApiException(status=None, reason="down") ) with pytest.raises(HTTPException) as exc_info: diff --git a/tests/unit/utils/test_conversations.py b/tests/unit/utils/test_conversations.py index f6d24dbbf..5f0f6653f 100644 --- a/tests/unit/utils/test_conversations.py +++ b/tests/unit/utils/test_conversations.py @@ -1,12 +1,14 @@ """Unit tests for conversation utility functions.""" +# pylint: disable=too-many-lines + from datetime import UTC, datetime from typing import Any import pytest from fastapi import HTTPException from ogx_api import OpenAIResponseMessage -from ogx_client import APIConnectionError, APIStatusError +from ogx_client import ApiException from ogx_client.models.add_items_request import AddItemsRequest from ogx_client.models.open_ai_response_input_function_tool_call_output import ( OpenAIResponseInputFunctionToolCallOutput as FunctionCallOutput, @@ -340,7 +342,7 @@ def test_mcp_approval_response_without_reason(self, mocker: MockerFixture) -> No assert tool_result is not None assert tool_result.content == "{}" - def test_function_call_output(self, mocker: MockerFixture) -> None: + def test_function_call_output(self) -> None: """Test parsing a function_call_output item.""" item = FunctionCallOutput.from_dict( { @@ -361,7 +363,7 @@ def test_function_call_output(self, mocker: MockerFixture) -> None: assert tool_result.type == "function_call_output" assert tool_result.round == 1 - def test_function_call_output_without_status(self, mocker: MockerFixture) -> None: + def test_function_call_output_without_status(self) -> None: """Test parsing a function_call_output item without status.""" item = FunctionCallOutput.from_dict( { @@ -376,9 +378,7 @@ def test_function_call_output_without_status(self, mocker: MockerFixture) -> Non assert tool_result is not None assert tool_result.status == "success" # Defaults to "success" - def test_function_call_output_with_structured_content( - self, mocker: MockerFixture - ) -> None: + def test_function_call_output_with_structured_content(self) -> None: """Test parsing function_call_output with mixed content parts.""" item = FunctionCallOutput.from_dict( { @@ -988,9 +988,7 @@ async def test_returns_all_items_across_pages(self, mocker: MockerFixture) -> No second_page.data = [item_2, item_3] second_page.has_more = False - mock_client.items.list = mocker.AsyncMock( - side_effect=[first_page, second_page] - ) + mock_client.items.list = mocker.AsyncMock(side_effect=[first_page, second_page]) result = await get_all_conversation_items(mock_client, "conv_abc") @@ -1012,12 +1010,10 @@ async def test_handles_empty_data(self, mocker: MockerFixture) -> None: @pytest.mark.asyncio async def test_handles_connection_error(self, mocker: MockerFixture) -> None: - """Test that APIConnectionError is converted to HTTPException 503.""" + """Test that ApiException is converted to HTTPException 503.""" mock_client = mocker.Mock() mock_client.items.list = mocker.AsyncMock( - side_effect=APIConnectionError( - message="connection refused", request=mocker.Mock() - ) + side_effect=ApiException(status=None, reason="connection refused") ) with pytest.raises(HTTPException) as exc_info: @@ -1028,14 +1024,10 @@ async def test_handles_connection_error(self, mocker: MockerFixture) -> None: @pytest.mark.asyncio async def test_handles_api_status_error(self, mocker: MockerFixture) -> None: - """Test that APIStatusError is converted to HTTPException 500.""" + """Test that ApiException is converted to HTTPException 500.""" mock_client = mocker.Mock() mock_client.items.list = mocker.AsyncMock( - side_effect=APIStatusError( - message="internal error", - response=mocker.Mock(request=None), - body=None, - ) + side_effect=ApiException(status=500, reason="internal error") ) with pytest.raises(HTTPException) as exc_info: diff --git a/tests/unit/utils/test_llama_stack_version.py b/tests/unit/utils/test_llama_stack_version.py index 30864622c..32bc7d4e5 100644 --- a/tests/unit/utils/test_llama_stack_version.py +++ b/tests/unit/utils/test_llama_stack_version.py @@ -3,7 +3,7 @@ from typing import Any import pytest -from ogx_client import APIConnectionError +from ogx_client import ApiException from ogx_client.models.version_info import VersionInfo from pytest_mock import MockerFixture from pytest_subtests import SubTests @@ -122,14 +122,14 @@ async def test_check_llama_stack_version_too_big_version( async def test_check_llama_stack_version_retries_on_connection_error( mocker: MockerFixture, ) -> None: - """Test that check_llama_stack_version retries on APIConnectionError.""" + """Test that check_llama_stack_version retries on ApiException.""" mock_client = mocker.AsyncMock() mock_sleep = mocker.patch("utils.llama_stack_version.asyncio.sleep") # Fail twice with connection error, then succeed mock_client.inspect.version.side_effect = [ - APIConnectionError(request=mocker.MagicMock()), - APIConnectionError(request=mocker.MagicMock()), + ApiException(status=None), + ApiException(status=None), VersionInfo(version=MINIMAL_SUPPORTED_LLAMA_STACK_VERSION), ] @@ -147,11 +147,9 @@ async def test_check_llama_stack_version_raises_after_max_retries( mock_client = mocker.AsyncMock() mock_sleep = mocker.patch("utils.llama_stack_version.asyncio.sleep") - mock_client.inspect.version.side_effect = APIConnectionError( - request=mocker.MagicMock() - ) + mock_client.inspect.version.side_effect = ApiException(status=None) - with pytest.raises(APIConnectionError): + with pytest.raises(ApiException): await check_llama_stack_version(mock_client, max_retries=3, retry_delay=1) assert mock_client.inspect.version.call_count == 3 diff --git a/tests/unit/utils/test_pydantic_ai.py b/tests/unit/utils/test_pydantic_ai.py index 82068743a..862544d14 100644 --- a/tests/unit/utils/test_pydantic_ai.py +++ b/tests/unit/utils/test_pydantic_ai.py @@ -22,6 +22,7 @@ ) from pydantic_ai_lightspeed.capabilities import QuestionValidity from pydantic_ai_lightspeed.capabilities.redaction import PiiRedactionCapability +from tests.unit.conftest import attach_mock_api_client from utils.pydantic_ai_helpers import ( _agent_capabilities, _shield_capability, @@ -29,7 +30,6 @@ build_agent, get_agent_capability_tools, ) -from tests.unit.conftest import attach_mock_api_client _QUESTION_VALIDITY_MODULE = ( "pydantic_ai_lightspeed.capabilities.question_validity._capability" diff --git a/tests/unit/utils/test_query.py b/tests/unit/utils/test_query.py index ae5984a23..c9031884b 100644 --- a/tests/unit/utils/test_query.py +++ b/tests/unit/utils/test_query.py @@ -9,6 +9,7 @@ import psycopg2 import pytest from fastapi import HTTPException +from ogx_client import ApiException from pydantic_ai.messages import ImageUrl from pytest_mock import MockerFixture from sqlalchemy.exc import SQLAlchemyError @@ -311,11 +312,10 @@ class TestHandleKnownApistatusErrors: def test_context_length_exceeded(self) -> None: """Test handling context length exceeded error.""" - error = type( - "APIStatusError", - (), - {"status_code": 400, "message": "context_length_exceeded: prompt too long"}, - )() + error = ApiException( + status=400, + reason="context_length_exceeded: prompt too long", + ) result = handle_known_apistatus_errors(error, "model1") assert isinstance(result, PromptTooLongResponse) detail = result.model_dump()["detail"] @@ -325,9 +325,7 @@ def test_context_length_exceeded(self) -> None: def test_quota_exceeded(self) -> None: """Test handling quota exceeded error.""" - error = type( - "APIStatusError", (), {"status_code": 429, "message": "Rate limit exceeded"} - )() + error = ApiException(status=429, reason="Rate limit exceeded") result = handle_known_apistatus_errors(error, "model1") assert isinstance(result, QuotaExceededResponse) detail = result.model_dump()["detail"] @@ -335,11 +333,7 @@ def test_quota_exceeded(self) -> None: def test_generic_error(self) -> None: """Test handling generic error.""" - error = type( - "APIStatusError", - (), - {"status_code": 500, "message": "Internal server error"}, - )() + error = ApiException(status=500, reason="Internal server error") result = handle_known_apistatus_errors(error, "model1") assert isinstance(result, InternalServerErrorResponse) detail = result.model_dump()["detail"] diff --git a/tests/unit/utils/test_responses.py b/tests/unit/utils/test_responses.py index 4879b60ff..46a2ecca5 100644 --- a/tests/unit/utils/test_responses.py +++ b/tests/unit/utils/test_responses.py @@ -50,12 +50,10 @@ from ogx_api.openai_responses import ( OpenAIResponseOutputMessageWebSearchToolCall as WebSearchCall, ) -from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client import ApiException, AsyncOgxClient from pydantic import AnyUrl, BaseModel from pytest_mock import MockerFixture -from tests.unit.conftest import make_openai_model, make_openai_models_list_response - import constants from models.api.requests import QueryRequest from models.common.query import Attachment @@ -66,6 +64,7 @@ InferenceConfiguration, ModelContextProtocolServer, ) +from tests.unit.conftest import make_openai_model, make_openai_models_list_response from utils.query import normalize_vertex_ai_model_id from utils.responses import ( _build_chunk_attributes, @@ -930,6 +929,10 @@ class TestGetTopicSummary: async def test_get_topic_summary_success(self, mocker: MockerFixture) -> None: """Test successful topic summary generation.""" mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_output_item = make_output_item( item_type="message", role="assistant", content="Topic Summary" ) @@ -952,6 +955,10 @@ async def test_get_topic_summary_empty_response( ) -> None: """Test topic summary with empty response.""" mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_response = mocker.Mock() mock_response.output = [] mock_client.responses.create = mocker.AsyncMock(return_value=mock_response) @@ -970,10 +977,12 @@ async def test_get_topic_summary_connection_error( ) -> None: """Test topic summary raises HTTPException on connection error.""" mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.responses.create = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", request=mocker.Mock() - ) + side_effect=ApiException(status=None, reason="Connection failed") ) mocker.patch( @@ -989,10 +998,12 @@ async def test_get_topic_summary_connection_error( async def test_get_topic_summary_api_error(self, mocker: MockerFixture) -> None: """Test topic summary raises HTTPException on API error.""" mock_client = mocker.AsyncMock(spec=AsyncOgxClient) - # Create a mock exception that will be caught by except APIStatusError - mock_error = APIStatusError( - message="API error", response=mocker.Mock(request=None), body=None - ) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") + # Create a mock exception that will be caught by except ApiException + mock_error = ApiException(status=500, reason="API error") mock_client.responses.create = mocker.AsyncMock(side_effect=mock_error) mocker.patch( @@ -1852,7 +1863,11 @@ async def test_prepare_responses_params_with_conversation_id( self, mocker: MockerFixture ) -> None: """Test prepare_responses_params with existing conversation ID.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( return_value=make_openai_models_list_response( make_openai_model( @@ -1889,7 +1904,11 @@ async def test_prepare_responses_params_create_conversation( self, mocker: MockerFixture ) -> None: """Test prepare_responses_params creates new conversation when ID not provided.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( return_value=make_openai_models_list_response( make_openai_model( @@ -1926,11 +1945,13 @@ async def test_prepare_responses_params_connection_error_on_models( self, mocker: MockerFixture ) -> None: """Test prepare_responses_params raises HTTPException on connection error when fetching models.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", request=mocker.Mock() - ) + side_effect=ApiException(status=None, reason="Connection failed") ) query_request = QueryRequest(query="test") # pyright: ignore[reportCallIssue] @@ -1947,7 +1968,11 @@ async def test_prepare_responses_params_connection_error_on_conversation( self, mocker: MockerFixture ) -> None: """Test prepare_responses_params raises HTTPException on connection error when creating conversation.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( return_value=make_openai_models_list_response( make_openai_model( @@ -1958,9 +1983,7 @@ async def test_prepare_responses_params_connection_error_on_conversation( ) ) mock_client.conversations.create = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", request=mocker.Mock() - ) + side_effect=ApiException(status=None, reason="Connection failed") ) query_request = QueryRequest(query="test") # pyright: ignore[reportCallIssue] @@ -1981,11 +2004,13 @@ async def test_prepare_responses_params_api_status_error_on_models( self, mocker: MockerFixture ) -> None: """Test prepare_responses_params raises HTTPException on API status error when fetching models.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( - side_effect=APIStatusError( - message="API error", response=mocker.Mock(request=None), body=None - ) + side_effect=ApiException(status=500, reason="API error") ) query_request = QueryRequest(query="test") # pyright: ignore[reportCallIssue] @@ -2002,7 +2027,11 @@ async def test_prepare_responses_params_includes_mcp_provider_data_headers( self, mocker: MockerFixture ) -> None: """Test that extra_headers with x-llamastack-provider-data is set when MCP tools have headers.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( return_value=make_openai_models_list_response( make_openai_model( @@ -2071,7 +2100,11 @@ async def test_prepare_responses_params_no_extra_headers_without_mcp_tools( self, mocker: MockerFixture ) -> None: """Test that extra_headers is None when no MCP tools have headers.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( return_value=make_openai_models_list_response( make_openai_model( @@ -2109,7 +2142,11 @@ async def test_prepare_responses_params_api_status_error_on_conversation( self, mocker: MockerFixture ) -> None: """Test prepare_responses_params raises HTTPException on API status error when creating conversation.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_client.openai.list = mocker.AsyncMock( return_value=make_openai_models_list_response( make_openai_model( @@ -2120,9 +2157,7 @@ async def test_prepare_responses_params_api_status_error_on_conversation( ) ) mock_client.conversations.create = mocker.AsyncMock( - side_effect=APIStatusError( - message="API error", response=mocker.Mock(request=None), body=None - ) + side_effect=ApiException(status=500, reason="API error") ) query_request = QueryRequest(query="test") # pyright: ignore[reportCallIssue] @@ -2143,7 +2178,11 @@ async def test_image_attachments_excluded_from_input( self, mocker: MockerFixture ) -> None: """Test that image attachments are excluded from text input.""" - mock_client = mocker.AsyncMock() + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_client.attach_mock(mocker.AsyncMock(), "responses") + mock_client.attach_mock(mocker.AsyncMock(), "items") + mock_client.attach_mock(mocker.AsyncMock(), "openai") + mock_client.attach_mock(mocker.AsyncMock(), "conversations") mock_conversation = mocker.Mock() mock_conversation.id = "new_conv_id" mock_client.conversations.create = mocker.AsyncMock( diff --git a/uv.lock b/uv.lock index 3a5596da6..00d9863fe 100644 --- a/uv.lock +++ b/uv.lock @@ -1846,7 +1846,7 @@ requires-dist = [ { name = "prometheus-client", specifier = ">=0.22.1" }, { name = "psycopg2-binary", specifier = ">=2.9.10" }, { name = "pyasn1", specifier = ">=0.6.3" }, - { name = "pydantic-ai", specifier = "==2.16.0" }, + { name = "pydantic-ai", specifier = ">=2.23.0" }, { name = "pydantic-ai-skills", specifier = ">=0.11.0" }, { name = "python-dotenv", specifier = ">=1.2.2" }, { name = "pyyaml", specifier = ">=6.0.0" }, @@ -3136,14 +3136,14 @@ email = [ [[package]] name = "pydantic-ai" -version = "2.16.0" +version = "2.23.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "pydantic-ai-slim", extra = ["anthropic", "cli", "evals", "google", "logfire", "mcp", "openai", "retries", "web"] }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ac/22/d3069491e73d2dda637afb87e3ad1957fe2924e9954d3cc0b9f1d312fd5e/pydantic_ai-2.16.0.tar.gz", hash = "sha256:89aed181b2f317d0bca958ba1885b8480559583321ee9e6bfd85ad45df0fe905", size = 18761, upload-time = "2026-07-23T02:47:18.419Z" } +sdist = { url = "https://files.pythonhosted.org/packages/1c/73/8dbd43b74f31c187a57fc2b7ae35d2596893b8f386abebadc2d136e62e7e/pydantic_ai-2.23.0.tar.gz", hash = "sha256:3da15a28e171cbb4548f3fffbd098dd9df44888c800dd7e48633795e18525a07", size = 19369, upload-time = "2026-08-04T01:58:18.18Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d7/b5/3412d12dbb7b88ee4c7905e6d8046150fe855aa5ff51924e0dc3c5259626/pydantic_ai-2.16.0-py3-none-any.whl", hash = "sha256:7f206e0860f547fd9a72cd95e4a6d10bdaf91e68818ac04d1ee044dd2c261946", size = 7729, upload-time = "2026-07-23T02:47:09.537Z" }, + { url = "https://files.pythonhosted.org/packages/3d/af/965fb83595ab34f5c0b6363323ac15c9aa478fc5de7cfc3360386eed7bd0/pydantic_ai-2.23.0-py3-none-any.whl", hash = "sha256:a9042f5880522565c36e716a983c196d57cc9e2c40e8fd1188ee40802fc8d104", size = 7740, upload-time = "2026-08-04T01:58:08.868Z" }, ] [[package]] @@ -3162,9 +3162,10 @@ wheels = [ [[package]] name = "pydantic-ai-slim" -version = "2.16.0" +version = "2.23.0" source = { registry = "https://pypi.org/simple" } dependencies = [ + { name = "anyio" }, { name = "genai-prices" }, { name = "griffelib" }, { name = "httpx" }, @@ -3173,9 +3174,9 @@ dependencies = [ { name = "pydantic-graph" }, { name = "typing-inspection" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/76/6a/3048579c646f4cea7966889009a1d7eef1365f9126393945d39cefad9e83/pydantic_ai_slim-2.16.0.tar.gz", hash = "sha256:36d17cb12edd72ffc62f9e06cc49ac5f23cb77cf6b665c4b1fd11c00dbe6a852", size = 889343, upload-time = "2026-07-23T02:47:20.696Z" } +sdist = { url = "https://files.pythonhosted.org/packages/0d/9f/53b19efefa041c1080f7c4ad41679a9293cce64f1265168a98cbe06a0ab7/pydantic_ai_slim-2.23.0.tar.gz", hash = "sha256:d16dcbfb2bfea0ee162bf0f499442fab5a4d69b41e4c54f3b694c2e90b983768", size = 965485, upload-time = "2026-08-04T01:58:20.668Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/2e/95/2c79b9f8e875562bae8141078af27428120e8415d493bb07ea36ca5905f9/pydantic_ai_slim-2.16.0-py3-none-any.whl", hash = "sha256:7cad27fb8f45ce4af4e8da83d7f206a3e1038dbc7535625c0ce5518fe378a55d", size = 1077057, upload-time = "2026-07-23T02:47:12.633Z" }, + { url = "https://files.pythonhosted.org/packages/9e/6f/539a255524178a8421d582271a8d7f8667b036f02b4ddc4f20abcc63888b/pydantic_ai_slim-2.23.0-py3-none-any.whl", hash = "sha256:a2fa3e56408bbf1b83900e3dd4ad9b137297742f450863c2f0f9a03a547d0e33", size = 1157486, upload-time = "2026-08-04T01:58:12.356Z" }, ] [package.optional-dependencies] @@ -3261,7 +3262,7 @@ wheels = [ [[package]] name = "pydantic-evals" -version = "2.16.0" +version = "2.23.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio" }, @@ -3271,24 +3272,25 @@ dependencies = [ { name = "pyyaml" }, { name = "rich" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/52/f5/7c2bf8ce45da52ed70c030e91ce43747317f61bd0615a40e91f1ec515c70/pydantic_evals-2.16.0.tar.gz", hash = "sha256:717e9615c7688650cdc716046f1d8edfbe0bac1748286c0a2744e3689fb518c7", size = 85147, upload-time = "2026-07-23T02:47:22.089Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ae/4e/ac3bcbbefe683991e8cbd3f69c624c2d550002e9f33fe03ba9e69309ef94/pydantic_evals-2.23.0.tar.gz", hash = "sha256:3f5e16708976c165ae23109f55143fa3d68a3f569b35bc70ca7c54cf737df63e", size = 85391, upload-time = "2026-08-04T01:58:21.933Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f0/f6/d7f20f976f81d073856138bd046bf1ce5dc96561391526753c1c9076d33c/pydantic_evals-2.16.0-py3-none-any.whl", hash = "sha256:6705b427ea7c77d7f6b7d152ced6f86728931eceb9c57a8aa07ab8225fb1a934", size = 100439, upload-time = "2026-07-23T02:47:14.523Z" }, + { url = "https://files.pythonhosted.org/packages/45/72/569511b3de588615a9151b727dedc181d2d54d5434d55b768dc26fa20ab0/pydantic_evals-2.23.0-py3-none-any.whl", hash = "sha256:8cde69fc2e126b20372488187b016f329fab710bd325d10dc4083a9e19ff04b2", size = 100540, upload-time = "2026-08-04T01:58:14.431Z" }, ] [[package]] name = "pydantic-graph" -version = "2.16.0" +version = "2.23.0" source = { registry = "https://pypi.org/simple" } dependencies = [ + { name = "anyio" }, { name = "httpx" }, { name = "logfire-api" }, { name = "pydantic" }, { name = "typing-inspection" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c3/22/6b6426e14607275f6b15f2a0c4530514c48ff01f86c53bbaac4041d09df9/pydantic_graph-2.16.0.tar.gz", hash = "sha256:f71e5c8e78a4ce56bc044861178e506ebd99881717ae0993924c457f5de6230e", size = 43979, upload-time = "2026-07-23T02:47:23.328Z" } +sdist = { url = "https://files.pythonhosted.org/packages/09/fc/273bac7d14fb62c060e0c20c51a9cc60e9e90a96d992fd20e77abf4b6ac1/pydantic_graph-2.23.0.tar.gz", hash = "sha256:54c9939f47fd8a268c96320d7d90e7cef037cbfd2625a675dc4028c1377f70ab", size = 45179, upload-time = "2026-08-04T01:58:23.085Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/80/c0/362a7fb50562d7b51e3d02d454826e0a7b5652f27e5be8e835ddfaf84da6/pydantic_graph-2.16.0-py3-none-any.whl", hash = "sha256:99c25852c436d4d510d1ecdcb53ce4a0b11aa4d1d9d81472727587f2b5a3bc13", size = 51661, upload-time = "2026-07-23T02:47:15.921Z" }, + { url = "https://files.pythonhosted.org/packages/0d/c4/875cf853d205dc55422bd44ff0fbfac82e6e34ff7693df016fc3a4088d32/pydantic_graph-2.23.0-py3-none-any.whl", hash = "sha256:b0f12b4f72adb2a5522b5962c95e1a7b140cb3f631a628036f935e291a9e50ba", size = 52662, upload-time = "2026-08-04T01:58:15.858Z" }, ] [[package]]