diff --git a/.konflux/requirements.hashes.source.txt b/.konflux/requirements.hashes.source.txt index 731aea0a5..29016231d 100644 --- a/.konflux/requirements.hashes.source.txt +++ b/.konflux/requirements.hashes.source.txt @@ -99,15 +99,6 @@ google-cloud-bigquery==3.42.2 \ google-cloud-resource-manager==1.18.0 \ --hash=sha256:69bab144acd75878ebe44b720903dc3d140cb7d3be3261eaa8bc81e48afaff33 \ --hash=sha256:db689f800a14c66d041196a7fbb8bb8aae8dc87f28c2929e101a5ec766b15512 -llama-stack==0.6.0 \ - --hash=sha256:b804830664dc91e54c7225a7a081cb1874c48fc18573569c19fac4a9397e8076 \ - --hash=sha256:d92711791633f5505a4473ffba3f3e26acb700716fddab5aec419d99e614c802 -llama-stack-api==0.6.0 \ - --hash=sha256:b99a03aba3659736b6b540c9e5e674b1daac2bf5eeb2a68795113d62b8250672 \ - --hash=sha256:f0f3a1a6239a5d3b8c7ef02cefdf817c96c6461dcd8a82c1689ac67ec3107270 -llama-stack-client==0.6.0 \ - --hash=sha256:3290aac36dcafbd1bc0baaf995522e2037f57056672b5a1516af112a4210f3ea \ - --hash=sha256:7e514a6ffd92f237aceb062dadc4db44e24a3cd9c4ea35e25173d1e0739beb8e oci==2.182.1 \ --hash=sha256:0c616a6bc3bc458464bc3456469d8da63a1a2d2277e9314b41a1c4e76d5df523 \ --hash=sha256:9862de221f2abe9cf8319393eec58ea59c014fd9b61afaf0a3cca163e2a508b0 diff --git a/.konflux/requirements.hashes.wheel.txt b/.konflux/requirements.hashes.wheel.txt index d23afc79a..dac4d632a 100644 --- a/.konflux/requirements.hashes.wheel.txt +++ b/.konflux/requirements.hashes.wheel.txt @@ -222,6 +222,12 @@ oauthlib==3.3.1 \ --hash=sha256:c6fbab4f1a77a539f01175e2ea74b9552806bd0a849c70e744fda4ab801031c0 openai==2.44.0 \ --hash=sha256:87429e9a4d15b2918a03b040b6330fa03fc175bbf0bc6eaff7ae61c93cd42c53 +ogx==1.0.2+rhaiv.0 \ + --hash=sha256:52c891af9dfb22f884d2dae150e5e94d04e492f452ec811bbc61ba84257de000 +ogx-api==1.0.2+rhaiv.0 \ + --hash=sha256:265dff54f2367f4f366952e7a17abe06ceeef777a809259a5fba4715f4df4e90 +ogx-client==1.0.2 \ + --hash=sha256:3ac249fb8365a4cb801815c95113fc6954a16193db5247d9b9f053398a80add9 opentelemetry-api==1.42.1 \ --hash=sha256:078a23234520ddf8654a48045b8f582c4a73c9bd5da2f0d23e05aba6e40a7c91 opentelemetry-distro==0.63b1 \ @@ -374,6 +380,8 @@ sse-starlette==3.4.5 \ --hash=sha256:c1f701f19f43c181be2d62c6164e90b7e8e4d5d67956fface4c794ac7a784962 starlette==1.3.1 \ --hash=sha256:19acd5bfb037734bb027cb753f36379bca5f72da9b79fa7fb3ef4639b1999f6b +structlog==26.1.0 \ + --hash=sha256:0054592356113edd68e13e9be1be5062566ddbf15a24bb3f400ab764a26b0b36 sympy==1.14.0 \ --hash=sha256:6cf6a1e379dbba75ca560e2cbd4394144f9e7e7335ee1be872c224f004eeb172 tenacity==9.1.4 \ @@ -439,3 +447,6 @@ yarl==1.24.2 \ --hash=sha256:c72aabbf544bceb3fa895376a188a90b7152efd47045b589fd358505ad8a85d4 zipp==4.1.0 \ --hash=sha256:f04782919048b93bac92fa1d20b2dd8a062b30ebee718e834cf663874884c645 +zstandard==0.25.0 \ + --hash=sha256:3f12f3143d2dba84081a79932a30e8179df5c655134d1ae613cc760aab9e4ff2 \ + --hash=sha256:c734ffe543d7aaf7ff9848fb06fcc4c251e2c1670eab845b8983b6ff24f993fd diff --git a/.tekton/lightspeed-stack-0-7-pull-request.yaml b/.tekton/lightspeed-stack-0-7-pull-request.yaml index 57b0f79a5..59a73d453 100644 --- a/.tekton/lightspeed-stack-0-7-pull-request.yaml +++ b/.tekton/lightspeed-stack-0-7-pull-request.yaml @@ -53,7 +53,7 @@ spec: ], "requirements_build_files": ["requirements-build.txt"], "binary": { - "packages": "a2a-sdk,accelerate,aiofile,aiohappyeyeballs,aiohttp,aiosignal,aiosqlite,annotated-doc,annotated-types,anthropic,anyio,argcomplete,asyncpg,attrs,authlib,autoevals,azure-core,azure-identity,beartype,cachetools,caio,certifi,cffi,chardet,charset-normalizer,chevron,click,cryptography,datasets,dill,distro,dnspython,docstring-parser,durationpy,einops,email-validator,emoji,exceptiongroup,executing,faiss-cpu,fastapi,fastmcp-slim,fastuuid,filelock,fire,frozenlist,fsspec,genai-prices,google-api-core,google-auth,google-cloud-core,google-cloud-storage,google-crc32c,google-genai,google-resumable-media,googleapis-common-protos,greenlet,griffelib,grpc-google-iam-v1,grpcio,grpcio-status,h11,hf-xet,httpcore,httpcore2,httpx,httpx-sse,httpx2,huggingface-hub,idna,importlib-metadata,jaraco-classes,jaraco-context,jaraco-functools,jeepney,jinja2,jiter,joblib,joserfc,jsonpath-ng,jsonschema,jsonschema-specifications,keyring,kubernetes,langdetect,litellm,logfire,logfire-api,markdown-it-py,markupsafe,maturin,mcp,mdurl,more-itertools,mpmath,msal,msal-extensions,multidict,multiprocess,narwhals,networkx,nltk,numpy,oauthlib,openai,opentelemetry-api,opentelemetry-distro,opentelemetry-exporter-otlp,opentelemetry-exporter-otlp-proto-common,opentelemetry-exporter-otlp-proto-grpc,opentelemetry-exporter-otlp-proto-http,opentelemetry-instrumentation,opentelemetry-instrumentation-httpx,opentelemetry-proto,opentelemetry-sdk,opentelemetry-semantic-conventions,opentelemetry-util-http,oracledb,packaging,pandas,peft,pip,platformdirs,polyleven,prometheus-client,prompt-toolkit,propcache,proto-plus,protobuf,psutil,psycopg2-binary,py-key-value-aio,pyaml,pyarrow,pyasn1,pyasn1-modules,pycparser,pydantic,pydantic-ai,pydantic-ai-slim,pydantic-core,pydantic-evals,pydantic-graph,pydantic-settings,pygments,pyjwt,pyopenssl,pypdf,pyperclip,python-dateutil,python-dotenv,python-multipart,pytz,pyyaml,referencing,regex,requests,requests-oauthlib,rich,rpds-py,safetensors,scikit-learn,scipy,secretstorage,semver,sentence-transformers,sentry-sdk,setuptools,shellingham,six,sniffio,sqlalchemy,sse-starlette,starlette,sympy,tenacity,termcolor,threadpoolctl,tiktoken,tokenizers,torch,tornado,tqdm,transformers,tree-sitter,triton,trl,truststore,typer,typing-extensions,typing-inspection,urllib3,uv,uv-build,uvicorn,wcwidth,websocket-client,websockets,wrapt,xxhash,yarl,zipp", + "packages": "a2a-sdk,accelerate,aiofile,aiohappyeyeballs,aiohttp,aiosignal,aiosqlite,annotated-doc,annotated-types,anthropic,anyio,argcomplete,asyncpg,attrs,authlib,autoevals,azure-core,azure-identity,beartype,cachetools,caio,certifi,cffi,chardet,charset-normalizer,chevron,click,cryptography,datasets,dill,distro,dnspython,docstring-parser,durationpy,einops,email-validator,emoji,exceptiongroup,executing,faiss-cpu,fastapi,fastmcp-slim,fastuuid,filelock,fire,frozenlist,fsspec,genai-prices,google-api-core,google-auth,google-cloud-core,google-cloud-storage,google-crc32c,google-genai,google-resumable-media,googleapis-common-protos,greenlet,griffelib,grpc-google-iam-v1,grpcio,grpcio-status,h11,hf-xet,httpcore,httpcore2,httpx,httpx-sse,httpx2,huggingface-hub,idna,importlib-metadata,jaraco-classes,jaraco-context,jaraco-functools,jeepney,jinja2,jiter,joblib,joserfc,jsonpath-ng,jsonschema,jsonschema-specifications,keyring,kubernetes,langdetect,litellm,logfire,logfire-api,markdown-it-py,markupsafe,maturin,mcp,mdurl,more-itertools,mpmath,msal,msal-extensions,multidict,multiprocess,narwhals,networkx,nltk,numpy,oauthlib,ogx,ogx-api,ogx-client,openai,opentelemetry-api,opentelemetry-distro,opentelemetry-exporter-otlp,opentelemetry-exporter-otlp-proto-common,opentelemetry-exporter-otlp-proto-grpc,opentelemetry-exporter-otlp-proto-http,opentelemetry-instrumentation,opentelemetry-instrumentation-httpx,opentelemetry-proto,opentelemetry-sdk,opentelemetry-semantic-conventions,opentelemetry-util-http,oracledb,packaging,pandas,peft,pip,platformdirs,polyleven,prometheus-client,prompt-toolkit,propcache,proto-plus,protobuf,psutil,psycopg2-binary,py-key-value-aio,pyaml,pyarrow,pyasn1,pyasn1-modules,pycparser,pydantic,pydantic-ai,pydantic-ai-slim,pydantic-core,pydantic-evals,pydantic-graph,pydantic-settings,pygments,pyjwt,pyopenssl,pypdf,pyperclip,python-dateutil,python-dotenv,python-multipart,pytz,pyyaml,referencing,regex,requests,requests-oauthlib,rich,rpds-py,safetensors,scikit-learn,scipy,secretstorage,semver,sentence-transformers,sentry-sdk,setuptools,shellingham,six,sniffio,sqlalchemy,sse-starlette,starlette,structlog,sympy,tenacity,termcolor,threadpoolctl,tiktoken,tokenizers,torch,tornado,tqdm,transformers,tree-sitter,triton,trl,truststore,typer,typing-extensions,typing-inspection,urllib3,uv,uv-build,uvicorn,wcwidth,websocket-client,websockets,wrapt,xxhash,yarl,zipp,zstandard", "os": "linux", "arch": "x86_64,aarch64", "py_version": 312 diff --git a/.tekton/lightspeed-stack-0-7-push.yaml b/.tekton/lightspeed-stack-0-7-push.yaml index ccfc1aa66..8cfde4276 100644 --- a/.tekton/lightspeed-stack-0-7-push.yaml +++ b/.tekton/lightspeed-stack-0-7-push.yaml @@ -54,7 +54,7 @@ spec: ], "requirements_build_files": ["requirements-build.txt"], "binary": { - "packages": "a2a-sdk,accelerate,aiofile,aiohappyeyeballs,aiohttp,aiosignal,aiosqlite,annotated-doc,annotated-types,anthropic,anyio,argcomplete,asyncpg,attrs,authlib,autoevals,azure-core,azure-identity,beartype,cachetools,caio,certifi,cffi,chardet,charset-normalizer,chevron,click,cryptography,datasets,dill,distro,dnspython,docstring-parser,durationpy,einops,email-validator,emoji,exceptiongroup,executing,faiss-cpu,fastapi,fastmcp-slim,fastuuid,filelock,fire,frozenlist,fsspec,genai-prices,google-api-core,google-auth,google-cloud-core,google-cloud-storage,google-crc32c,google-genai,google-resumable-media,googleapis-common-protos,greenlet,griffelib,grpc-google-iam-v1,grpcio,grpcio-status,h11,hf-xet,httpcore,httpcore2,httpx,httpx-sse,httpx2,huggingface-hub,idna,importlib-metadata,jaraco-classes,jaraco-context,jaraco-functools,jeepney,jinja2,jiter,joblib,joserfc,jsonpath-ng,jsonschema,jsonschema-specifications,keyring,kubernetes,langdetect,litellm,logfire,logfire-api,markdown-it-py,markupsafe,maturin,mcp,mdurl,more-itertools,mpmath,msal,msal-extensions,multidict,multiprocess,narwhals,networkx,nltk,numpy,oauthlib,openai,opentelemetry-api,opentelemetry-distro,opentelemetry-exporter-otlp,opentelemetry-exporter-otlp-proto-common,opentelemetry-exporter-otlp-proto-grpc,opentelemetry-exporter-otlp-proto-http,opentelemetry-instrumentation,opentelemetry-instrumentation-httpx,opentelemetry-proto,opentelemetry-sdk,opentelemetry-semantic-conventions,opentelemetry-util-http,oracledb,packaging,pandas,peft,pip,platformdirs,polyleven,prometheus-client,prompt-toolkit,propcache,proto-plus,protobuf,psutil,psycopg2-binary,py-key-value-aio,pyaml,pyarrow,pyasn1,pyasn1-modules,pycparser,pydantic,pydantic-ai,pydantic-ai-slim,pydantic-core,pydantic-evals,pydantic-graph,pydantic-settings,pygments,pyjwt,pyopenssl,pypdf,pyperclip,python-dateutil,python-dotenv,python-multipart,pytz,pyyaml,referencing,regex,requests,requests-oauthlib,rich,rpds-py,safetensors,scikit-learn,scipy,secretstorage,semver,sentence-transformers,sentry-sdk,setuptools,shellingham,six,sniffio,sqlalchemy,sse-starlette,starlette,sympy,tenacity,termcolor,threadpoolctl,tiktoken,tokenizers,torch,tornado,tqdm,transformers,tree-sitter,triton,trl,truststore,typer,typing-extensions,typing-inspection,urllib3,uv,uv-build,uvicorn,wcwidth,websocket-client,websockets,wrapt,xxhash,yarl,zipp", + "packages": "a2a-sdk,accelerate,aiofile,aiohappyeyeballs,aiohttp,aiosignal,aiosqlite,annotated-doc,annotated-types,anthropic,anyio,argcomplete,asyncpg,attrs,authlib,autoevals,azure-core,azure-identity,beartype,cachetools,caio,certifi,cffi,chardet,charset-normalizer,chevron,click,cryptography,datasets,dill,distro,dnspython,docstring-parser,durationpy,einops,email-validator,emoji,exceptiongroup,executing,faiss-cpu,fastapi,fastmcp-slim,fastuuid,filelock,fire,frozenlist,fsspec,genai-prices,google-api-core,google-auth,google-cloud-core,google-cloud-storage,google-crc32c,google-genai,google-resumable-media,googleapis-common-protos,greenlet,griffelib,grpc-google-iam-v1,grpcio,grpcio-status,h11,hf-xet,httpcore,httpcore2,httpx,httpx-sse,httpx2,huggingface-hub,idna,importlib-metadata,jaraco-classes,jaraco-context,jaraco-functools,jeepney,jinja2,jiter,joblib,joserfc,jsonpath-ng,jsonschema,jsonschema-specifications,keyring,kubernetes,langdetect,litellm,logfire,logfire-api,markdown-it-py,markupsafe,maturin,mcp,mdurl,more-itertools,mpmath,msal,msal-extensions,multidict,multiprocess,narwhals,networkx,nltk,numpy,oauthlib,ogx,ogx-api,ogx-client,openai,opentelemetry-api,opentelemetry-distro,opentelemetry-exporter-otlp,opentelemetry-exporter-otlp-proto-common,opentelemetry-exporter-otlp-proto-grpc,opentelemetry-exporter-otlp-proto-http,opentelemetry-instrumentation,opentelemetry-instrumentation-httpx,opentelemetry-proto,opentelemetry-sdk,opentelemetry-semantic-conventions,opentelemetry-util-http,oracledb,packaging,pandas,peft,pip,platformdirs,polyleven,prometheus-client,prompt-toolkit,propcache,proto-plus,protobuf,psutil,psycopg2-binary,py-key-value-aio,pyaml,pyarrow,pyasn1,pyasn1-modules,pycparser,pydantic,pydantic-ai,pydantic-ai-slim,pydantic-core,pydantic-evals,pydantic-graph,pydantic-settings,pygments,pyjwt,pyopenssl,pypdf,pyperclip,python-dateutil,python-dotenv,python-multipart,pytz,pyyaml,referencing,regex,requests,requests-oauthlib,rich,rpds-py,safetensors,scikit-learn,scipy,secretstorage,semver,sentence-transformers,sentry-sdk,setuptools,shellingham,six,sniffio,sqlalchemy,sse-starlette,starlette,structlog,sympy,tenacity,termcolor,threadpoolctl,tiktoken,tokenizers,torch,tornado,tqdm,transformers,tree-sitter,triton,trl,truststore,typer,typing-extensions,typing-inspection,urllib3,uv,uv-build,uvicorn,wcwidth,websocket-client,websockets,wrapt,xxhash,yarl,zipp,zstandard", "os": "linux", "arch": "x86_64,aarch64", "py_version": 312 diff --git a/AGENTS.md b/AGENTS.md index fc3fa8d83..413616ffb 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -234,7 +234,7 @@ src/ #### Imports & Dependencies - Use absolute imports for internal modules: `from authentication import get_auth_dependency` - FastAPI dependencies: `from fastapi import APIRouter, HTTPException, Request, status, Depends` -- Llama Stack imports: `from llama_stack_client import AsyncLlamaStackClient` +- Llama Stack imports: `from ogx_client import AsyncOgxClient` - **ALWAYS** check `pyproject.toml` for existing dependencies before adding new ones - **ALWAYS** verify current library versions in `pyproject.toml` rather than assuming versions - Check `constants.py` for shared constants before defining new ones diff --git a/Makefile b/Makefile index 245cde000..223acf7a0 100644 --- a/Makefile +++ b/Makefile @@ -108,7 +108,7 @@ start-llama-stack-container: build-llama-stack-image ## Start llama-stack contai -e WATSONX_API_KEY \ -e LITELLM_DROP_PARAMS=true \ -e AWS_BEARER_TOKEN_BEDROCK \ - -e LLAMA_STACK_LOGGING=$${LLAMA_STACK_LOGGING:-} \ + -e OGX_LOGGING=$${OGX_LOGGING:-} \ -e FAISS_VECTOR_STORE_ID=$${FAISS_VECTOR_STORE_ID:-} \ -e RH_SERVER_OKP \ -e SOLR_URL \ @@ -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 llama stack run $(LLAMA_STACK_CONFIG) + uv run ogx stack run $(LLAMA_STACK_CONFIG) test-unit: ## Run the unit tests @echo "Running unit tests..." diff --git a/README.md b/README.md index 316098aca..556c59a08 100644 --- a/README.md +++ b/README.md @@ -770,22 +770,29 @@ For the configuration guide, skill authoring instructions, and examples, see the ## Safety Shields -A single Llama Stack configuration file can include multiple safety shields, which are utilized in agent -configurations to monitor input and/or output streams. LCS uses the following naming convention to specify how each safety shield is -utilized: +Safety shields used by `/query`, `/streaming_query`, `/responses`, and `/rlsapi` +are **owned by Lightspeed Core Stack** and configured in `lightspeed-stack.yaml` +(not via the Llama Stack / OGX Safety or Moderations APIs). -1. If the `shield_id` starts with `input_`, it will be used for input only. -1. If the `shield_id` starts with `output_`, it will be used for output only. -1. If the `shield_id` starts with `inout_`, it will be used both for input and output. -1. Otherwise, it will be used for input only. +Supported shield types (`provider_id`): -Additionally, an optional list parameter `shield_ids` can be specified in `/query` and `/streaming_query` endpoints to override which shields are applied. You can use this config to disable shield overrides: +- `question_validity` — topic / off-topic classification +- `redaction` — regex-based PII redaction + +List configured shields with `GET /v1/shields`. Optionally override which +shields apply with the `shield_ids` request field (`null` = all, `[]` = none, +or a list of configured `identifier` values). To forbid client overrides on +`/query` and `/streaming_query`: ```yaml customization: disable_shield_ids_override: true ``` +For configuration details, endpoint application (direct-run vs agent +capabilities), and examples, see the +[Safety Shields Guide](docs/user_doc/shields_guide.md). + ## Authentication See [authentication and authorization](docs/auth.md). diff --git a/docker-compose-library.yaml b/docker-compose-library.yaml index d9293a85b..f4fe486b6 100755 --- a/docker-compose-library.yaml +++ b/docker-compose-library.yaml @@ -56,7 +56,7 @@ services: # AWS Bedrock - AWS_BEARER_TOKEN_BEDROCK=${AWS_BEARER_TOKEN_BEDROCK:-} # Enable debug logging if needed - - LLAMA_STACK_LOGGING=${LLAMA_STACK_LOGGING:-} + - OGX_LOGGING=${OGX_LOGGING:-} # FAISS test and inline RAG config - FAISS_VECTOR_STORE_ID=${FAISS_VECTOR_STORE_ID:-} # Prevent HuggingFace Hub update checks (HTTP 429 rate-limiting in CI from parallel jobs). diff --git a/docker-compose.yaml b/docker-compose.yaml index 01508941a..18e43be2d 100755 --- a/docker-compose.yaml +++ b/docker-compose.yaml @@ -54,7 +54,7 @@ services: # AWS Bedrock - AWS_BEARER_TOKEN_BEDROCK=${AWS_BEARER_TOKEN_BEDROCK:-} # Enable debug logging if needed - - LLAMA_STACK_LOGGING=${LLAMA_STACK_LOGGING:-} + - OGX_LOGGING=${OGX_LOGGING:-} # FAISS test - FAISS_VECTOR_STORE_ID=${FAISS_VECTOR_STORE_ID:-} # Prevent HuggingFace Hub update checks (HTTP 429 rate-limiting in CI from parallel jobs). diff --git a/docs/README.md b/docs/README.md index faf582109..b8cfebb23 100644 --- a/docs/README.md +++ b/docs/README.md @@ -19,6 +19,8 @@ See the full documentation at [`../README.md`](../README.md) or browse sub-pages [Agent skills](https://lightspeed-core.github.io/lightspeed-stack/user_doc/skills_guide.html) +[Safety shields](https://lightspeed-core.github.io/lightspeed-stack/user_doc/shields_guide.html) + [A2A [Agent-to-Agent] Protocol](https://lightspeed-core.github.io/lightspeed-stack/user_doc/a2a_protocol.html) [RAG configuration guide](https://lightspeed-core.github.io/lightspeed-stack/user_doc/rag_guide.html) diff --git a/docs/basic_info/overview.md b/docs/basic_info/overview.md index 981866f07..621a2a5de 100644 --- a/docs/basic_info/overview.md +++ b/docs/basic_info/overview.md @@ -6,9 +6,11 @@ **Lightspeed Core Stack (LCore)** is an enterprise-grade middleware service that provides a robust layer between client applications and AI Large Language Model (LLM) backends. It adds essential enterprise features such as authentication, authorization, quota management, caching, and observability to LLM interactions. -Current version of LCore is built on **Llama Stack** - open-source framework that provides standardized APIs for building LLM applications. Llama Stack offers a unified interface for models, RAG (vector stores), tools, and safety (shields) across different providers. LCore communicates with Llama Stack to orchestrate all LLM operations. +Current version of LCore is built on **OGX (Llama Stack)** - open-source framework that provides standardized APIs for building LLM applications. OGX offers a unified interface for models, RAG (vector stores), and tools across different providers. LCore communicates with OGX to orchestrate all LLM operations. -To enhance LLM responses, LCore leverages **RAG (Retrieval-Augmented Generation)**, which retrieves relevant context from vector databases before generating answers. Llama Stack manages the vector stores, and LCore queries them to inject relevant documentation, knowledge bases, or previous conversations into the LLM prompt. +To enhance LLM responses, LCore leverages **RAG (Retrieval-Augmented Generation)**, which retrieves relevant context from vector databases before generating answers. OGX manages the vector stores, and LCore queries them to inject relevant documentation, knowledge bases, or previous conversations into the LLM prompt. + +LCore also provides **safety shields** such as topic validation and PII redaction. These are configured in LCore and applied on request endpoints before or during agent processing. ### Key Features diff --git a/docs/design/conversation-compaction/conversation-compaction.md b/docs/design/conversation-compaction/conversation-compaction.md index 0e23e7a1e..4098b30db 100644 --- a/docs/design/conversation-compaction/conversation-compaction.md +++ b/docs/design/conversation-compaction/conversation-compaction.md @@ -356,7 +356,7 @@ Example config files go in `examples/`. ## Test patterns - Framework: pytest + pytest-asyncio + pytest-mock. unittest is banned by ruff. -- Mock Llama Stack client: `mocker.AsyncMock(spec=AsyncLlamaStackClient)`. +- Mock Llama Stack client: `mocker.AsyncMock(spec=AsyncOgxClient)`. - Patch at module level: `mocker.patch("utils.responses.compact_conversation_if_needed", ...)`. - Async mocking pattern: see `tests/unit/utils/test_shields.py`. - Config validation tests: see `tests/unit/models/config/`. diff --git a/docs/design/human-in-the-loop/human-in-the-loop.md b/docs/design/human-in-the-loop/human-in-the-loop.md index cb9a6b90d..43299ebc8 100644 --- a/docs/design/human-in-the-loop/human-in-the-loop.md +++ b/docs/design/human-in-the-loop/human-in-the-loop.md @@ -564,7 +564,7 @@ Example config files go in `examples/`. ### Test patterns - Framework: pytest + pytest-asyncio + pytest-mock. unittest is banned by ruff. -- Mock Llama Stack client: `mocker.AsyncMock(spec=AsyncLlamaStackClient)`. +- Mock Llama Stack client: `mocker.AsyncMock(spec=AsyncOgxClient)`. - Patch at module level: `mocker.patch("utils.module.function_name", ...)`. - Async mocking pattern: see `tests/unit/utils/test_shields.py`. - Config validation tests: see `tests/unit/models/config/`. diff --git a/docs/design/llama-stack-config-merge/llama-stack-config-merge-spike.md b/docs/design/llama-stack-config-merge/llama-stack-config-merge-spike.md index 27d02a26a..d013f923a 100644 --- a/docs/design/llama-stack-config-merge/llama-stack-config-merge-spike.md +++ b/docs/design/llama-stack-config-merge/llama-stack-config-merge-spike.md @@ -893,10 +893,10 @@ Summary of validation: ### Findings discovered during PoC -- **`AsyncLlamaStackAsLibraryClient` takes a file path, not a dict.** The +- **`AsyncOGXAsLibraryClient` takes a file path, not a dict.** The initial design assumed we could pass the synthesized configuration to the library client in memory and avoid touching the filesystem. In practice - `llama_stack.core.library_client.AsyncLlamaStackAsLibraryClient` accepts + `ogx.core.library_client.AsyncOGXAsLibraryClient` accepts only a string path (or, in newer versions, a `StackRunConfig` object that is itself built from a parsed YAML file). There is no dict-only entry point in the public API. Consequences for the implementation: @@ -1085,7 +1085,7 @@ new list — they don't need to know a patch syntax. ### Process-model recap (no LCORE supervision of LS) **Library mode**: LCORE process embeds the Llama Stack library client. LCORE -synthesizes `run.yaml` to a file, calls `AsyncLlamaStackAsLibraryClient(path)`, +synthesizes `run.yaml` to a file, calls `AsyncOGXAsLibraryClient(path)`, initializes, serves. One process. **Server mode**: Llama Stack runs as a separate process (container). LCORE diff --git a/docs/design/llama-stack-config-merge/llama-stack-config-merge.md b/docs/design/llama-stack-config-merge/llama-stack-config-merge.md index 59d9a5480..7a315a0fc 100644 --- a/docs/design/llama-stack-config-merge/llama-stack-config-merge.md +++ b/docs/design/llama-stack-config-merge/llama-stack-config-merge.md @@ -182,7 +182,7 @@ lightspeed-stack.yaml (unified mode) Library mode Server mode ──────────── ─────────── Write to deterministic path. Written by LS container's entrypoint - AsyncLlamaStackAsLibraryClient script (same synthesizer, same CLI, + AsyncOGXAsLibraryClient script (same synthesizer, same CLI, reads the path and initializes. auto-detects unified via Python). `llama stack run ` starts LS. LCORE connects by URL. @@ -381,7 +381,7 @@ removed `-g/-i/-o` flags is cleaned up as part of the docs JIRA. from scratch with only high-level keys produces synthesized output with no literal secrets on disk. LS itself resolves env refs to values in-memory at startup via `replace_env_vars()` in - `llama_stack.core.library_client`. + `ogx.core.library_client`. - **`native_override` (and dumb-mode migration output) MAY carry literal secrets**: `native_override` is whatever raw YAML the operator drops in, and `migrate_config_dumb()` lifts an existing diff --git a/docs/devel_doc/ARCHITECTURE.md b/docs/devel_doc/ARCHITECTURE.md index b3bf18635..c0c3224a4 100644 --- a/docs/devel_doc/ARCHITECTURE.md +++ b/docs/devel_doc/ARCHITECTURE.md @@ -24,10 +24,12 @@ **Lightspeed Core Stack (LCORE)** is an enterprise-grade middleware service that provides a robust layer between client applications and AI Large Language Model (LLM) backends. It adds essential enterprise features such as authentication, authorization, quota management, caching, and observability to LLM interactions. -LCore is built on **Llama Stack** - Meta's open-source framework that provides standardized APIs for building LLM applications. Llama Stack offers a unified interface for models, RAG (vector stores), tools, and safety (shields) across different providers. LCore communicates with Llama Stack to orchestrate all LLM operations. +LCore is built on **Llama Stack / OGX** - an open-source framework that provides standardized APIs for building LLM applications. It offers a unified interface for models, RAG (vector stores), and tools across different providers. LCore communicates with the stack to orchestrate all LLM operations. To enhance LLM responses, LCore leverages **RAG (Retrieval-Augmented Generation)**, which retrieves relevant context from vector databases before generating answers. Llama Stack manages the vector stores, and LCore queries them to inject relevant documentation, knowledge bases, or previous conversations into the LLM prompt. +To keep requests on-topic and protect sensitive data, LCore applies **safety shields**, which validate user questions and redact PII from model traffic. Shields are owned by LCore and configured in the service configuration. + ### 1.2 Key Features - **Multi-Provider Support**: Works with multiple LLM providers (Ollama, OpenAI, Watsonx, etc.) @@ -61,6 +63,7 @@ To enhance LLM responses, LCore leverages **RAG (Retrieval-Augmented Generation) │ ┌───────────────────────────────────────────────────┐ │ │ │ Request Processing │ │ │ │ • LLM Orchestration (via Llama Stack) │ │ +│ │ • Safety Shields │ │ │ │ • Tool Integration (MCP servers) │ │ │ │ • RAG & Context Management │ │ │ └───────────────────────────────────────────────────┘ │ @@ -79,8 +82,8 @@ To enhance LLM responses, LCore leverages **RAG (Retrieval-Augmented Generation) │ │ │ • Models & LLMs │ │ • RAG Stores │ - │ • Shields │ - └────────┬─────────┘ + └──────────────────┘ + │ │ (manages & invokes) ▼ ┌──────────────────┐ @@ -232,7 +235,6 @@ The system defines 30+ actions that can be authorized. Examples (see `docs/auth. - **Models**: List available LLM models - **Responses**: Generate LLM responses (OpenAI-compatible) - **Conversations**: Manage conversation history -- **Shields**: List and apply guardrails (content filtering, safety checks) - **Vector Stores**: Access RAG databases for context injection - **Toolgroups**: Register MCP servers as tools @@ -397,11 +399,12 @@ Here's how a real query flows through the system: 4. **Quota Check** - User has 50,000 tokens available ✅ 5. **Model Selection** - Use configured default model (e.g., `meta-llama/Llama-3.1-8B-Instruct`) 6. **Context Building** - Retrieve conversation history, query RAG vector stores for relevant docs, determine available MCP tools -7. **Llama Stack Call** - Send complete request with system prompt, RAG context, MCP tools, and shields -8. **LLM Processing** - Llama Stack generates response, may invoke MCP tools, returns token counts -9. **Post-Processing** - Apply shields, generate conversation summary if new -10. **Store Results** - Save to Cache DB, User DB, consume quota, update metrics -11. **Return Response** - Complete LLM response with referenced documents, token usage, and remaining quota +7. **Shield moderation** - LCore-owned direct-run moderation (and agent capabilities where applicable) using shields configured in LCORE config +8. **Llama Stack / agent call** - Send request with system prompt, RAG context, and MCP tools +9. **LLM Processing** - Stack / agent generates response, may invoke MCP tools, returns token counts +10. **Post-Processing** - Generate conversation summary if new +11. **Store Results** - Save to Cache DB, User DB, consume quota, update metrics +12. **Return Response** - Complete LLM response with referenced documents, token usage, and remaining quota **Key Takeaways:** - RAG enhances responses with relevant documentation @@ -500,8 +503,8 @@ This section documents the REST API endpoints exposed by LCore for client intera - `GET /v1/mcp-servers` - List all registered MCP servers (static and dynamic) - `DELETE /v1/mcp-servers/{name}` - Unregister a dynamically registered MCP server -**List Shields:** `GET /shields` -- Returns available guardrails +**List Shields:** `GET /v1/shields` +- Returns list of shields configured in LCORE **List RAG Databases:** `GET /rags` - Returns configured vector stores diff --git a/docs/devel_doc/openapi.json b/docs/devel_doc/openapi.json index f66ea5cbd..073109e41 100644 --- a/docs/devel_doc/openapi.json +++ b/docs/devel_doc/openapi.json @@ -294,11 +294,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -494,11 +494,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -693,11 +693,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -1041,26 +1041,6 @@ } } } - }, - "503": { - "description": "Service unavailable", - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ServiceUnavailableResponse" - }, - "examples": { - "kubernetes api": { - "value": { - "detail": { - "cause": "Failed to connect to Kubernetes API: Service Unavailable (status 503)", - "response": "Unable to connect to Kubernetes API" - } - } - } - } - } - } } } }, @@ -1069,7 +1049,7 @@ "mcp-servers" ], "summary": "Register Mcp Server Handler", - "description": "Register an MCP server dynamically at runtime.\n\nAdds the MCP server to the runtime configuration and registers it\nas a toolgroup with Llama Stack so it becomes available for queries.\n\n### Parameters:\n- request: Model containing attributes to dynamically registering an MCP server.\n- auth: Authentication tuple from the auth dependency (used by middleware).\n- body: Headers that should be passed to MCP servers.\n\n### Raises:\n- HTTPException: On duplicate name, Llama Stack connection error, or\n registration failure.\n\n### Returns:\n- MCPServerRegistrationResponse: Details of the newly registered server.", + "description": "Register an MCP server dynamically at runtime.\n\nAdds the MCP server to the runtime configuration so it becomes available\nfor queries.\n\n### Parameters:\n- request: Model containing attributes to dynamically registering an MCP server.\n- auth: Authentication tuple from the auth dependency (used by middleware).\n- body: Headers that should be passed to MCP servers.\n\n### Raises:\n- HTTPException: On duplicate name or registration failure.\n\n### Returns:\n- MCPServerRegistrationResponse: Details of the newly registered server.", "operationId": "register_mcp_server_handler_v1_mcp_servers_post", "requestBody": { "content": { @@ -1242,34 +1222,6 @@ } } }, - "503": { - "description": "Service unavailable", - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ServiceUnavailableResponse" - }, - "examples": { - "llama stack": { - "value": { - "detail": { - "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" - } - } - }, - "kubernetes api": { - "value": { - "detail": { - "cause": "Failed to connect to Kubernetes API: Service Unavailable (status 503)", - "response": "Unable to connect to Kubernetes API" - } - } - } - } - } - } - }, "422": { "description": "Validation Error", "content": { @@ -1289,7 +1241,7 @@ "mcp-servers" ], "summary": "Delete Mcp Server Handler", - "description": "Unregister a dynamically registered MCP server.\n\nRemoves the MCP server from the runtime configuration and unregisters\nits toolgroup from Llama Stack. Only servers registered via the API\ncan be deleted; statically configured servers cannot be removed.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- auth: Authentication tuple from the auth dependency (used by middleware).\n- name: MCP server name\n\n### Raises:\n- HTTPException: If the server is not found, is statically configured, or\n Llama Stack unregistration fails.\n\n### Returns:\n- MCPServerDeleteResponse: Confirmation of the deletion.", + "description": "Unregister a dynamically registered MCP server.\n\nRemoves the MCP server from the runtime configuration. Only servers\nregistered via the API can be deleted; statically configured servers\ncannot be removed.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- auth: Authentication tuple from the auth dependency (used by middleware).\n- name: MCP server name\n\n### Raises:\n- HTTPException: If the server is not found or is statically configured.\n\n### Returns:\n- MCPServerDeleteResponse: Confirmation of the deletion.", "operationId": "delete_mcp_server_handler_v1_mcp_servers__name__delete", "parameters": [ { @@ -1453,34 +1405,6 @@ } } }, - "503": { - "description": "Service unavailable", - "content": { - "application/json": { - "examples": { - "llama stack": { - "value": { - "detail": { - "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" - } - } - }, - "kubernetes api": { - "value": { - "detail": { - "cause": "Failed to connect to Kubernetes API: Service Unavailable (status 503)", - "response": "Unable to connect to Kubernetes API" - } - } - } - }, - "schema": { - "$ref": "#/components/schemas/ServiceUnavailableResponse" - } - } - } - }, "422": { "description": "Validation Error", "content": { @@ -1500,7 +1424,7 @@ "shields" ], "summary": "Shields Endpoint Handler", - "description": "Handle requests to the /shields endpoint.\n\nProcess GET requests to the /shields endpoint, returning a list of available\nshields from the Llama Stack service.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- auth: Authentication tuple from the auth dependency (used by middleware).\n\n### Raises:\n- HTTPException: with status 401 for unauthorized access.\n- HTTPException: with status 403 if permission is denied.\n- HTTPException: with status 500 and a detail object containing `response`\n and `cause` when service configuration is wrong or incomplete.\n- HTTPException: with status 503 and a detail object containing `response`\n and `cause` when unable to connect to Llama Stack.\n\n### Returns:\n- ShieldsResponse: An object containing the list of available shields.", + "description": "Handle requests to the /shields endpoint.\n\nProcess GET requests to the /shields endpoint, returning a list of available\nshields from Lightspeed Core Stack configuration.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- auth: Authentication tuple from the auth dependency (used by middleware).\n\n### Raises:\n- HTTPException: with status 401 for unauthorized access.\n- HTTPException: with status 403 if permission is denied.\n- HTTPException: with status 500 and a detail object containing `response`\n and `cause` when service configuration is wrong or incomplete.\n\n### Returns:\n- ShieldsResponse: An object containing the list of available shields.", "operationId": "shields_endpoint_handler_v1_shields_get", "responses": { "200": { @@ -1513,10 +1437,27 @@ "example": { "shields": [ { - "identifier": "lightspeed_question_validity-shield", - "params": {}, - "provider_id": "lightspeed_question_validity", - "provider_resource_id": "lightspeed_question_validity-shield", + "config": { + "invalid_question_response": "I can only answer questions about the product.", + "model_id": "openai/gpt-4o-mini", + "model_prompt": "Is this question valid?" + }, + "name": "question-validity", + "provider_id": "question_validity", + "type": "shield" + }, + { + "config": { + "case_sensitive": false, + "rules": [ + { + "pattern": "\\b\\d{3}-\\d{2}-\\d{4}\\b", + "replacement": "[REDACTED]" + } + ] + }, + "name": "pii-redaction", + "provider_id": "redaction", "type": "shield" } ] @@ -1639,34 +1580,6 @@ } } } - }, - "503": { - "description": "Service unavailable", - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ServiceUnavailableResponse" - }, - "examples": { - "llama stack": { - "value": { - "detail": { - "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" - } - } - }, - "kubernetes api": { - "value": { - "detail": { - "cause": "Failed to connect to Kubernetes API: Service Unavailable (status 503)", - "response": "Unable to connect to Kubernetes API" - } - } - } - } - } - } } } } @@ -1834,11 +1747,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -2040,11 +1953,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -2240,11 +2153,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -2431,11 +2344,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -2688,11 +2601,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -2940,11 +2853,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -3169,11 +3082,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -3355,11 +3268,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -3558,11 +3471,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -3600,7 +3513,7 @@ "vector-stores" ], "summary": "List Vector Stores", - "description": "List all vector stores.\n\nParameters:\n request: The incoming HTTP request.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoresListResponse: List of all vector stores.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "List all vector stores.\n\nParameters:\n request: The incoming HTTP request.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoresListResponse: List of all vector stores.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "list_vector_stores_v1_vector_stores_get", "responses": { "200": { @@ -3761,11 +3674,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -3788,7 +3701,7 @@ "vector-stores" ], "summary": "Create Vector Store", - "description": "Create a new vector store.\n\nParameters:\n request: The incoming HTTP request.\n auth: Authentication tuple from the auth dependency.\n body: Vector store creation parameters.\n\nReturns:\n VectorStoreResponse: The created vector store object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "Create a new vector store.\n\nParameters:\n request: The incoming HTTP request.\n auth: Authentication tuple from the auth dependency.\n body: Vector store creation parameters.\n\nReturns:\n VectorStoreResponse: The created vector store object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "create_vector_store_v1_vector_stores_post", "requestBody": { "content": { @@ -3970,11 +3883,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -4009,7 +3922,7 @@ "vector-stores" ], "summary": "Get Vector Store", - "description": "Retrieve a vector store by ID.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store to retrieve.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoreResponse: The vector store object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "Retrieve a vector store by ID.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store to retrieve.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoreResponse: The vector store object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "get_vector_store_v1_vector_stores__vector_store_id__get", "parameters": [ { @@ -4189,11 +4102,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -4229,7 +4142,7 @@ "vector-stores" ], "summary": "Update Vector Store", - "description": "Update a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store to update.\n auth: Authentication tuple from the auth dependency.\n body: Vector store update parameters.\n\nReturns:\n VectorStoreResponse: The updated vector store object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "Update a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store to update.\n auth: Authentication tuple from the auth dependency.\n body: Vector store update parameters.\n\nReturns:\n VectorStoreResponse: The updated vector store object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "update_vector_store_v1_vector_stores__vector_store_id__put", "parameters": [ { @@ -4419,11 +4332,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -4459,7 +4372,7 @@ "vector-stores" ], "summary": "Delete Vector Store", - "description": "Delete a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store to delete.\n auth: Authentication tuple from the auth dependency.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack\n\nReturns:\n VectorStoreDeleteResponse: Delete outcome for the requested vector store.", + "description": "Delete a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store to delete.\n auth: Authentication tuple from the auth dependency.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX\n\nReturns:\n VectorStoreDeleteResponse: Delete outcome for the requested vector store.", "operationId": "delete_vector_store_v1_vector_stores__vector_store_id__delete", "parameters": [ { @@ -4620,11 +4533,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -4662,7 +4575,7 @@ "vector-stores" ], "summary": "Create File", - "description": "Upload a file.\n\nParameters:\n request: The incoming HTTP request.\n auth: Authentication tuple from the auth dependency.\n file: The file to upload.\n\nReturns:\n FileResponse: The uploaded file object.\n\nRaises:\n HTTPException:\n - 400: Bad request (e.g., file too large, invalid format)\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "Upload a file.\n\nParameters:\n request: The incoming HTTP request.\n auth: Authentication tuple from the auth dependency.\n file: The file to upload.\n\nReturns:\n FileResponse: The uploaded file object.\n\nRaises:\n HTTPException:\n - 400: Bad request (e.g., file too large, invalid format)\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "create_file_v1_files_post", "requestBody": { "content": { @@ -4845,11 +4758,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -4884,7 +4797,7 @@ "vector-stores" ], "summary": "Add File To Vector Store", - "description": "Add a file to a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n auth: Authentication tuple from the auth dependency.\n body: File addition parameters.\n\nReturns:\n VectorStoreFileResponse: The vector store file object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store or file not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "Add a file to a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n auth: Authentication tuple from the auth dependency.\n body: File addition parameters.\n\nReturns:\n VectorStoreFileResponse: The vector store file object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store or file not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "add_file_to_vector_store_v1_vector_stores__vector_store_id__files_post", "parameters": [ { @@ -5069,11 +4982,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -5109,7 +5022,7 @@ "vector-stores" ], "summary": "List Vector Store Files", - "description": "List files in a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoreFilesListResponse: List of files in the vector store.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "List files in a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoreFilesListResponse: List of files in the vector store.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: Vector store not found\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "list_vector_store_files_v1_vector_stores__vector_store_id__files_get", "parameters": [ { @@ -5294,11 +5207,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -5336,7 +5249,7 @@ "vector-stores" ], "summary": "Get Vector Store File", - "description": "Retrieve a file from a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n file_id: ID of the file.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoreFileResponse: The vector store file object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: File not found in vector store\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack", + "description": "Retrieve a file from a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n file_id: ID of the file.\n auth: Authentication tuple from the auth dependency.\n\nReturns:\n VectorStoreFileResponse: The vector store file object.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 404: File not found in vector store\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX", "operationId": "get_vector_store_file_v1_vector_stores__vector_store_id__files__file_id__get", "parameters": [ { @@ -5520,11 +5433,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -5560,7 +5473,7 @@ "vector-stores" ], "summary": "Delete Vector Store File", - "description": "Delete a file from a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n file_id: ID of the file to delete.\n auth: Authentication tuple from the auth dependency.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to Llama Stack\n\nReturns:\n VectorStoreFileDeleteResponse: Delete outcome for the requested file.", + "description": "Delete a file from a vector store.\n\nParameters:\n request: The incoming HTTP request.\n vector_store_id: ID of the vector store.\n file_id: ID of the file to delete.\n auth: Authentication tuple from the auth dependency.\n\nRaises:\n HTTPException:\n - 401: Authentication failed\n - 403: Authorization failed\n - 500: Lightspeed Stack configuration not loaded\n - 503: Unable to connect to OGX\n\nReturns:\n VectorStoreFileDeleteResponse: Delete outcome for the requested file.", "operationId": "delete_vector_store_file_v1_vector_stores__vector_store_id__files__file_id__delete", "parameters": [ { @@ -5730,11 +5643,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -5772,7 +5685,7 @@ "query" ], "summary": "Query Endpoint Handler", - "description": "Handle request to the /query endpoint using Responses API.\n\nProcesses a POST request to a query endpoint, forwarding the\nuser's query to a selected Llama Stack LLM and returning the generated response.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- query_request: Request to the LLM.\n- auth: Auth context tuple resolved from the authentication dependency.\n- mcp_headers: Headers that should be passed to MCP servers.\n\n### Returns:\n- QueryResponse: Contains the conversation ID and the LLM-generated response.\n\n### Raises:\n- HTTPException:\n- 401: Unauthorized - Missing or invalid credentials\n- 403: Forbidden - Insufficient permissions or model override not allowed\n- 404: Not Found - Conversation, model, or provider not found\n- 413: Prompt too long - Prompt exceeded model's context window size\n- 422: Unprocessable Entity - Request validation failed\n- 429: Quota limit exceeded - The token quota for model or user has been exceeded\n- 500: Internal Server Error - Configuration not loaded or other server errors\n- 503: Service Unavailable - Unable to connect to Llama Stack backend", + "description": "Handle request to the /query endpoint using Responses API.\n\nProcesses a POST request to a query endpoint, forwarding the\nuser's query to a selected Llama Stack LLM and returning the generated response.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- query_request: Request to the LLM.\n- auth: Auth context tuple resolved from the authentication dependency.\n- mcp_headers: Headers that should be passed to MCP servers.\n\n### Returns:\n- QueryResponse: Contains the conversation ID and the LLM-generated response.\n\n### Raises:\n- HTTPException:\n- 401: Unauthorized - Missing or invalid credentials\n- 403: Forbidden - Insufficient permissions or model override not allowed\n- 404: Not Found - Conversation, model, or provider not found\n- 413: Prompt too long - Prompt exceeded model's context window size\n- 422: Unprocessable Entity - Request validation failed\n- 429: Quota limit exceeded - The token quota for model or user has been exceeded\n- 500: Internal Server Error - Configuration not loaded or other server errors\n- 503: Service Unavailable - Unable to connect to OGX backend", "operationId": "query_endpoint_handler_v1_query_post", "requestBody": { "content": { @@ -6153,11 +6066,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -6182,7 +6095,7 @@ "streaming_query" ], "summary": "Streaming Query Endpoint Handler", - "description": "Handle request to the /streaming_query endpoint using Responses API.\n\nReturns a streaming response using Server-Sent Events (SSE) format with\ncontent type text/event-stream.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- query_request: Request to the LLM.\n- auth: Auth context tuple resolved from the authentication dependency.\n- mcp_headers: Headers that should be passed to MCP servers.\n\n### Returns:\n- SSE-formatted events for the query lifecycle.\n\n### Raises:\n- HTTPException:\n- 401: Unauthorized - Missing or invalid credentials\n- 403: Forbidden - Insufficient permissions or model override not allowed\n- 404: Not Found - Conversation, model, or provider not found\n- 413: Prompt too long - Prompt exceeded model's context window size\n- 422: Unprocessable Entity - Request validation failed\n- 429: Quota limit exceeded - The token quota for model or user has been exceeded\n- 500: Internal Server Error - Configuration not loaded or other server errors\n- 503: Service Unavailable - Unable to connect to Llama Stack backend", + "description": "Handle request to the /streaming_query endpoint using Responses API.\n\nReturns a streaming response using Server-Sent Events (SSE) format with\ncontent type text/event-stream.\n\n### Parameters:\n- request: The incoming HTTP request (used by middleware).\n- query_request: Request to the LLM.\n- auth: Auth context tuple resolved from the authentication dependency.\n- mcp_headers: Headers that should be passed to MCP servers.\n\n### Returns:\n- SSE-formatted events for the query lifecycle.\n\n### Raises:\n- HTTPException:\n- 401: Unauthorized - Missing or invalid credentials\n- 403: Forbidden - Insufficient permissions or model override not allowed\n- 404: Not Found - Conversation, model, or provider not found\n- 413: Prompt too long - Prompt exceeded model's context window size\n- 422: Unprocessable Entity - Request validation failed\n- 429: Quota limit exceeded - The token quota for model or user has been exceeded\n- 500: Internal Server Error - Configuration not loaded or other server errors\n- 503: Service Unavailable - Unable to connect to OGX backend", "operationId": "streaming_query_endpoint_handler_v1_streaming_query_post", "requestBody": { "content": { @@ -6530,11 +6443,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -8313,11 +8226,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -8566,11 +8479,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -8803,11 +8716,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -9051,11 +8964,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -9979,7 +9892,7 @@ "responses" ], "summary": "Responses Endpoint Handler", - "description": "Handle request to the /responses endpoint using Responses API (LCORE specification).\n\nProcesses a POST request to the responses endpoint, forwarding the\nuser's request to a selected Llama Stack LLM and returning the generated response\nfollowing the LCORE OpenAPI specification.\n\nReturns:\n ResponsesResponse: Contains the response following LCORE specification (non-streaming).\n StreamingResponse: SSE-formatted streaming response with enriched events (streaming).\n - response.created event includes conversation attribute\n - response.completed event includes available_quotas attribute\n\nRaises:\n HTTPException:\n - 401: Unauthorized - Missing or invalid credentials\n - 403: Forbidden - Insufficient permissions or model override not allowed\n - 404: Not Found - Conversation, model, or provider not found\n - 413: Prompt too long - Prompt exceeded model's context window size\n - 422: Unprocessable Entity - Request validation failed\n - 429: Quota limit exceeded - The token quota for model or user has been exceeded\n - 500: Internal Server Error - Configuration not loaded or other server errors\n - 503: Service Unavailable - Unable to connect to Llama Stack backend", + "description": "Handle request to the /responses endpoint using Responses API (LCORE specification).\n\nProcesses a POST request to the responses endpoint, forwarding the\nuser's request to a selected Llama Stack LLM and returning the generated response\nfollowing the LCORE OpenAPI specification.\n\nReturns:\n ResponsesResponse: Contains the response following LCORE specification (non-streaming).\n StreamingResponse: SSE-formatted streaming response with enriched events (streaming).\n - response.created event includes conversation attribute\n - response.completed event includes available_quotas attribute\n\nRaises:\n HTTPException:\n - 401: Unauthorized - Missing or invalid credentials\n - 403: Forbidden - Insufficient permissions or model override not allowed\n - 404: Not Found - Conversation, model, or provider not found\n - 413: Prompt too long - Prompt exceeded model's context window size\n - 422: Unprocessable Entity - Request validation failed\n - 429: Quota limit exceeded - The token quota for model or user has been exceeded\n - 500: Internal Server Error - Configuration not loaded or other server errors\n - 503: Service Unavailable - Unable to connect to OGX backend", "operationId": "responses_endpoint_handler_v1_responses_post", "requestBody": { "content": { @@ -10406,11 +10319,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -10748,11 +10661,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -10900,11 +10813,11 @@ "$ref": "#/components/schemas/ServiceUnavailableResponse" }, "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -11353,11 +11266,11 @@ "content": { "application/json": { "examples": { - "llama stack": { + "ogx": { "value": { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" } } }, @@ -12819,6 +12732,181 @@ "title": "CORSConfiguration", "description": "CORS configuration.\n\nCORS or 'Cross-Origin Resource Sharing' refers to the situations when a\nfrontend running in a browser has JavaScript code that communicates with a\nbackend, and the backend is in a different 'origin' than the frontend.\n\nUseful resources:\n\n - [CORS in FastAPI](https://fastapi.tiangolo.com/tutorial/cors/)\n - [Wikipedia article](https://en.wikipedia.org/wiki/Cross-origin_resource_sharing)\n - [What is CORS?](https://dev.to/akshay_chauhan/what-is-cors-explained-8f1)" }, + "CatalogModel": { + "properties": { + "identifier": { + "type": "string", + "title": "Identifier", + "description": "Model identifier" + }, + "metadata": { + "additionalProperties": true, + "type": "object", + "title": "Metadata", + "description": "Provider-specific metadata excluding core catalog fields" + }, + "api_model_type": { + "type": "string", + "title": "Api Model Type", + "description": "API model type (typically mirrors model_type)" + }, + "provider_id": { + "type": "string", + "title": "Provider Id", + "description": "Provider identifier" + }, + "type": { + "type": "string", + "title": "Type", + "description": "Object type, always 'model'", + "default": "model" + }, + "provider_resource_id": { + "type": "string", + "title": "Provider Resource Id", + "description": "Provider-native resource identifier for the model", + "default": "" + }, + "model_type": { + "type": "string", + "title": "Model Type", + "description": "Model type such as 'llm' or 'embedding'" + } + }, + "type": "object", + "required": [ + "identifier", + "api_model_type", + "provider_id", + "model_type" + ], + "title": "CatalogModel", + "description": "Normalized model entry used by ``/models`` and internal model resolution.\n\nUnifies OpenAI-style, Anthropic, and Google ``models.list()`` payloads into\none catalog shape." + }, + "CatalogShield": { + "properties": { + "name": { + "type": "string", + "title": "Name", + "description": "Unique, user-facing name of the shield instance" + }, + "provider_id": { + "type": "string", + "enum": [ + "question_validity", + "redaction" + ], + "title": "Provider Id", + "description": "Shield provider / type discriminator" + }, + "type": { + "type": "string", + "const": "shield", + "title": "Type", + "description": "Catalog entry type; always shield", + "default": "shield" + }, + "config": { + "additionalProperties": true, + "type": "object", + "title": "Config", + "description": "Type-specific shield configuration" + } + }, + "type": "object", + "required": [ + "name", + "provider_id", + "config" + ], + "title": "CatalogShield", + "description": "Shield entry in the ``/shields`` catalog response.\n\nAttributes:\n name: Unique, user-facing name identifying this shield instance.\n provider_id: Shield provider / type discriminator.\n type: Catalog entry type; always shield.\n config: Type-specific shield configuration." + }, + "CatalogTool": { + "properties": { + "identifier": { + "type": "string", + "title": "Identifier" + }, + "description": { + "type": "string", + "title": "Description" + }, + "parameters": { + "items": { + "$ref": "#/components/schemas/CatalogToolParameter" + }, + "type": "array", + "title": "Parameters" + }, + "provider_id": { + "type": "string", + "title": "Provider Id" + }, + "toolgroup_id": { + "type": "string", + "title": "Toolgroup Id" + }, + "server_source": { + "type": "string", + "title": "Server Source" + }, + "type": { + "type": "string", + "title": "Type", + "default": "tool" + } + }, + "type": "object", + "required": [ + "identifier", + "description", + "parameters", + "provider_id", + "toolgroup_id", + "server_source" + ], + "title": "CatalogTool", + "description": "Tool entry in the ``/tools`` catalog response." + }, + "CatalogToolParameter": { + "properties": { + "name": { + "type": "string", + "title": "Name" + }, + "description": { + "type": "string", + "title": "Description" + }, + "parameter_type": { + "type": "string", + "title": "Parameter Type" + }, + "required": { + "type": "boolean", + "title": "Required", + "default": false + }, + "default": { + "anyOf": [ + {}, + { + "type": "null" + } + ], + "title": "Default" + } + }, + "type": "object", + "required": [ + "name", + "description", + "parameter_type" + ], + "title": "CatalogToolParameter", + "description": "Parameter entry for a tool in the ``/tools`` catalog response." + }, "ClientCredentialsOAuthFlow": { "properties": { "refreshUrl": { @@ -13067,6 +13155,28 @@ "$ref": "#/components/schemas/SavedPromptsConfiguration", "title": "Saved prompts configuration", "description": "Configuration for saved prompts feature limits including maximum prompts per user, display name length, and content length." + }, + "shields": { + "items": { + "oneOf": [ + { + "$ref": "#/components/schemas/QuestionValidityShieldConfiguration" + }, + { + "$ref": "#/components/schemas/RedactionShieldConfiguration" + } + ], + "discriminator": { + "propertyName": "provider_id", + "mapping": { + "question_validity": "#/components/schemas/QuestionValidityShieldConfiguration", + "redaction": "#/components/schemas/RedactionShieldConfiguration" + } + } + }, + "type": "array", + "title": "Shields configuration", + "description": "List of pydantic-ai-lightspeed agent guardrail shields (question validity and PII redaction). Each entry has a unique 'name', a 'provider_id' ('question_validity' or 'redaction'), and a type-specific 'config'." } }, "additionalProperties": false, @@ -14496,8 +14606,7 @@ "computer_call_output.output.image_url", "file_search_call.results", "message.input_image.image_url", - "message.output_text.logprobs", - "reasoning.encrypted_content" + "message.output_text.logprobs" ] }, "InferenceConfiguration": { @@ -15566,8 +15675,7 @@ "properties": { "models": { "items": { - "additionalProperties": true, - "type": "object" + "$ref": "#/components/schemas/CatalogModel" }, "type": "array", "title": "Models", @@ -15921,7 +16029,8 @@ "filename", "start_index" ], - "title": "OpenAIResponseAnnotationContainerFileCitation" + "title": "OpenAIResponseAnnotationContainerFileCitation", + "description": "Container file citation annotation referencing a file within a container." }, "OpenAIResponseAnnotationFileCitation": { "properties": { @@ -15975,7 +16084,8 @@ "file_id", "index" ], - "title": "OpenAIResponseAnnotationFilePath" + "title": "OpenAIResponseAnnotationFilePath", + "description": "File path annotation referencing a generated file in response content." }, "OpenAIResponseContentPartRefusal": { "properties": { @@ -16347,7 +16457,8 @@ "required", "none" ], - "title": "OpenAIResponseInputToolChoiceMode" + "title": "OpenAIResponseInputToolChoiceMode", + "description": "Enumeration of simple tool choice modes for response generation." }, "OpenAIResponseInputToolChoiceWebSearch": { "properties": { @@ -16788,7 +16899,8 @@ "required": [ "text" ], - "title": "OpenAIResponseOutputMessageContentOutputText" + "title": "OpenAIResponseOutputMessageContentOutputText", + "description": "Text content within an output message of an OpenAI response." }, "OpenAIResponseOutputMessageFileSearchToolCall": { "properties": { @@ -17014,6 +17126,113 @@ "title": "OpenAIResponseOutputMessageMCPListTools", "description": "MCP list tools output message containing available tools from an MCP server.\n\n:param id: Unique identifier for this MCP list tools operation\n:param type: Tool call type identifier, always \"mcp_list_tools\"\n:param server_label: Label identifying the MCP server providing the tools\n:param tools: List of available tools provided by the MCP server" }, + "OpenAIResponseOutputMessageReasoningContent": { + "properties": { + "text": { + "type": "string", + "title": "Text", + "description": "The reasoning text content from the model." + }, + "type": { + "type": "string", + "const": "reasoning_text", + "title": "Type", + "description": "The type identifier, always 'reasoning_text'.", + "default": "reasoning_text" + } + }, + "type": "object", + "required": [ + "text" + ], + "title": "OpenAIResponseOutputMessageReasoningContent", + "description": "Reasoning text from the model." + }, + "OpenAIResponseOutputMessageReasoningItem": { + "properties": { + "id": { + "type": "string", + "title": "Id", + "description": "Unique identifier for the reasoning output item." + }, + "summary": { + "items": { + "$ref": "#/components/schemas/OpenAIResponseOutputMessageReasoningSummary" + }, + "type": "array", + "title": "Summary", + "description": "Summary of the reasoning output." + }, + "type": { + "type": "string", + "const": "reasoning", + "title": "Type", + "description": "The type identifier, always 'reasoning'.", + "default": "reasoning" + }, + "content": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/OpenAIResponseOutputMessageReasoningContent" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Content", + "description": "The reasoning content from the model." + }, + "status": { + "anyOf": [ + { + "type": "string", + "enum": [ + "in_progress", + "completed", + "incomplete" + ] + }, + { + "type": "null" + } + ], + "title": "Status", + "description": "The status of the reasoning output." + } + }, + "type": "object", + "required": [ + "id", + "summary" + ], + "title": "OpenAIResponseOutputMessageReasoningItem", + "description": "Reasoning output from the model, representing the model's thinking process." + }, + "OpenAIResponseOutputMessageReasoningSummary": { + "properties": { + "text": { + "type": "string", + "title": "Text", + "description": "The summary text of the reasoning output." + }, + "type": { + "type": "string", + "const": "summary_text", + "title": "Type", + "description": "The type identifier, always 'summary_text'.", + "default": "summary_text" + } + }, + "type": "object", + "required": [ + "text" + ], + "title": "OpenAIResponseOutputMessageReasoningSummary", + "description": "A summary of reasoning output from the model." + }, "OpenAIResponseOutputMessageWebSearchToolCall": { "properties": { "id": { @@ -17116,6 +17335,23 @@ } ], "title": "Effort" + }, + "summary": { + "anyOf": [ + { + "type": "string", + "enum": [ + "auto", + "concise", + "detailed" + ] + }, + { + "type": "null" + } + ], + "title": "Summary", + "description": "Summary mode for reasoning output. One of 'auto', 'concise', or 'detailed'." } }, "type": "object", @@ -17133,11 +17369,27 @@ "type": "null" } ] + }, + "verbosity": { + "anyOf": [ + { + "type": "string", + "enum": [ + "low", + "medium", + "high" + ] + }, + { + "type": "null" + } + ], + "title": "Verbosity" } }, "type": "object", "title": "OpenAIResponseText", - "description": "Text response configuration for OpenAI responses.\n\n:param format: (Optional) Text format configuration specifying output format requirements" + "description": "Text response configuration for OpenAI responses.\n\n:param format: (Optional) Text format configuration specifying output format requirements\n:param verbosity: (Optional) Controls response verbosity level" }, "OpenAIResponseTextFormat": { "properties": { @@ -18307,10 +18559,10 @@ } ], "title": "Shield Ids", - "description": "Optional list of safety shield IDs to apply. If None, all configured shields are used. ", + "description": "Optional list of configured shield names to apply. If None, all configured shields are used.", "examples": [ - "llama-guard", - "custom-shield" + "topic-guard", + "pii-redaction" ] }, "solr": { @@ -18349,7 +18601,7 @@ "query" ], "title": "QueryRequest", - "description": "Model representing a request for the LLM (Language Model).\n\nAttributes:\n query: The query string.\n conversation_id: The optional conversation ID (UUID).\n provider: The optional provider.\n model: The optional model.\n system_prompt: The optional system prompt.\n attachments: The optional attachments.\n no_tools: Whether to bypass all tools and MCP servers (default: False).\n generate_topic_summary: Whether to generate topic summary for new conversations.\n media_type: The optional media type for response format (application/json or text/plain).\n vector_store_ids: The optional list of specific vector store IDs to query for RAG.\n shield_ids: The optional list of safety shield IDs to apply.\n solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict.", + "description": "Model representing a request for the LLM (Language Model).\n\nAttributes:\n query: The query string.\n conversation_id: The optional conversation ID (UUID).\n provider: The optional provider.\n model: The optional model.\n system_prompt: The optional system prompt.\n attachments: The optional attachments.\n no_tools: Whether to bypass all tools and MCP servers (default: False).\n generate_topic_summary: Whether to generate topic summary for new conversations.\n media_type: The optional media type for response format (application/json or text/plain).\n vector_store_ids: The optional list of specific vector store IDs to query for RAG.\n shield_ids: The optional list of configured shield names to apply.\n solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict.", "examples": [ { "attachments": [ @@ -18538,6 +18790,63 @@ } ] }, + "QuestionValidityConfig": { + "properties": { + "model_id": { + "type": "string", + "title": "Model id", + "description": "The model_id to use for the guard" + }, + "model_prompt": { + "type": "string", + "title": "Model prompt", + "description": "The default prompt sent to the LLM used to validate the Users' question.", + "default": "\nInstructions:\n- You are a question classifying tool\n- You are an expert in kubernetes and openshift\n- Your job is to determine where or a user's question is related to kubernetes and/or openshift technologies and to provide a one-word response.\n- If a question appears to be related to kubernetes or openshift technologies, answer with the word ${allowed}, otherwise answer with the word ${rejected}.\n- Do not explain your answer, just provide the one-word response. Do not give any other response.\n- If the given question is an empty string, answer with the word ${rejected}\n\n\nExample Question:\nWhy is the sky blue?\nExample Response:\n${rejected}\n\nExample Question:\nWhy is the grass green?\nExample Response:\n${rejected}\n\nExample Question:\nWhy is sand yellow?\nExample Response:\n${rejected}\n\nExample Question:\nCan you help configure my cluster to automatically scale?\nExample Response:\n${allowed}\n\nQuestion:\n${message}\nResponse:\n" + }, + "invalid_question_response": { + "type": "string", + "title": "Invalid question response", + "description": "The default response when the Users' question is determined to be invalid.", + "default": "\nHi, I'm the OpenShift Lightspeed assistant, I can help you with questions about OpenShift, \nplease ask me a question related to OpenShift.\n" + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "model_id" + ], + "title": "QuestionValidityConfig", + "description": "Configuration for the question validity guardrail." + }, + "QuestionValidityShieldConfiguration": { + "properties": { + "name": { + "type": "string", + "title": "Shield name", + "description": "Unique, user-facing name identifying this shield instance." + }, + "provider_id": { + "type": "string", + "const": "question_validity", + "title": "Shield provider id", + "description": "Discriminator identifying this as a question-validity shield." + }, + "config": { + "$ref": "#/components/schemas/QuestionValidityConfig", + "title": "Shield configuration", + "description": "Question-validity-specific configuration for this shield." + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "name", + "provider_id", + "config" + ], + "title": "QuestionValidityShieldConfiguration", + "description": "Configuration for a named question-validity guardrail shield.\n\nAttributes:\n name: Unique, user-facing name identifying this shield instance.\n provider_id: Discriminator identifying this as a question-validity shield.\n config: Question-validity-specific configuration." + }, "QuotaExceededResponse": { "properties": { "status_code": { @@ -19071,6 +19380,91 @@ } ] }, + "RedactionConfig": { + "properties": { + "rules": { + "items": { + "$ref": "#/components/schemas/RedactionRule" + }, + "type": "array", + "title": "Redaction rules", + "description": "Ordered list of PII redaction rules" + }, + "case_sensitive": { + "type": "boolean", + "title": "Case sensitive", + "description": "When False, patterns are compiled with re.IGNORECASE", + "default": false + } + }, + "additionalProperties": false, + "type": "object", + "title": "RedactionConfig", + "description": "Configuration for PII redaction with regex-based rules.\n\nRules are validated and compiled at construction time. Invalid\nregex patterns raise a ``ValueError`` immediately.\n\nAttributes:\n rules: Ordered list of redaction rules applied sequentially.\n case_sensitive: When False, patterns are compiled with\n ``re.IGNORECASE``. Defaults to False." + }, + "RedactionRule": { + "properties": { + "pattern": { + "type": "string", + "title": "Pattern", + "description": "Regex pattern to match sensitive data" + }, + "replacement": { + "type": "string", + "title": "Replacement", + "description": "Replacement string for matched text" + }, + "case_sensitive": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "title": "Case sensitive", + "description": "Per-rule case sensitivity override. When None, the global config flag applies." + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "pattern", + "replacement" + ], + "title": "RedactionRule", + "description": "A single regex-based redaction rule.\n\nAttributes:\n pattern: Raw regex pattern string to match sensitive data.\n replacement: Text to substitute for each match.\n case_sensitive: Per-rule override for case sensitivity.\n When None, the global ``RedactionConfig.case_sensitive``\n flag applies." + }, + "RedactionShieldConfiguration": { + "properties": { + "name": { + "type": "string", + "title": "Shield name", + "description": "Unique, user-facing name identifying this shield instance." + }, + "provider_id": { + "type": "string", + "const": "redaction", + "title": "Shield provider id", + "description": "Discriminator identifying this as a redaction shield." + }, + "config": { + "$ref": "#/components/schemas/RedactionConfig", + "title": "Shield configuration", + "description": "Redaction-specific configuration for this shield." + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "name", + "provider_id", + "config" + ], + "title": "RedactionShieldConfiguration", + "description": "Configuration for a named PII-redaction guardrail shield.\n\nAttributes:\n name: Unique, user-facing name identifying this shield instance.\n provider_id: Discriminator identifying this as a redaction shield.\n config: Redaction-specific configuration." + }, "ReferencedDocument": { "properties": { "doc_url": { @@ -19189,6 +19583,9 @@ }, { "$ref": "#/components/schemas/OpenAIResponseMCPApprovalResponse" + }, + { + "$ref": "#/components/schemas/OpenAIResponseOutputMessageReasoningItem" } ] }, @@ -19504,7 +19901,7 @@ "input" ], "title": "ResponsesRequest", - "description": "Model representing a request for the Responses API following LCORE specification.\n\nAttributes:\n input: Input text or structured input items containing the query.\n model: Model identifier in format \"provider/model\". Auto-selected if not provided.\n conversation: Conversation ID linking to an existing conversation. Accepts both\n OpenAI and LCORE formats. Mutually exclusive with previous_response_id.\n include: Explicitly specify output item types that are excluded by default but\n should be included in the response.\n instructions: System instructions or guidelines provided to the model (acts as\n the system prompt).\n max_infer_iters: Maximum number of inference iterations the model can perform.\n max_output_tokens: Maximum number of tokens allowed in the response.\n max_tool_calls: Maximum number of tool calls allowed in a single response.\n metadata: Custom metadata dictionary with key-value pairs for tracking or logging.\n parallel_tool_calls: Whether the model can make multiple tool calls in parallel.\n previous_response_id: Identifier of the previous response in a multi-turn\n conversation. Mutually exclusive with conversation.\n prompt: Prompt object containing a template with variables for dynamic\n substitution.\n reasoning: Reasoning configuration for the response.\n safety_identifier: Safety identifier for the response.\n store: Whether to store the response in conversation history. Defaults to True.\n stream: Whether to stream the response as it is generated. Defaults to False.\n temperature: Sampling temperature controlling randomness (typically 0.0\u20132.0).\n text: Text response configuration specifying output format constraints (JSON\n schema, JSON object, or plain text).\n tool_choice: Tool selection strategy (\"auto\", \"required\", \"none\", or specific\n tool configuration).\n tools: List of tools available to the model (file search, web search, function\n calls, MCP tools). Defaults to all tools available to the model.\n generate_topic_summary: LCORE-specific flag indicating whether to generate a\n topic summary for new conversations. Defaults to True.\n shield_ids: LCORE-specific list of safety shield IDs to apply. If None, all\n configured shields are used.\n solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict.", + "description": "Model representing a request for the Responses API following LCORE specification.\n\nAttributes:\n input: Input text or structured input items containing the query.\n model: Model identifier in format \"provider/model\". Auto-selected if not provided.\n conversation: Conversation ID linking to an existing conversation. Accepts both\n OpenAI and LCORE formats. Mutually exclusive with previous_response_id.\n include: Explicitly specify output item types that are excluded by default but\n should be included in the response.\n instructions: System instructions or guidelines provided to the model (acts as\n the system prompt).\n max_infer_iters: Maximum number of inference iterations the model can perform.\n max_output_tokens: Maximum number of tokens allowed in the response.\n max_tool_calls: Maximum number of tool calls allowed in a single response.\n metadata: Custom metadata dictionary with key-value pairs for tracking or logging.\n parallel_tool_calls: Whether the model can make multiple tool calls in parallel.\n previous_response_id: Identifier of the previous response in a multi-turn\n conversation. Mutually exclusive with conversation.\n prompt: Prompt object containing a template with variables for dynamic\n substitution.\n reasoning: Reasoning configuration for the response.\n safety_identifier: Safety identifier for the response.\n store: Whether to store the response in conversation history. Defaults to True.\n stream: Whether to stream the response as it is generated. Defaults to False.\n temperature: Sampling temperature controlling randomness (typically 0.0\u20132.0).\n text: Text response configuration specifying output format constraints (JSON\n schema, JSON object, or plain text).\n tool_choice: Tool selection strategy (\"auto\", \"required\", \"none\", or specific\n tool configuration).\n tools: List of tools available to the model (file search, web search, function\n calls, MCP tools). Defaults to all tools available to the model.\n generate_topic_summary: LCORE-specific flag indicating whether to generate a\n topic summary for new conversations. Defaults to True.\n shield_ids: LCORE-specific list of configured shield names to apply.\n If None, all configured shields are used.\n solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict.", "examples": [ { "generate_topic_summary": true, @@ -19580,6 +19977,9 @@ }, { "$ref": "#/components/schemas/OpenAIResponseMCPApprovalRequest" + }, + { + "$ref": "#/components/schemas/OpenAIResponseOutputMessageReasoningItem" } ], "discriminator": { @@ -19591,6 +19991,7 @@ "mcp_call": "#/components/schemas/OpenAIResponseOutputMessageMCPCall", "mcp_list_tools": "#/components/schemas/OpenAIResponseOutputMessageMCPListTools", "message": "#/components/schemas/OpenAIResponseMessage", + "reasoning": "#/components/schemas/OpenAIResponseOutputMessageReasoningItem", "web_search_call": "#/components/schemas/OpenAIResponseOutputMessageWebSearchToolCall" } } @@ -20670,7 +21071,7 @@ }, "type": "object", "title": "SearchRankingOptions", - "description": "Options for ranking and filtering search results.\n\nThis class configures how search results are ranked and filtered. You can use algorithm-based\nrerankers (weighted, RRF) or neural rerankers. Defaults from VectorStoresConfig are\nused when parameters are not provided.\n\nExamples:\n # Weighted ranker with custom alpha\n SearchRankingOptions(ranker=\"weighted\", alpha=0.7)\n\n # RRF ranker with custom impact factor\n SearchRankingOptions(ranker=\"rrf\", impact_factor=50.0)\n\n # Use config defaults (just specify ranker type)\n SearchRankingOptions(ranker=\"weighted\") # Uses alpha from VectorStoresConfig\n\n # Score threshold filtering\n SearchRankingOptions(ranker=\"weighted\", score_threshold=0.5)\n\n:param ranker: (Optional) Name of the ranking algorithm to use. Supported values:\n - \"weighted\": Weighted combination of vector and keyword scores\n - \"rrf\": Reciprocal Rank Fusion algorithm\n - \"neural\": Neural reranking model (requires model parameter, Part II)\n Note: For OpenAI API compatibility, any string value is accepted, but only the above values are supported.\n:param score_threshold: (Optional) Minimum relevance score threshold for results. Default: 0.0\n:param alpha: (Optional) Weight factor for weighted ranker (0-1).\n - 0.0 = keyword only\n - 0.5 = equal weight (default)\n - 1.0 = vector only\n Only used when ranker=\"weighted\" and weights is not provided.\n Falls back to VectorStoresConfig.chunk_retrieval_params.weighted_search_alpha if not provided.\n:param impact_factor: (Optional) Impact factor (k) for RRF algorithm.\n Lower values emphasize higher-ranked results. Default: 60.0 (optimal from research).\n Only used when ranker=\"rrf\".\n Falls back to VectorStoresConfig.chunk_retrieval_params.rrf_impact_factor if not provided.\n:param weights: (Optional) Dictionary of weights for combining different signal types.\n Keys can be \"vector\", \"keyword\", \"neural\". Values should sum to 1.0.\n Used when combining algorithm-based reranking with neural reranking (Part II).\n Example: {\"vector\": 0.3, \"keyword\": 0.3, \"neural\": 0.4}\n:param model: (Optional) Model identifier for neural reranker (e.g., \"vllm/Qwen3-Reranker-0.6B\").\n Required when ranker=\"neural\" or when weights contains \"neural\" (Part II)." + "description": "Options for ranking and filtering search results.\n\nThis class configures how search results are ranked and filtered. You can use algorithm-based\nrerankers (weighted, RRF) or neural rerankers. Defaults from VectorStoresConfig are\nused when parameters are not provided.\n\nExamples:\n # Weighted ranker with custom alpha\n SearchRankingOptions(ranker=\"weighted\", alpha=0.7)\n\n # RRF ranker with custom impact factor\n SearchRankingOptions(ranker=\"rrf\", impact_factor=50.0)\n\n # Use config defaults (just specify ranker type)\n SearchRankingOptions(ranker=\"weighted\") # Uses alpha from VectorStoresConfig\n\n # Score threshold filtering\n SearchRankingOptions(ranker=\"weighted\", score_threshold=0.5)\n\n:param ranker: (Optional) Name of the ranking algorithm to use. Supported values:\n - \"weighted\": Weighted combination of vector and keyword scores\n - \"rrf\": Reciprocal Rank Fusion algorithm\n - \"neural\": Neural reranking model (requires model parameter)\n Note: For OpenAI API compatibility, any string value is accepted, but only the above values are supported.\n:param score_threshold: (Optional) Minimum relevance score threshold for results. Default: 0.0\n:param alpha: (Optional) Weight factor for weighted ranker (0-1).\n - 0.0 = keyword only\n - 0.5 = equal weight (default)\n - 1.0 = vector only\n Only used when ranker=\"weighted\" and weights is not provided.\n Falls back to VectorStoresConfig.chunk_retrieval_params.weighted_search_alpha if not provided.\n:param impact_factor: (Optional) Impact factor (k) for RRF algorithm.\n Lower values emphasize higher-ranked results. Default: 60.0 (optimal from research).\n Only used when ranker=\"rrf\".\n Falls back to VectorStoresConfig.chunk_retrieval_params.rrf_impact_factor if not provided.\n:param weights: (Optional) Dictionary of weights for combining different signal types.\n Keys can be \"vector\", \"keyword\", \"neural\". Values should sum to 1.0.\n Used when combining algorithm-based reranking with neural reranking.\n Example: {\"vector\": 0.3, \"keyword\": 0.3, \"neural\": 0.4}\n:param model: (Optional) Model identifier for neural reranker (e.g., \"transformers/Qwen/Qwen3-Reranker-0.6B\").\n Required when ranker=\"neural\" or when weights contains \"neural\"." }, "SecurityScheme": { "anyOf": [ @@ -20789,9 +21190,9 @@ { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" }, - "label": "llama stack" + "label": "ogx" }, { "detail": { @@ -20806,12 +21207,11 @@ "properties": { "shields": { "items": { - "additionalProperties": true, - "type": "object" + "$ref": "#/components/schemas/CatalogShield" }, "type": "array", "title": "Shields", - "description": "List of shields available" + "description": "List of shields configured in Lightspeed Core Stack" } }, "type": "object", @@ -20824,10 +21224,27 @@ { "shields": [ { - "identifier": "lightspeed_question_validity-shield", - "params": {}, - "provider_id": "lightspeed_question_validity", - "provider_resource_id": "lightspeed_question_validity-shield", + "config": { + "invalid_question_response": "I can only answer questions about the product.", + "model_id": "openai/gpt-4o-mini", + "model_prompt": "Is this question valid?" + }, + "name": "question-validity", + "provider_id": "question_validity", + "type": "shield" + }, + { + "config": { + "case_sensitive": false, + "rules": [ + { + "pattern": "\\b\\d{3}-\\d{2}-\\d{4}\\b", + "replacement": "[REDACTED]" + } + ] + }, + "name": "pii-redaction", + "provider_id": "redaction", "type": "shield" } ] @@ -21227,8 +21644,7 @@ "properties": { "tools": { "items": { - "additionalProperties": true, - "type": "object" + "$ref": "#/components/schemas/CatalogTool" }, "type": "array", "title": "Tools", diff --git a/docs/devel_doc/providers.md b/docs/devel_doc/providers.md index b4601d2e1..d133731ab 100644 --- a/docs/devel_doc/providers.md +++ b/docs/devel_doc/providers.md @@ -198,14 +198,12 @@ make run CONFIG=examples/lightspeed-stack-azure-entraid-service.yaml ## Safety Providers -| Name | Type | Pip Dependencies | Supported in LCS | -|--------------|--------|--------------------------------------------------------------------------------------|:----------------:| -| code-scanner | inline | `codeshield` | ❌ | -| llama-guard | inline | — | ❌ | -| prompt-guard | inline | `transformers[accelerate]`, `torch --index-url https://download.pytorch.org/whl/cpu` | ❌ | -| bedrock | remote | `boto3` | ❌ | -| nvidia | remote | `requests` | ❌ | -| sambanova | remote | `litellm`, `requests` | ❌ | +Shields are owned by LCORE (configured under `shields:` block), not as OGX `providers.safety` entries. + +| Name | Type | Pip Dependencies | Supported in LCS | +|--------------------|-------|------------------|:----------------:| +| question_validity | lcore | — | ✅ | +| redaction | lcore | — | ✅ | --- @@ -342,7 +340,6 @@ make run CONFIG=examples/lightspeed-stack-azure-entraid-service.yaml Some of APIs are associated with a set of **Resources**. Here is the mapping of APIs to resources: - **Inference**, **Eval** and **Post Training** are associated with **Model** resources. - - **Safety** is associated with **Shield** resources. - **Tool Runtime** is associated with **ToolGroup** resources. - **DatasetIO** is associated with **Dataset** resources. - **VectorIO** is associated with **VectorDB** resources. @@ -359,9 +356,6 @@ make run CONFIG=examples/lightspeed-stack-azure-entraid-service.yaml provider_id: openai model_type: llm provider_model_id: gpt-4-turbo # provider label - - shields: - ... ``` **Note** It is necessary for llama-stack to know which resources to use for a given provider. This means you need to explicitly register resources (including models) before you can use them with the associated APIs. diff --git a/docs/devel_doc/responses.md b/docs/devel_doc/responses.md index 707658b0d..bde420988 100644 --- a/docs/devel_doc/responses.md +++ b/docs/devel_doc/responses.md @@ -107,7 +107,7 @@ The following fields are LCORE-specific request extensions and are not part of t | Field | Type | Description | Required | |-------|------|-------------|----------| | `generate_topic_summary` | boolean | Generate topic summary for new conversations. Default: true | No | -| `shield_ids` | array[string] | Shield IDs to apply. If omitted, all configured shields in LCORE are used | No | +| `shield_ids` | array[string] | LCORE-configured shield `name` values to apply. If omitted, all configured shields are used. Not Llama Stack Safety resource names. | No | | `solr` | object | Optional `mode` and `filters`. Legacy top-level filter-only objects are still accepted. | No | @@ -154,7 +154,7 @@ All input item objects have a common `type` discriminator that determines the su Optional. List of output item types to include in the response that are excluded by default. -Allowed values (literal strings): `web_search_call.action.sources`, `code_interpreter_call.outputs`, `computer_call_output.output.image_url`, `file_search_call.results`, `message.input_image.image_url`, `message.output_text.logprobs`, `reasoning.encrypted_content`. +Allowed values (literal strings): `web_search_call.action.sources`, `code_interpreter_call.outputs`, `computer_call_output.output.image_url`, `file_search_call.results`, `message.input_image.image_url`, `message.output_text.logprobs`. **Examples:** @@ -600,7 +600,7 @@ In streaming mode, server-deployed MCP events (e.g. `mcp_call`, `mcp_list_tools` The API introduces extensions that are not part of the OpenResponses specification: - `generate_topic_summary` (request) — When set to `true` and a new conversation is created, a topic summary is automatically generated and stored in conversation metadata. -- `shield_ids` (request) — Optional list of safety shield IDs to apply. If omitted, all configured shields are used. +- `shield_ids` (request) — Optional list of LCORE shield `name` values to apply. If omitted, all shields from `lightspeed-stack.yaml` are used. See the [Safety Shields Guide](../user_doc/shields_guide.md). - `solr` (request) — Object with optional `mode` (`semantic`, `hybrid`, or `lexical`) and `filters` (Solr vector_io provider payload). Legacy filter-only objects (no `mode`/`filters` wrapper) still work. - `available_quotas` (response) — Provides real-time quota information from all configured quota limiters. diff --git a/docs/models/requests.md b/docs/models/requests.md index 022368741..c77d72f11 100644 --- a/docs/models/requests.md +++ b/docs/models/requests.md @@ -844,7 +844,7 @@ Attributes: generate_topic_summary: Whether to generate topic summary for new conversations. media_type: The optional media type for response format (application/json or text/plain). vector_store_ids: The optional list of specific vector store IDs to query for RAG. - shield_ids: The optional list of safety shield IDs to apply. + shield_ids: The optional list of configured shield names to apply. solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict. @@ -860,7 +860,7 @@ Attributes: | generate_topic_summary | boolean | Whether to generate topic summary for new conversations | | media_type | string | Media type for the response format | | vector_store_ids | array | Optional list of specific vector store IDs to query for RAG. If not provided, all available vector stores will be queried. | -| shield_ids | array | Optional list of safety shield IDs to apply. If None, all configured shields are used. | +| shield_ids | array | Optional list of configured shield names to apply. If None, all configured shields are used. | | solr | | Solr inline RAG config: mode (semantic, hybrid, lexical) and filters; a legacy filter-only object (e.g. fq) is still accepted. | @@ -912,8 +912,8 @@ Attributes: calls, MCP tools). Defaults to all tools available to the model. generate_topic_summary: LCORE-specific flag indicating whether to generate a topic summary for new conversations. Defaults to True. - shield_ids: LCORE-specific list of safety shield IDs to apply. If None, all - configured shields are used. + shield_ids: LCORE-specific list of configured shield names to apply. + If None, all configured shields are used. solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict. diff --git a/docs/testing/e2e_scenarios.md b/docs/testing/e2e_scenarios.md index 4b3b4e8e0..d63c6ccdf 100644 --- a/docs/testing/e2e_scenarios.md +++ b/docs/testing/e2e_scenarios.md @@ -102,8 +102,7 @@ * Check if the OpenAPI endpoint works as expected * Check if info endpoint is working * Check if info endpoint reports error when llama-stack connection is not working -* Check if shields endpoint is working -* Check if shields endpoint reports error when llama-stack is unreachable +* Check if shields endpoint is working (lists LCORE-configured shields) * Check if tools endpoint is working * Check if tools endpoint reports error when llama-stack is unreachable * Check if metrics endpoint is working diff --git a/docs/testing/e2e_testing.md b/docs/testing/e2e_testing.md index 2cba4eea6..6bddfd3bb 100644 --- a/docs/testing/e2e_testing.md +++ b/docs/testing/e2e_testing.md @@ -24,7 +24,7 @@ This guide describes how to run, extend, and understand the Lightspeed Core Stac - **Framework**: [Behave](https://behave.readthedocs.io/) (Python BDD). - **Scope**: REST API of the Lightspeed Core Stack (query, streaming_query, models, info, health, feedback, conversations, RBAC, MCP, etc.). -- **Execution**: Tests run in a **separate process** from the app. They send HTTP requests to the service and (in server mode) optionally talk to the Llama Stack service for shield setup. +- **Execution**: Tests run in a **separate process** from the app. They send HTTP requests to the service. LCORE shields are configured in `lightspeed-stack.yaml` (not via Llama Stack Safety APIs). - **Environments**: Local (Docker Compose) or Prow/OpenShift (containers/pods). Mode is detected via `E2E_DEPLOYMENT_MODE` and `RUNNING_PROW`. --- @@ -206,8 +206,8 @@ You can put several tags on one scenario. To document why a scenario is skipped, - **before_all**: Sets `deployment_mode`, `is_library_mode`, detects or overrides `default_model` / `default_provider`, sets `faiss_vector_store_id`. - **before_feature**: Applies feature-level config and restarts container for `Authorized`, `RBAC`, `RHIdentity`, `Feedback`, `MCP`. -- **before_scenario**: Skips scenarios for `@skip`, `@local`, `@skip-in-library-mode`; applies scenario config for `InvalidFeedbackStorageConfig` / `NoCacheConfig`; for `@disable-shields` (server mode) unregisters the shield. -- **after_scenario**: Restores Llama Stack if it was disrupted; restores config and restarts for scenario config tags; for `@disable-shields` re-registers the shield. +- **before_scenario**: Skips scenarios for `@skip`, `@local`, `@skip-in-library-mode`; applies scenario config for `InvalidFeedbackStorageConfig` / `NoCacheConfig`. +- **after_scenario**: Restores Llama Stack if it was disrupted; restores config and restarts for scenario config tags. - **after_feature**: Restores config and restarts for `Authorized`, `RBAC`, `RHIdentity`, `MCP`; deletes feedback conversations for `Feedback`. --- @@ -353,7 +353,7 @@ Here, **Given** sets state, **When** performs the HTTP call, **Then** and **And* - **"Container state improper" / restart fails**: Usually the llama-stack container is in a bad state. Ensure it is started (or recreated) before restarting lightspeed-stack; see Docker/Podman and compose usage in the project. - **Readonly database (SQLite) in Llama Stack**: If the RAG KV DB is on a bind-mounted path that becomes read-only (e.g. after restart), move it to a named volume (e.g. via `KV_RAG_PATH` in docker-compose) so writes succeed. - **ChunkedEncodingError on streaming_query**: The step for streaming_query uses `stream=True` and consumes the stream; if you add new streaming steps, avoid reading the full response with `response.content` and use the same stream-reading pattern so a server close after an error event does not raise. -- **Event loop is closed (httpx/AsyncClient)**: In E2E, any code that creates an `AsyncLlamaStackClient` (e.g. for shields) must close it (e.g. `await client.close()`) in a `finally` block before the event loop is torn down (e.g. before `asyncio.run()` returns). +- **Event loop is closed (httpx/AsyncClient)**: In E2E, any code that creates an `AsyncOgxClient` (e.g. for shields) must close it (e.g. `await client.close()`) in a `finally` block before the event loop is torn down (e.g. before `asyncio.run()` returns). - **Scenarios skipped**: Check tags (`@skip`, `@skip-in-library-mode`, `@local`) and `E2E_DEPLOYMENT_MODE`; ensure the scenario is not excluded by `--tags=-skip` (or the opposite if you intend to run only skipped scenarios for debugging). For more on test structure and commands, see the main project guide (`CLAUDE.md`) and `tests/e2e/features/steps/README.md`. diff --git a/docs/user_doc/a2a_protocol.md b/docs/user_doc/a2a_protocol.md index af6648d8f..a70f3b3e4 100644 --- a/docs/user_doc/a2a_protocol.md +++ b/docs/user_doc/a2a_protocol.md @@ -43,7 +43,7 @@ The A2A protocol is an open standard for agent-to-agent communication that allow │ ┌──────────────────────────────────────────────────────────┐ │ │ │ Llama Stack Client │ │ │ │ - Responses API (streaming responses) │ │ -│ │ - Tools, Shields, RAG integration │ │ +│ │ - Tools, RAG integration │ │ │ └──────────────────────────────────────────────────────────┘ │ └─────────────────────────────────────────────────────────────────┘ ``` diff --git a/docs/user_doc/config.md b/docs/user_doc/config.md index 66f8c1f16..bf8d52d2b 100644 --- a/docs/user_doc/config.md +++ b/docs/user_doc/config.md @@ -258,6 +258,7 @@ Global service configuration. | okp | | OKP provider settings. Only used when 'okp' is listed in rag.inline or rag.tool. | | reranker | | Configuration for neural reranking of RAG chunks using cross-encoder. | | skills | | Agent skills configuration. Specifies paths to skill directories. | +| shields | array | Configuration for a single named guardrail shield (question validity or redaction). | ## ConversationHistoryConfiguration @@ -757,6 +758,70 @@ Paths are validated at startup to ensure they exist and contain valid SKILL.md f | paths | array | Paths to skill directories or directories containing skill subdirectories. | +## QuestionValidityConfig + + +Configuration for the question validity guardrail. + + +| Field | Type | Description | +|---------------------------|--------|---------------------------------------------------------------| +| model_id | string | The model_id to use for the guard | +| model_prompt | string | Prompt sent to the LLM used to validate the user's question | +| invalid_question_response | string | Response when the user's question is determined to be invalid | + + +## QuestionValidityShieldConfiguration + + +Configuration for a named question-validity guardrail shield. + + +| Field | Type | Description | +|-------------|--------|--------------------------------------------------------------| +| name | string | Unique, user-facing name identifying this shield instance | +| provider_id | string | Discriminator identifying this as a question-validity shield | +| config | | Question-validity-specific configuration for this shield | + + +## RedactionRule + + +A single regex-based redaction rule. + + +| Field | Type | Description | +|----------------|---------|-----------------------------------------------------------------------| +| pattern | string | Regex pattern to match sensitive data | +| replacement | string | Replacement string for matched text | +| case_sensitive | boolean | Per-rule override; when null, the global RedactionConfig flag applies | + + +## RedactionConfig + + +Configuration for PII redaction with regex-based rules. + + +| Field | Type | Description | +|----------------|---------|--------------------------------------------------------| +| rules | array | Ordered list of PII redaction rules | +| case_sensitive | boolean | When false, patterns are compiled with `re.IGNORECASE` | + + +## RedactionShieldConfiguration + + +Configuration for a named PII-redaction guardrail shield. + + +| Field | Type | Description | +|-------------|--------|-----------------------------------------------------------| +| name | string | Unique, user-facing name identifying this shield instance | +| provider_id | string | Discriminator identifying this as a redaction shield | +| config | | Redaction-specific configuration for this shield | + + ## SplunkConfiguration diff --git a/docs/user_doc/shields_guide.md b/docs/user_doc/shields_guide.md new file mode 100644 index 000000000..5e1bdeeca --- /dev/null +++ b/docs/user_doc/shields_guide.md @@ -0,0 +1,203 @@ +# Safety Shields Guide + +This guide covers LCORE-owned safety shields: how to configure them in +`lightspeed-stack.yaml`, which shield types are supported, how they apply on +request endpoints, how to list them via `/v1/shields`, and how `shield_ids` +request overrides work. + +> [!IMPORTANT] +> Shields used by `/query`, `/streaming_query`, `/responses`, and `/rlsapi` are +> **owned and configured by Lightspeed Core Stack**, not by the Llama Stack / +> OGX Safety or Moderations APIs anymore. Do not configure LCORE request guardrails +> under `providers.safety` / `registered_resources.shields` in the stack +> `run.yaml`. + +--- + +- [Introduction](#introduction) +- [Configuration](#configuration) + - [Supported shield types](#supported-shield-types) + - [question_validity](#question_validity) + - [redaction](#redaction) +- [How shields apply at runtime](#how-shields-apply-at-runtime) + - [Agent-based endpoints](#agent-based-endpoints) + - [Responses-based endpoints](#responses-based-endpoints) + - [Per-endpoint behavior](#per-endpoint-behavior) +- [Listing shields (`GET /v1/shields`)](#listing-shields-get-v1shields) +- [Request overrides (`shield_ids`)](#request-overrides-shield_ids) +- [Disabling overrides](#disabling-overrides) +- [References](#references) + +--- + +# Introduction + +LCORE shields are guardrails declared in the Lightspeed Core Stack +configuration. Each entry has: + +| Field | Meaning | +|-------|---------| +| `name` | Unique shield name used in `/v1/shields` and in `shield_ids` overrides | +| `provider_id` | Shield type discriminator (`question_validity` or `redaction`) | +| `config` | Type-specific settings | + +Names must be unique across the `shields` list. + +# Configuration + +Add a `shields` list to `lightspeed-stack.yaml`: + +```yaml +shields: + - name: topic-guard + provider_id: question_validity + config: + model_id: openai/gpt-4o-mini + # optional: + # model_prompt: "..." + # invalid_question_response: "..." + + - name: pii-redaction + provider_id: redaction + config: + rules: + - pattern: '\b\d{3}-\d{2}-\d{4}\b' + replacement: '[REDACTED]' + case_sensitive: false +``` + +See [examples/lightspeed-stack-shields.yaml](../../examples/lightspeed-stack-shields.yaml) +for a complete example. + +## Supported shield types + +| `provider_id` | Purpose | Typical application | +|---------------|---------|---------------------| +| `question_validity` | Classify whether the user question is in-topic; reject off-topic input with a fixed reply | Agent capability on agent-based endpoints; also considered by direct-run input moderation | +| `redaction` | Regex-based PII / sensitive-data redaction of model messages | Agent capability on agent-based endpoints | + +## question_validity + +| Config field | Required | Description | +|--------------|----------|-------------| +| `model_id` | Yes | Model used for the validity check (for example `openai/gpt-4o-mini`) | +| `model_prompt` | No | Classifier prompt (has a built-in default) | +| `invalid_question_response` | No | Reply returned when the question is rejected | + +## redaction + +| Config field | Required | Description | +|--------------|----------|-------------| +| `rules` | No (default `[]`) | Ordered list of `{pattern, replacement, case_sensitive?}` rules | +| `case_sensitive` | No (default `false`) | Global case sensitivity when a rule does not override it | + +Invalid regex patterns are rejected at configuration load time. + +# How shields apply at runtime + +The same shield logic (`question_validity` and `redaction`) is used on both +agent-based and responses-based endpoints; only the integration point differs. + +## Agent-based endpoints + +On agent-based endpoints (for example `/v1/query` and `/v1/streaming_query`), +shields run as **pydantic-ai capabilities** attached when the agent is built. +Those capabilities wrap the agent pipeline — for example rejecting off-topic +questions or redacting PII from model messages — using the configured shields. + +## Responses-based endpoints + +On pure responses-based endpoints (for example `/v1/responses` and `/v1/infer`), +there is no agent capability layer. Instead, LCORE runs the **same core shield +functionality directly** through a custom API (`run_shield_moderation`) before +each request. When moderation blocks the input, the endpoint returns a refusal +(and may persist the blocked turn) without calling the model. + +## Per-endpoint behavior + +| Endpoint | How shields run | `shield_ids` | +|----------|-----------------|--------------| +| `POST /v1/query` | Agent capabilities (via `build_agent`) | Yes; subject to `disable_shield_ids_override` | +| `POST /v1/streaming_query` | Agent capabilities (via `build_agent`) | Yes; subject to `disable_shield_ids_override` | +| `POST /v1/responses` | Direct custom API before the request; agent capabilities when the request uses the agent path | Yes (`shield_ids` is an LCORE extension). Override disable gate is not applied on this endpoint today | +| `POST /v1/infer` (rlsapi v1) | Direct custom API before the request | No `shield_ids` field — always uses all configured shields | + +# Listing shields (`GET /v1/shields`) + +`GET /v1/shields` returns shields from **LCORE configuration only**. It does +not call Llama Stack / OGX to list Safety or Moderations resources. + +Each catalog entry has this shape: + +| Field | Description | +|-------|-------------| +| `name` | Configured shield name | +| `provider_id` | `question_validity` or `redaction` | +| `type` | Always `"shield"` | +| `config` | Type-specific shield configuration | + +Example response body: + +```json +{ + "shields": [ + { + "name": "pii-redaction", + "provider_id": "redaction", + "type": "shield", + "config": { + "rules": [ + { + "pattern": "\\d+", + "replacement": "[NUM]", + "case_sensitive": null + } + ], + "case_sensitive": false + } + } + ] +} +``` + +# Request overrides (`shield_ids`) + +Optional request field on `/v1/query`, `/v1/streaming_query`, and +`/v1/responses`: + +| `shield_ids` value | Behavior | +|--------------------|----------| +| omitted / `null` | Apply **all** configured shields | +| `[]` | Apply **no** shields | +| `["topic-guard", ...]` | Apply only those names; unknown IDs yield HTTP **404** | + +Values must match configured `name` strings (as returned by +`GET /v1/shields`), not Llama Stack shield resource names. + +Example: + +```json +{ + "query": "How do I scale a Deployment?", + "shield_ids": ["topic-guard"] +} +``` + +# Disabling overrides + +To ignore client-provided `shield_ids` on `/v1/query` and +`/v1/streaming_query` (always use the configured set), set: + +```yaml +customization: + disable_shield_ids_override: true +``` + +When this flag is set and the client still sends `shield_ids` (including an +empty list), the endpoint returns HTTP **422**. + +# References + +- [Configuration options](config.md) — schema tables for shield-related models +- [OpenResponses /responses](../devel_doc/responses.md) — `shield_ids` LCORE extension +- [Example configuration](../../examples/lightspeed-stack-shields.yaml) diff --git a/examples/azure-run.yaml b/examples/azure-run.yaml index 91cc92bdf..17195a8d0 100644 --- a/examples/azure-run.yaml +++ b/examples/azure-run.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: azure-configuration +distro_name: azure-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -40,19 +35,10 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - config: {} - provider_id: basic - provider_type: inline::basic tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol @@ -63,45 +49,21 @@ providers: backend: kv_rag provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -128,8 +90,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: gpt-4o-mini @@ -142,25 +107,13 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: - embedding_dimension: 768 embedding_model: sentence-transformers/all-mpnet-base-v2 provider_id: faiss vector_store_id: ${env.FAISS_VECTOR_STORE_ID} - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/examples/bedrock-run.yaml b/examples/bedrock-run.yaml index b793da3e9..04db8ece4 100644 --- a/examples/bedrock-run.yaml +++ b/examples/bedrock-run.yaml @@ -1,20 +1,15 @@ version: 2 apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io - -benchmarks: [] -datasets: [] -image_name: starter + +distro_name: starter # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -39,26 +34,10 @@ providers: storage_dir: ${env.SQLITE_STORE_DIR:=~/.llama/storage/files} provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol @@ -69,45 +48,21 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -131,8 +86,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: custom-bedrock-model @@ -145,17 +103,7 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: diff --git a/examples/lightspeed-stack-shields.yaml b/examples/lightspeed-stack-shields.yaml new file mode 100644 index 000000000..2400f6ab7 --- /dev/null +++ b/examples/lightspeed-stack-shields.yaml @@ -0,0 +1,38 @@ +name: Lightspeed Core Service (LCS) with Shields +service: + host: localhost + port: 8080 + auth_enabled: false + workers: 1 + color_log: true + access_log: true +llama_stack: + use_as_library_client: true + library_client_config_path: run.yaml +user_data_collection: + feedback_enabled: true + feedback_storage: "/tmp/data/feedback" + transcripts_enabled: true + transcripts_storage: "/tmp/data/transcripts" +authentication: + module: "noop" +# LCORE-owned safety shields (not Llama Stack / OGX Safety API resources). +# Listed via GET /v1/shields; selected per request with optional shield_ids. +shields: + - identifier: topic-guard + provider_id: question_validity + config: + model_id: openai/gpt-4o-mini + # Optional; omit to use built-in defaults: + model_prompt: "Classify whether the question is about OpenShift. Reply ALLOWED or REJECTED." + invalid_question_response: "I can only answer questions about OpenShift." + - identifier: pii-redaction + provider_id: redaction + config: + rules: + - pattern: '\b\d{3}-\d{2}-\d{4}\b' + replacement: "[REDACTED]" + case_sensitive: false +# Optional: reject client shield_ids overrides on /query and /streaming_query +# customization: +# disable_shield_ids_override: true diff --git a/examples/profiles/inline-faiss.yaml b/examples/profiles/inline-faiss.yaml index 4cf10bf2e..070d841b0 100644 --- a/examples/profiles/inline-faiss.yaml +++ b/examples/profiles/inline-faiss.yaml @@ -28,13 +28,13 @@ version: 2 apis: -- agents +- responses - files - inference - tool_runtime - vector_io -image_name: starter +distro_name: starter external_providers_dir: ${env.EXTERNAL_PROVIDERS_DIR:=~/.llama/providers.d} providers: @@ -51,8 +51,8 @@ providers: provider_type: inline::localfs tool_runtime: - config: {} - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search vector_io: - config: persistence: @@ -60,17 +60,14 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default provider_id: meta-reference - provider_type: inline::meta-reference + provider_type: inline::builtin storage: backends: @@ -93,19 +90,18 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: [] - shields: [] vector_stores: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime # REQUIRED for file_search tool calls to work. Without it, llama-stack's -# rag-runtime silently fails all file_search operations with no error logged. +# file-search runtime silently fails all file_search operations with no error logged. vector_stores: annotation_prompt_params: enable_annotations: false diff --git a/examples/profiles/openai-remote.yaml b/examples/profiles/openai-remote.yaml index f4315139b..0058a092d 100644 --- a/examples/profiles/openai-remote.yaml +++ b/examples/profiles/openai-remote.yaml @@ -19,14 +19,13 @@ version: 2 apis: -- agents +- responses - files - inference -- safety - tool_runtime - vector_io -image_name: starter +distro_name: starter external_providers_dir: ${env.EXTERNAL_PROVIDERS_DIR:=~/.llama/providers.d} providers: @@ -46,15 +45,10 @@ providers: storage_dir: ${env.SQLITE_STORE_DIR:=~/.llama/storage/files} provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard tool_runtime: - config: {} - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search vector_io: - config: persistence: @@ -62,17 +56,14 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default provider_id: meta-reference - provider_type: inline::meta-reference + provider_type: inline::builtin storage: backends: @@ -95,24 +86,18 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: [] - # No shield is registered: a real safety gate needs a real guard model - # (e.g. meta-llama/Llama-Guard-3-8B) registered per deployment, with - # provider_shield_id pointing at that guard model and a matching - # safety.default_shield_id. A placeholder shield backed by a chat model - # would only *look* like moderation. - shields: [] vector_stores: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime # REQUIRED for file_search tool calls to work. Without it, llama-stack's -# rag-runtime silently fails all file_search operations with no error logged. +# file-search runtime silently fails all file_search operations with no error logged. vector_stores: annotation_prompt_params: enable_annotations: false diff --git a/examples/run.yaml b/examples/run.yaml index 5b330b2fc..63cc35941 100644 --- a/examples/run.yaml +++ b/examples/run.yaml @@ -7,20 +7,15 @@ version: 2 apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] -image_name: starter +distro_name: starter # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -41,68 +36,28 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable the MCP tool provider_id: model-context-protocol provider_type: remote::model-context-protocol - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -126,24 +81,17 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: gpt-4o-mini provider_id: openai model_type: llm provider_model_id: gpt-4o-mini - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: annotation_prompt_params: # Override the default Llama Stack annotation that adds <| file-xyz |> to responses enable_annotations: true @@ -155,7 +103,5 @@ vector_stores: default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: nomic-ai/nomic-embed-text-v1.5 -safety: - default_shield_id: llama-guard telemetry: enabled: true diff --git a/examples/vertexai-run.yaml b/examples/vertexai-run.yaml index 6a49e350f..69f0a8a28 100644 --- a/examples/vertexai-run.yaml +++ b/examples/vertexai-run.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: vertexai-configuration +distro_name: vertexai-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -40,19 +35,10 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - config: {} - provider_id: basic - provider_type: inline::basic tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol @@ -63,45 +49,21 @@ providers: backend: kv_rag provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -128,8 +90,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: all-mpnet-base-v2 @@ -138,25 +103,13 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: - embedding_dimension: 768 embedding_model: sentence-transformers/all-mpnet-base-v2 provider_id: faiss vector_store_id: ${env.FAISS_VECTOR_STORE_ID} - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/examples/vllm-rhaiis.yaml b/examples/vllm-rhaiis.yaml index f28c7fd82..c23b9ff73 100644 --- a/examples/vllm-rhaiis.yaml +++ b/examples/vllm-rhaiis.yaml @@ -1,20 +1,14 @@ version: 2 -image_name: rhaiis-configuration +distro_name: rhaiis-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -41,19 +35,10 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - config: {} - provider_id: basic - provider_type: inline::basic tool_runtime: - config: {} - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search vector_io: - config: persistence: @@ -61,45 +46,21 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default provider_id: meta-reference - provider_type: inline::meta-reference + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -123,30 +84,21 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: ${env.RHAIIS_MODEL} provider_id: vllm model_type: llm provider_model_id: ${env.RHAIIS_MODEL} - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: provider_id: sentence-transformers model_id: nomic-ai/nomic-embed-text-v1.5 -safety: - default_shield_id: llama-guard telemetry: enabled: true diff --git a/examples/vllm-rhelai.yaml b/examples/vllm-rhelai.yaml index 4fa263387..c2b33ead1 100644 --- a/examples/vllm-rhelai.yaml +++ b/examples/vllm-rhelai.yaml @@ -1,20 +1,14 @@ version: 2 -image_name: rhelai-configuration +distro_name: rhelai-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -42,69 +36,29 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol vector_io: [] - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default provider_id: meta-reference - provider_type: inline::meta-reference + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -128,8 +82,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: ${env.VLLM_MODEL} @@ -142,21 +99,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/examples/vllm-rhoai.yaml b/examples/vllm-rhoai.yaml index 38ca8d14a..81d91aa6d 100644 --- a/examples/vllm-rhoai.yaml +++ b/examples/vllm-rhoai.yaml @@ -1,20 +1,14 @@ version: 2 -image_name: rhoai-configuration +distro_name: rhoai-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -41,19 +35,10 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - provider_id: llama-guard - provider_type: inline::llama-guard - config: - excluded_categories: [] - scoring: - - config: {} - provider_id: basic - provider_type: inline::basic tool_runtime: - config: {} - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search vector_io: - config: persistence: @@ -61,45 +46,21 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default provider_id: meta-reference - provider_type: inline::meta-reference + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -123,28 +84,19 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: meta-llama/Llama-3.2-1B-Instruct provider_id: vllm model_type: llm provider_model_id: null - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: provider_id: sentence-transformers model_id: nomic-ai/nomic-embed-text-v1.5 -safety: - default_shield_id: llama-guard diff --git a/examples/watsonx-run.yaml b/examples/watsonx-run.yaml index ec7c988c4..da78b1561 100644 --- a/examples/watsonx-run.yaml +++ b/examples/watsonx-run.yaml @@ -1,20 +1,15 @@ version: 2 apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io - -benchmarks: [] -datasets: [] -image_name: starter + +distro_name: starter # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -41,26 +36,10 @@ providers: storage_dir: ${env.SQLITE_STORE_DIR:=~/.llama/storage/files} provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol @@ -71,45 +50,21 @@ providers: backend: kv_rag provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -136,8 +91,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: custom-watsonx-model @@ -150,21 +108,11 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: - embedding_dimension: 768 embedding_model: sentence-transformers/all-mpnet-base-v2 provider_id: faiss vector_store_id: ${env.FAISS_VECTOR_STORE_ID} - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG diff --git a/pyproject.toml b/pyproject.toml index 40979e208..8c91f1e00 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,9 +28,9 @@ dependencies = [ # Used by authentication/k8s integration "kubernetes>=30.1.0", # Used to call Llama Stack APIs - "llama-stack==0.6.0", - "llama-stack-client==0.6.0", - "llama-stack-api==0.6.0", + "ogx==1.0.2", + "ogx-client==1.0.2", + "ogx-api==1.0.2", # Used by Logger "rich>=14.0.0", # Used by JWK token auth handler @@ -154,7 +154,7 @@ dev = [ llslibdev = [ # To check llama-stack API provider dependecies: # - # $ uv run llama stack list-providers + # $ uv run ogx stack list-providers # # API agents: inline::meta-reference # commented out because they are not used in the project diff --git a/run.yaml b/run.yaml index e9bd5ef75..e4d5aad47 100644 --- a/run.yaml +++ b/run.yaml @@ -1,21 +1,15 @@ version: 2 apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io - -benchmarks: [] -datasets: [] -image_name: starter -external_providers_dir: /opt/app-root/providers/resources/external_providers + +distro_name: starter providers: inference: @@ -34,26 +28,13 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search + - config: {} # Enable MCP (Model Context Protocol) support + provider_id: model-context-protocol + provider_type: remote::model-context-protocol vector_io: - config: persistence: @@ -61,45 +42,21 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -123,23 +80,16 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: [] - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime # REQUIRED: This section is necessary for file_search tool calls to work. -# Without it, llama-stack's rag-runtime silently fails all file_search operations +# Without it, llama-stack's file-search runtime silently fails all file_search operations # with no error logged. vector_stores: # LCORE-1498: Disables Llama Stack RAG annotation generation @@ -150,6 +100,3 @@ vector_stores: default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: nomic-ai/nomic-embed-text-v1.5 -safety: - default_shield_id: llama-guard - diff --git a/scripts/generate_openapi_schema.py b/scripts/generate_openapi_schema.py index 509f14aeb..d9c10e9af 100644 --- a/scripts/generate_openapi_schema.py +++ b/scripts/generate_openapi_schema.py @@ -7,7 +7,7 @@ from fastapi.openapi.utils import get_openapi -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder # it is needed to read proper configuration in order to start the app to generate schema from configuration import configuration @@ -18,7 +18,7 @@ # Llama Stack client needs to be loaded before REST API is fully initialized import asyncio # noqa: E402 pylint: disable=C0411,C0413 -asyncio.run(AsyncLlamaStackClientHolder().load(configuration.configuration.llama_stack)) +asyncio.run(AsyncOgxClientHolder().load(configuration.configuration.llama_stack)) from app.main import app # noqa: E402 pylint: disable=C0413 diff --git a/scripts/llama-stack-entrypoint.sh b/scripts/llama-stack-entrypoint.sh index e3360c3b6..2ddcfd2e8 100755 --- a/scripts/llama-stack-entrypoint.sh +++ b/scripts/llama-stack-entrypoint.sh @@ -19,9 +19,9 @@ if [ -f "$LIGHTSPEED_CONFIG" ]; then if [ -f "$ENRICHED_CONFIG" ] && [ "$ENRICHMENT_FAILED" -eq 0 ]; then echo "Using enriched config: $ENRICHED_CONFIG" - exec llama stack run "$ENRICHED_CONFIG" + exec ogx stack run "$ENRICHED_CONFIG" fi fi echo "Using original config: $INPUT_CONFIG" -exec llama stack run "$INPUT_CONFIG" +exec ogx stack run "$INPUT_CONFIG" diff --git a/src/app/endpoints/a2a.py b/src/app/endpoints/a2a.py index f183cbf4f..1ff7a31d7 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 llama_stack_client import APIConnectionError, APIStatusError +from ogx_client import APIConnectionError, APIStatusError from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError from pydantic_ai.messages import ( @@ -54,13 +54,13 @@ from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from constants import MEDIA_TYPE_EVENT_STREAM from log import get_logger from models.api.requests import QueryRequest from models.config import Action -from utils.agents.query import map_agent_inference_error +from utils.agents.error_handler import map_agent_inference_error from utils.conversation_compaction import apply_compaction_blocking from utils.mcp_headers import McpHeaders, mcp_headers_dependency from utils.pydantic_ai_helpers import build_agent @@ -346,7 +346,7 @@ async def _process_task_streaming( # pylint: disable=too-many-locals ) # Get LLM client and select model - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() try: responses_params = await prepare_responses_params( client, @@ -372,7 +372,12 @@ async def _process_task_streaming( # pylint: disable=too-many-locals ) responses_params = compaction.params - agent = build_agent(client, responses_params, configuration.skills) + agent = build_agent( + client, + responses_params, + configuration, + shields=query_request.shield_ids, + ) except (AgentRunError, APIStatusError, APIConnectionError, 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) diff --git a/src/app/endpoints/conversations_v1.py b/src/app/endpoints/conversations_v1.py index 4bc9237cb..6ab693658 100644 --- a/src/app/endpoints/conversations_v1.py +++ b/src/app/endpoints/conversations_v1.py @@ -3,8 +3,8 @@ from typing import Any from fastapi import APIRouter, Depends, HTTPException, Request -from llama_stack_api import ConversationNotFoundError -from llama_stack_client import ( +from ogx_api import ConversationNotFoundError, InvalidParameterError +from ogx_client import ( APIConnectionError, APIStatusError, ) @@ -13,7 +13,7 @@ from app.database import get_session from authentication import get_auth_dependency from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.requests import ConversationUpdateRequest @@ -68,7 +68,7 @@ examples=["database", "configuration"] ), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -83,7 +83,7 @@ examples=["database", "configuration"] ), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -95,7 +95,7 @@ examples=["database", "configuration"] ), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -109,7 +109,7 @@ examples=["database", "configuration"] ), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -235,7 +235,7 @@ async def get_conversation_endpoint_handler( # pylint: disable=too-many-locals, ) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Convert to llama-stack format (add 'conv_' prefix if needed) llama_stack_conv_id = to_llama_stack_conversation_id(normalized_conv_id) @@ -276,7 +276,7 @@ async def get_conversation_endpoint_handler( # pylint: disable=too-many-locals, except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) response = ServiceUnavailableResponse( - backend_name="Llama Stack", cause=str(e) + backend_name="OGX", cause=str(e) ).model_dump() raise HTTPException(**response) from e @@ -369,7 +369,7 @@ async def delete_conversation_endpoint_handler( try: # Get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Convert to llama-stack format (add 'conv_' prefix if needed) llama_stack_conv_id = to_llama_stack_conversation_id(normalized_conv_id) @@ -385,10 +385,10 @@ async def delete_conversation_endpoint_handler( ) except APIConnectionError as e: - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e - except (APIStatusError, ConversationNotFoundError): + except (APIStatusError, ConversationNotFoundError, InvalidParameterError): # In library mode, ConversationNotFoundError is raised instead of APIStatusError logger.warning( "Conversation %s in LlamaStack not found. Treating as already deleted.", @@ -482,7 +482,7 @@ async def update_conversation_endpoint_handler( try: # Get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Convert to llama-stack format (add 'conv_' prefix if needed) llama_stack_conv_id = to_llama_stack_conversation_id(normalized_conv_id) @@ -522,7 +522,7 @@ async def update_conversation_endpoint_handler( except APIConnectionError as e: response = ServiceUnavailableResponse( - backend_name="Llama Stack", cause=str(e) + backend_name="OGX", cause=str(e) ).model_dump() raise HTTPException(**response) from e diff --git a/src/app/endpoints/health.py b/src/app/endpoints/health.py index b718dc178..0294df8f4 100644 --- a/src/app/endpoints/health.py +++ b/src/app/endpoints/health.py @@ -8,12 +8,12 @@ from typing import Annotated, Any from fastapi import APIRouter, Depends, Response, status -from llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.responses.constants import UNAUTHORIZED_OPENAPI_EXAMPLES @@ -42,7 +42,7 @@ 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -64,7 +64,7 @@ async def get_providers_health_statuses() -> list[ProviderHealthStatus]: determined, returns a single entry indicating an error. """ try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() providers = await client.providers.list() logger.debug("Found %d providers", len(providers)) @@ -110,7 +110,7 @@ async def check_default_model_available() -> tuple[bool, str]: expected_model_id = f"{inference.default_provider}/{inference.default_model}" - client_holder = AsyncLlamaStackClientHolder() + client_holder = AsyncOgxClientHolder() return await client_holder.check_model_available(expected_model_id) diff --git a/src/app/endpoints/info.py b/src/app/endpoints/info.py index 52490e611..569966af7 100644 --- a/src/app/endpoints/info.py +++ b/src/app/endpoints/info.py @@ -3,12 +3,12 @@ from typing import Annotated, Any from fastapi import APIRouter, Depends, HTTPException, Request -from llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.responses.constants import UNAUTHORIZED_OPENAPI_EXAMPLES @@ -30,7 +30,7 @@ 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -70,7 +70,7 @@ async def info_endpoint_handler( try: # try to get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # retrieve version llama_stack_version_object = await client.inspect.version() llama_stack_version = llama_stack_version_object.version @@ -85,5 +85,5 @@ async def info_endpoint_handler( # connection to Llama Stack server except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e diff --git a/src/app/endpoints/mcp_servers.py b/src/app/endpoints/mcp_servers.py index 045334a49..8ec6c5583 100644 --- a/src/app/endpoints/mcp_servers.py +++ b/src/app/endpoints/mcp_servers.py @@ -3,13 +3,10 @@ from typing import Annotated, Any from fastapi import APIRouter, Depends, HTTPException, Request, status -from llama_stack_api.common.errors import ToolGroupNotFoundError -from llama_stack_client import APIConnectionError, NotFoundError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder from configuration import configuration from log import get_logger from models.api.requests import MCPServerRegistrationRequest @@ -18,7 +15,6 @@ ConflictResponse, ForbiddenResponse, InternalServerErrorResponse, - ServiceUnavailableResponse, UnauthorizedResponse, ) from models.api.responses.successful import ( @@ -42,9 +38,6 @@ 500: InternalServerErrorResponse.openapi_response( examples=["configuration", "mcp server registration"] ), - 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] - ), } @@ -61,8 +54,8 @@ async def register_mcp_server_handler( ) -> MCPServerRegistrationResponse: """Register an MCP server dynamically at runtime. - Adds the MCP server to the runtime configuration and registers it - as a toolgroup with Llama Stack so it becomes available for queries. + Adds the MCP server to the runtime configuration so it becomes available + for queries. ### Parameters: - request: Model containing attributes to dynamically registering an MCP server. @@ -70,8 +63,7 @@ async def register_mcp_server_handler( - body: Headers that should be passed to MCP servers. ### Raises: - - HTTPException: On duplicate name, Llama Stack connection error, or - registration failure. + - HTTPException: On duplicate name or registration failure. ### Returns: - MCPServerRegistrationResponse: Details of the newly registered server. @@ -81,13 +73,8 @@ async def register_mcp_server_handler( check_configuration_loaded(configuration) - mcp_server = ModelContextProtocolServer( - name=body.name, - url=body.url, - provider_id=body.provider_id, - authorization_headers=body.authorization_headers or {}, - headers=body.headers or [], - timeout=body.timeout, + mcp_server = ModelContextProtocolServer.model_validate( + body.model_dump(exclude_none=True) ) try: @@ -96,24 +83,6 @@ async def register_mcp_server_handler( response = ConflictResponse(resource="MCP server", resource_id=body.name) raise HTTPException(**response.model_dump()) from e - try: - client = AsyncLlamaStackClientHolder().get_client() - await client.toolgroups.register( # pyright: ignore[reportDeprecated] - toolgroup_id=mcp_server.name, - provider_id=mcp_server.provider_id, - mcp_endpoint={"uri": mcp_server.url}, - ) - except APIConnectionError as e: - configuration.remove_mcp_server(body.name) - logger.error("Failed to register MCP server with Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) - raise HTTPException(**response.model_dump()) from e - except Exception as e: # pylint: disable=broad-exception-caught - configuration.remove_mcp_server(body.name) - logger.error("Failed to register MCP toolgroup: %s", e) - error_response = InternalServerErrorResponse.mcp_server_registration_failed() - raise HTTPException(**error_response.model_dump()) from e - logger.info("Dynamically registered MCP server: %s at %s", body.name, body.url) return MCPServerRegistrationResponse( @@ -129,7 +98,6 @@ async def register_mcp_server_handler( 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), - 503: ServiceUnavailableResponse.openapi_response(examples=["kubernetes api"]), } @@ -159,17 +127,15 @@ async def list_mcp_servers_handler( check_configuration_loaded(configuration) - servers = [] - for mcp in configuration.mcp_servers: - source = "api" if configuration.is_dynamic_mcp_server(mcp.name) else "config" - servers.append( - MCPServerInfo( - name=mcp.name, - url=mcp.url, - provider_id=mcp.provider_id, - source=source, - ) + servers = [ + MCPServerInfo( + name=mcp.name, + url=mcp.url, + provider_id=mcp.provider_id, + source="api" if configuration.is_dynamic_mcp_server(mcp.name) else "config", ) + for mcp in configuration.mcp_servers + ] return MCPServerListResponse(servers=servers) @@ -179,9 +145,6 @@ async def list_mcp_servers_handler( 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint", "mcp server static"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), - 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] - ), } @@ -194,9 +157,9 @@ async def delete_mcp_server_handler( ) -> MCPServerDeleteResponse: """Unregister a dynamically registered MCP server. - Removes the MCP server from the runtime configuration and unregisters - its toolgroup from Llama Stack. Only servers registered via the API - can be deleted; statically configured servers cannot be removed. + Removes the MCP server from the runtime configuration. Only servers + registered via the API can be deleted; statically configured servers + cannot be removed. ### Parameters: - request: The incoming HTTP request (used by middleware). @@ -204,8 +167,7 @@ async def delete_mcp_server_handler( - name: MCP server name ### Raises: - - HTTPException: If the server is not found, is statically configured, or - Llama Stack unregistration fails. + - HTTPException: If the server is not found or is statically configured. ### Returns: - MCPServerDeleteResponse: Confirmation of the deletion. @@ -221,20 +183,6 @@ async def delete_mcp_server_handler( response = ForbiddenResponse.mcp_server_static_config(name) raise HTTPException(**response.model_dump()) - try: - client = AsyncLlamaStackClientHolder().get_client() - await client.toolgroups.unregister( # pyright: ignore[reportDeprecated] - toolgroup_id=name - ) - except APIConnectionError as e: - logger.error("Failed to connect to Llama Stack: %s", e) - svc_response = ServiceUnavailableResponse( - backend_name="Llama Stack", cause=str(e) - ) - raise HTTPException(**svc_response.model_dump()) from e - except (ToolGroupNotFoundError, NotFoundError): - logger.warning("MCP server not found, treating as already deleted.") - try: configuration.remove_mcp_server(name) local_deleted = True diff --git a/src/app/endpoints/metrics.py b/src/app/endpoints/metrics.py index 4b44799c3..f984292ab 100644 --- a/src/app/endpoints/metrics.py +++ b/src/app/endpoints/metrics.py @@ -29,7 +29,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } diff --git a/src/app/endpoints/models.py b/src/app/endpoints/models.py index fa435a6f9..37491f700 100644 --- a/src/app/endpoints/models.py +++ b/src/app/endpoints/models.py @@ -4,12 +4,12 @@ from fastapi import APIRouter, HTTPException, Query, Request from fastapi.params import Depends -from llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.requests.catalog import ModelFilter @@ -23,51 +23,19 @@ from models.api.responses.successful import ModelsResponse from models.config import Action from utils.endpoints import check_configuration_loaded +from utils.model_list import parse_model_list_response logger = get_logger(__name__) router = APIRouter(tags=["models"]) -def parse_llama_stack_model(model: Any) -> dict[str, Any]: - """ - Parse llama-stack model. - - Converting the new llama-stack model format (0.4.x) with custom_metadata. - - Parameters: - model: Model object from llama-stack (has id, custom_metadata, object fields) - - Returns: - dict: Model in legacy format with identifier, provider_id, model_type, etc. - """ - custom_metadata = getattr(model, "custom_metadata", {}) or {} - - model_type = str(custom_metadata.get("model_type", "unknown")) - - metadata = { - k: v - for k, v in custom_metadata.items() - if k not in ("provider_id", "provider_resource_id", "model_type") - } - - return { - "identifier": getattr(model, "id", ""), - "metadata": metadata, - "api_model_type": model_type, - "provider_id": str(custom_metadata.get("provider_id", "")), - "type": getattr(model, "object", "model"), - "provider_resource_id": str(custom_metadata.get("provider_resource_id", "")), - "model_type": model_type, - } - - models_responses: dict[int | str, dict[str, Any]] = { 200: ModelsResponse.openapi_response(), 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -119,23 +87,20 @@ async def models_endpoint_handler( check_configuration_loaded(configuration) llama_stack_configuration = configuration.llama_stack_configuration - logger.info("Llama stack config: %s", llama_stack_configuration) + logger.info("Llama Stack config: %s", llama_stack_configuration) try: # try to get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() - # retrieve models - models = await client.models.list() - - # parse models to legacy format - parsed_models = [parse_llama_stack_model(model) for model in models] + client = AsyncOgxClientHolder().get_client() + # retrieve and normalize models across OpenAI/Anthropic/Google list shapes + parsed_models = parse_model_list_response(await client.models.list()) # optional filtering by model type if model_type.model_type is not None: parsed_models = [ model for model in parsed_models - if model["model_type"] == model_type.model_type + if model.model_type == model_type.model_type ] return ModelsResponse(models=parsed_models) @@ -143,5 +108,5 @@ async def models_endpoint_handler( # Connection to Llama Stack server failed except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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 6c85603c8..d270b3966 100644 --- a/src/app/endpoints/prompts.py +++ b/src/app/endpoints/prompts.py @@ -3,14 +3,14 @@ from typing import Annotated, Any, Optional from fastapi import APIRouter, Depends, HTTPException, Request -from llama_stack_client import APIConnectionError, BadRequestError -from llama_stack_client import APIStatusError as LLSApiStatusError +from ogx_client import APIConnectionError, BadRequestError +from ogx_client import APIStatusError as LLSApiStatusError from openai._exceptions import APIStatusError as OpenAIAPIStatusError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.requests import PromptCreateRequest, PromptUpdateRequest @@ -44,7 +44,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint", "prompt manage"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -54,7 +54,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint", "prompt read"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -66,7 +66,7 @@ 404: NotFoundResponse.openapi_response(examples=["prompt"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -78,7 +78,7 @@ 404: NotFoundResponse.openapi_response(examples=["prompt"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -89,7 +89,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint", "prompt manage"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -135,13 +135,13 @@ async def create_prompt_handler( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() 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="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e except (LLSApiStatusError, OpenAIAPIStatusError) as e: logger.error("API status error while creating prompt: %s", e) @@ -184,13 +184,13 @@ async def list_prompts_handler( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() 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="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e except (LLSApiStatusError, OpenAIAPIStatusError) as e: logger.error("API status error while listing prompts: %s", e) @@ -244,7 +244,7 @@ async def get_prompt_handler( raise HTTPException(**response.model_dump()) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() if version is not None: retrieved = await client.prompts.retrieve(prompt_id, version=version) else: @@ -252,7 +252,7 @@ async def get_prompt_handler( 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="Llama Stack", cause=str(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) @@ -315,13 +315,13 @@ async def update_prompt_handler( raise HTTPException(**response.model_dump()) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() 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="Llama Stack", cause=str(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) @@ -381,12 +381,12 @@ async def delete_prompt_handler( raise HTTPException(**response.model_dump()) try: - client = AsyncLlamaStackClientHolder().get_client() + 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="Llama Stack", cause=str(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) diff --git a/src/app/endpoints/providers.py b/src/app/endpoints/providers.py index e6cb8ed07..6ff11447f 100644 --- a/src/app/endpoints/providers.py +++ b/src/app/endpoints/providers.py @@ -4,13 +4,13 @@ from fastapi import APIRouter, HTTPException, Request from fastapi.params import Depends -from llama_stack_client import APIConnectionError, BadRequestError -from llama_stack_client.types import ProviderListResponse +from ogx_client import APIConnectionError, BadRequestError +from ogx_client.types import ProviderListResponse from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.responses.constants import UNAUTHORIZED_OPENAPI_EXAMPLES @@ -38,7 +38,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -49,7 +49,7 @@ 404: NotFoundResponse.openapi_response(examples=["provider"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -87,14 +87,14 @@ async def providers_endpoint_handler( check_configuration_loaded(configuration) llama_stack_configuration = configuration.llama_stack_configuration - logger.info("Llama stack config: %s", llama_stack_configuration) + logger.info("Llama Stack config: %s", llama_stack_configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() providers: ProviderListResponse = await client.providers.list() except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e return ProvidersListResponse(providers=group_providers(providers)) @@ -157,16 +157,16 @@ async def get_provider_endpoint_handler( check_configuration_loaded(configuration) llama_stack_configuration = configuration.llama_stack_configuration - logger.info("Llama stack config: %s", llama_stack_configuration) + logger.info("Llama Stack config: %s", llama_stack_configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() 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="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e except BadRequestError as e: diff --git a/src/app/endpoints/query.py b/src/app/endpoints/query.py index 92312fc51..b7f73ed3c 100644 --- a/src/app/endpoints/query.py +++ b/src/app/endpoints/query.py @@ -1,27 +1,15 @@ """Handler for REST API call to provide answer to query using Response API.""" import datetime -from typing import Annotated, Any, Optional, cast +from typing import Annotated, Any -from fastapi import APIRouter, Depends, HTTPException, Request -from llama_stack_api.openai_responses import OpenAIResponseObject -from llama_stack_client import ( - APIConnectionError, - AsyncLlamaStackClient, -) -from llama_stack_client import ( - APIStatusError as LLSApiStatusError, -) -from openai._exceptions import ( - APIStatusError as OpenAIAPIStatusError, -) -from typing_extensions import deprecated +from fastapi import APIRouter, Depends, Request from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.azure_token_manager import AzureEntraIDManager from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from constants import ENDPOINT_PATH_QUERY, IMAGE_CONTENT_TYPES from log import get_logger @@ -38,18 +26,12 @@ UnprocessableEntityResponse, ) from models.api.responses.successful import QueryResponse -from models.common.moderation import ShieldModerationResult -from models.common.responses.responses_api_params import ResponsesApiParams -from models.common.responses.types import ResponseInput -from models.common.turn_summary import TurnSummary from models.config import Action from utils.agents.query import retrieve_agent_response from utils.conversation_compaction import ( apply_compaction_blocking, configured_conversation_cache, - store_compacted_turn, ) -from utils.conversations import append_turn_items_to_conversation from utils.endpoints import ( check_configuration_loaded, validate_and_retrieve_conversation, @@ -58,8 +40,6 @@ from utils.mcp_oauth_probe import check_mcp_auth from utils.query import ( consume_query_tokens, - handle_known_apistatus_errors, - is_context_length_error, prepare_input, store_query_results, validate_attachments_metadata, @@ -67,9 +47,7 @@ ) from utils.quota_utils import check_tokens_available, get_available_quotas from utils.responses import ( - build_turn_summary, deduplicate_referenced_documents, - extract_vector_store_ids_from_tools, maybe_get_topic_summary, prepare_responses_params, ) @@ -96,7 +74,7 @@ 429: QuotaExceededResponse.openapi_response(), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -133,7 +111,7 @@ async def query_endpoint_handler( - 422: Unprocessable Entity - Request validation failed - 429: Quota limit exceeded - The token quota for model or user has been exceeded - 500: Internal Server Error - Configuration not loaded or other server errors - - 503: Service Unavailable - Unable to connect to Llama Stack backend + - 503: Service Unavailable - Unable to connect to OGX backend """ check_configuration_loaded(configuration) @@ -172,7 +150,7 @@ async def query_endpoint_handler( in request.state.authorized_actions, ) - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Moderation input is the raw user content (query + attachments) without injected RAG # context, to avoid false positives from retrieved document content. @@ -225,7 +203,7 @@ async def query_endpoint_handler( and AzureEntraIDManager().is_token_expired and AzureEntraIDManager().refresh_token() ): - client = await AsyncLlamaStackClientHolder().update_azure_token() + client = await AsyncOgxClientHolder().update_azure_token() # Extract image attachments for multimodal support image_attachments = [ @@ -241,6 +219,7 @@ async def query_endpoint_handler( moderation_result, endpoint_path, compaction.original_input if compaction.compacted else None, + shield_ids=query_request.shield_ids, no_tools=bool(query_request.no_tools), image_attachments=image_attachments, ) @@ -312,94 +291,3 @@ async def query_endpoint_handler( output_tokens=turn_summary.token_usage.output_tokens, available_quotas=available_quotas, ) - - -@deprecated( - "Deprecated in favor of utils.agents.query.retrieve_agent_response.", - stacklevel=2, -) -async def retrieve_response( - client: AsyncLlamaStackClient, - responses_params: ResponsesApiParams, - moderation_result: ShieldModerationResult, - endpoint_path: str = "", - original_input: Optional[ResponseInput] = None, -) -> TurnSummary: - """ - Retrieve response from LLMs and agents. - - Retrieves a response from the Llama Stack LLM using the Responses API. - This function processes the prepared request and returns the LLM response. - - Parameters: - ---------- - client: The AsyncLlamaStackClient to use for the request. - responses_params: The Responses API parameters. - moderation_result: The moderation result. - endpoint_path: The request path, for metrics/telemetry. - original_input: Set only in compacted mode (LCORE-1572). It is the new - user query before the explicit-input rewrite. When provided, the - turn is appended to the conversation here, because the conversation - parameter is no longer passed to Llama Stack and so the turn is not - stored automatically. - - Returns: - ------- - TurnSummary: Summary of the LLM response content - """ - response: Optional[OpenAIResponseObject] = None - # In compacted mode, the new turn must be stored against the original user - # query, not the explicit summaries-plus-recent input we send to inference. - turn_input = ( - original_input if original_input is not None else responses_params.input - ) - if moderation_result.decision == "blocked": - await append_turn_items_to_conversation( - client, - responses_params.conversation, - turn_input, - [moderation_result.refusal_response], - ) - return TurnSummary( - id=moderation_result.moderation_id, llm_response=moderation_result.message - ) - try: - response = await client.responses.create( - **responses_params.model_dump(exclude_none=True) - ) - response = cast(OpenAIResponseObject, response) - - except RuntimeError as e: # library mode wraps 413 into runtime error - if is_context_length_error(str(e)): - error_response = PromptTooLongResponse(model=responses_params.model) - raise HTTPException(**error_response.model_dump()) from e - raise e - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: - error_response = handle_known_apistatus_errors(e, responses_params.model) - raise HTTPException(**error_response.model_dump()) from e - - # In compacted mode, store the completed turn ourselves (the conversation - # parameter was not sent, so Llama Stack did not persist it). - if original_input is not None: - await store_compacted_turn( - client, - responses_params.conversation, - original_input, - response.output, - ) - - vector_store_ids = extract_vector_store_ids_from_tools(responses_params.tools) - rag_id_mapping = configuration.rag_id_mapping - return build_turn_summary( - response, - responses_params.model, - endpoint_path, - vector_store_ids, - rag_id_mapping, - ) diff --git a/src/app/endpoints/rags.py b/src/app/endpoints/rags.py index 8cf5c7679..5e6c1d55e 100644 --- a/src/app/endpoints/rags.py +++ b/src/app/endpoints/rags.py @@ -4,12 +4,12 @@ from fastapi import APIRouter, HTTPException, Request from fastapi.params import Depends -from llama_stack_client import APIConnectionError, BadRequestError +from ogx_client import APIConnectionError, BadRequestError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.responses.constants import UNAUTHORIZED_OPENAPI_EXAMPLES @@ -37,7 +37,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -48,7 +48,7 @@ 404: NotFoundResponse.openapi_response(examples=["rag"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -87,11 +87,11 @@ async def rags_endpoint_handler( check_configuration_loaded(configuration) llama_stack_configuration = configuration.llama_stack_configuration - logger.info("Llama stack config: %s", llama_stack_configuration) + logger.info("Llama Stack config: %s", llama_stack_configuration) try: # try to get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # retrieve list of RAGs rags = await client.vector_stores.list() logger.info("List of rags: %d", len(rags.data)) @@ -108,7 +108,7 @@ async def rags_endpoint_handler( # connection to Llama Stack server except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e @@ -174,7 +174,7 @@ async def get_rag_endpoint_handler( check_configuration_loaded(configuration) llama_stack_configuration = configuration.llama_stack_configuration - logger.info("Llama stack config: %s", llama_stack_configuration) + logger.info("Llama Stack config: %s", llama_stack_configuration) # Resolve user-facing rag_id to llama-stack vector_db_id vector_db_id = _resolve_rag_id_to_vector_db_id( @@ -183,7 +183,7 @@ async def get_rag_endpoint_handler( try: # try to get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # retrieve info about RAG rag_info = await client.vector_stores.retrieve(vector_db_id) @@ -204,7 +204,7 @@ async def get_rag_endpoint_handler( ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index 568a99ae1..47035eb9b 100644 --- a/src/app/endpoints/responses.py +++ b/src/app/endpoints/responses.py @@ -10,21 +10,21 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request from fastapi.responses import StreamingResponse -from llama_stack_api import ( +from ogx_api import ( OpenAIResponseObject, OpenAIResponseObjectStream, OpenAIResponseOutput, ) -from llama_stack_api import ( +from ogx_api import ( OpenAIResponseObjectStreamResponseOutputItemAdded as OutputItemAddedChunk, ) -from llama_stack_api import ( +from ogx_api import ( OpenAIResponseObjectStreamResponseOutputItemDone as OutputItemDoneChunk, ) -from llama_stack_client import ( +from ogx_client import ( APIConnectionError, ) -from llama_stack_client import ( +from ogx_client import ( APIStatusError as LLSApiStatusError, ) from openai._exceptions import ( @@ -40,7 +40,7 @@ from authentication.interface import AuthTuple from authorization.azure_token_manager import AzureEntraIDManager from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from constants import ENDPOINT_PATH_RESPONSES, SUBSTITUTED_INSTRUCTIONS_PLACEHOLDER from log import get_logger @@ -105,7 +105,7 @@ select_model_for_responses, ) from utils.rh_identity import get_rh_identity_context -from utils.shields import run_shield_moderation +from utils.shields import run_shield_moderation_v2 from utils.suid import ( normalize_conversation_id, ) @@ -161,7 +161,7 @@ def _get_user_agent(request: Request) -> Optional[str]: 429: QuotaExceededResponse.openapi_response(), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -185,7 +185,7 @@ def _http_exception_for_response_api_error( error_response = PromptTooLongResponse(model=api_params.model) elif isinstance(error, APIConnectionError): error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(error), ) elif isinstance(error, (LLSApiStatusError, OpenAIAPIStatusError)): @@ -352,7 +352,7 @@ async def responses_endpoint_handler( - 422: Unprocessable Entity - Request validation failed - 429: Quota limit exceeded - The token quota for model or user has been exceeded - 500: Internal Server Error - Configuration not loaded or other server errors - - 503: Service Unavailable - Unable to connect to Llama Stack backend + - 503: Service Unavailable - Unable to connect to OGX backend """ original_request = responses_request # read-only request updated_request = responses_request.model_copy(deep=True) @@ -395,7 +395,7 @@ async def responses_endpoint_handler( ) updated_request.conversation = response_context.conversation updated_request.generate_topic_summary = response_context.generate_topic_summary - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # LCORE-specific: Automatically select model if not provided in request # This extends the base LLS API which requires model to be specified. @@ -414,7 +414,7 @@ async def responses_endpoint_handler( and AzureEntraIDManager().is_token_expired and AzureEntraIDManager().refresh_token() ): - client = await AsyncLlamaStackClientHolder().update_azure_token() + client = await AsyncOgxClientHolder().update_azure_token() input_text = ( original_request.input @@ -424,11 +424,11 @@ async def responses_endpoint_handler( attachments_text = extract_attachments_text(original_request.input) endpoint_path = ENDPOINT_PATH_RESPONSES - moderation_result = await run_shield_moderation( - client, + + moderation_result = await run_shield_moderation_v2( input_text + "\n\n" + attachments_text, - endpoint_path, - original_request.shield_ids, + configuration.configuration.shields, + responses_request.shield_ids, ) filter_server_tools = ( @@ -559,7 +559,7 @@ async def handle_streaming_response( """Handle streaming response from Responses API. Args: - client: The AsyncLlamaStackClient instance + client: The AsyncOgxClient instance original_request: Original request (read-only) api_params: API parameters responses_context: Responses context @@ -582,7 +582,9 @@ async def handle_streaming_response( inference_start_time = time.monotonic() try: response = await context.client.responses.create( - **api_params.model_dump(exclude_none=True) + **api_params.model_dump( + exclude_none=True, exclude={"safety_identifier"} + ) ) generator = response_generator( stream=cast(AsyncIterator[OpenAIResponseObjectStream], response), @@ -903,6 +905,9 @@ async def response_generator( chunk_dict["response"]["conversation"] = normalize_conversation_id( api_params.conversation ) + chunk_dict["response"][ + "safety_identifier" + ] = api_params.safety_identifier _sanitize_response_dict( chunk_dict["response"], configured_mcp_labels, @@ -1082,7 +1087,9 @@ async def handle_non_streaming_response( api_response = cast( OpenAIResponseObject, await context.client.responses.create( - **api_params.model_dump(exclude_none=True) + **api_params.model_dump( + exclude_none=True, exclude={"safety_identifier"} + ) ), ) _record_response_inference_result( @@ -1182,6 +1189,7 @@ async def handle_non_streaming_response( response = ResponsesResponse.model_validate( { **response_dict, + "safety_identifier": api_params.safety_identifier, "available_quotas": available_quotas, "conversation": normalize_conversation_id(api_params.conversation), "completed_at": int(completed_at.timestamp()), diff --git a/src/app/endpoints/rlsapi_v1.py b/src/app/endpoints/rlsapi_v1.py index febc7caeb..8e9a3d446 100644 --- a/src/app/endpoints/rlsapi_v1.py +++ b/src/app/endpoints/rlsapi_v1.py @@ -12,8 +12,8 @@ import jinja2 from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request from jinja2.sandbox import SandboxedEnvironment -from llama_stack_api.openai_responses import OpenAIResponseObject -from llama_stack_client import APIConnectionError, APIStatusError, RateLimitError +from ogx_api.openai_responses import OpenAIResponseObject +from ogx_client import APIConnectionError, APIStatusError, RateLimitError from openai._exceptions import APIStatusError as OpenAIAPIStatusError import constants @@ -21,7 +21,7 @@ from authentication.interface import AuthTuple from authorization.azure_token_manager import AzureEntraIDManager from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from constants import ENDPOINT_PATH_INFER from log import get_logger @@ -42,9 +42,11 @@ RlsapiV1InferData, RlsapiV1InferResponse, ) -from models.config import Action +from models.config import Action, RedactionConfig from observability import InferenceEventData, build_inference_event, send_splunk_event +from pydantic_ai_lightspeed.capabilities.redaction.core import redact_text from utils.endpoints import check_configuration_loaded +from utils.model_list import parse_model_list_response from utils.query import ( consume_query_tokens, extract_provider_and_model_from_model_id, @@ -61,7 +63,7 @@ get_mcp_tools, ) from utils.rh_identity import AUTH_DISABLED, get_rh_identity_context -from utils.shields import run_shield_moderation +from utils.shields import run_shield_moderation_v2 from utils.suid import get_suid logger = get_logger(__name__) @@ -94,7 +96,7 @@ class TemplateRenderError(Exception): 429: QuotaExceededResponse.openapi_response(), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -184,12 +186,12 @@ async def _get_default_model_id() -> str: "No complete default model configured for rlsapi v1, " "auto-discovering LLM model" ) - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() try: - models = await client.models.list() + models = parse_model_list_response(await client.models.list()) except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e @@ -197,11 +199,7 @@ async def _get_default_model_id() -> str: error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e - llm_models = [ - m - for m in models - if m.custom_metadata and m.custom_metadata.get("model_type") == "llm" - ] + llm_models = [m for m in models if m.model_type == "llm"] if not llm_models: msg = "No LLM model found in available models" logger.error(msg) @@ -212,8 +210,8 @@ async def _get_default_model_id() -> str: raise HTTPException(**error_response.model_dump()) model = llm_models[0] - logger.info("Auto-discovered LLM model for rlsapi v1: %s", model.id) - return model.id + logger.info("Auto-discovered LLM model for rlsapi v1: %s", model.identifier) + return model.identifier async def _resolve_validated_model_id() -> str: @@ -230,7 +228,7 @@ async def _resolve_validated_model_id() -> str: HTTPException: 503 if Llama Stack is unreachable during resolution or validation. """ model_id = await _get_default_model_id() - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() if not await check_model_configured(client, model_id): _, model_name = extract_provider_and_model_from_model_id(model_id) error_response = NotFoundResponse(resource="model", resource_id=model_name) @@ -264,7 +262,7 @@ async def _call_llm( APIConnectionError: If the Llama Stack service is unreachable. HTTPException: 503 if no default model is configured. """ - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() resolved_model_id = model_id or await _get_default_model_id() # Handle Azure token refresh if needed @@ -274,7 +272,7 @@ async def _call_llm( and AzureEntraIDManager().is_token_expired and AzureEntraIDManager().refresh_token() ): - client = await AsyncLlamaStackClientHolder().update_azure_token() + client = await AsyncOgxClientHolder().update_azure_token() logger.debug("Using model %s for rlsapi v1 inference", resolved_model_id) @@ -354,12 +352,14 @@ async def _check_shield_moderation( # pylint: disable=too-many-arguments,too-ma background_tasks: BackgroundTasks, infer_request: RlsapiV1InferRequest, request: Request, - endpoint_path: str, -) -> Optional[RlsapiV1InferResponse]: - """Run shield moderation and return a refusal response if blocked. +) -> tuple[Optional[RlsapiV1InferResponse], str]: + """Run shield moderation and return the moderation outcome. - Uses all configured shields in Llama Stack. When no shields are - registered, moderation is a no-op and returns None immediately. + Iterates ``configuration.shields`` in order. Redaction shields apply + PII substitution to the input text (the redacted text is forwarded + to inference). All other shields (e.g. question validity) are run + via ``run_shield_moderation_v2``; the first block short-circuits + with a refusal response and Splunk telemetry event. Args: input_text: The combined user input to moderate. @@ -367,19 +367,39 @@ async def _check_shield_moderation( # pylint: disable=too-many-arguments,too-ma background_tasks: FastAPI background tasks for async Splunk event sending. infer_request: The original inference request (for Splunk event context). request: The FastAPI request object (for Splunk event context). - endpoint_path: The API endpoint path for metric labeling. Returns: - An RlsapiV1InferResponse containing the refusal message if the input - was blocked, or None if moderation passed. + A tuple of (refusal_response, moderated_input). refusal_response is + None when moderation passed; moderated_input is the (possibly + redacted) text to forward to inference. """ - client = AsyncLlamaStackClientHolder().get_client() logger.info("Running shield moderation for rlsapi v1 request %s", request_id) - moderation_result = await run_shield_moderation(client, input_text, endpoint_path) + + moderated_input = input_text + non_redaction_shields = [] + + for shield_config in configuration.shields: + if isinstance(shield_config.config, RedactionConfig): + result = redact_text( + moderated_input, shield_config.config.compiled_patterns + ) + if result.redacted: + logger.info( + "PII redaction applied for rlsapi v1 request %s (%d substitutions)", + request_id, + result.redaction_count, + ) + moderated_input = result.content + else: + non_redaction_shields.append(shield_config) + + moderation_result = await run_shield_moderation_v2( + moderated_input, non_redaction_shields + ) if moderation_result.decision != "blocked": logger.info("Shield moderation passed for rlsapi v1 request %s", request_id) - return None + return None, moderated_input logger.info("Shield moderation blocked rlsapi v1 request %s", request_id) _queue_splunk_event( @@ -391,17 +411,20 @@ async def _check_shield_moderation( # pylint: disable=too-many-arguments,too-ma 0.0, "infer_shield_blocked", ) - return RlsapiV1InferResponse( - data=RlsapiV1InferData( - text=moderation_result.message, - request_id=request_id, - tool_calls=None, - tool_results=None, - rag_chunks=None, - referenced_documents=None, - input_tokens=None, - output_tokens=None, - ) + return ( + RlsapiV1InferResponse( + data=RlsapiV1InferData( + text=moderation_result.message, + request_id=request_id, + tool_calls=None, + tool_results=None, + rag_chunks=None, + referenced_documents=None, + input_tokens=None, + output_tokens=None, + ) + ), + moderated_input, ) @@ -595,12 +618,12 @@ def _map_inference_error_to_http_exception( # pylint: disable=too-many-return-s if isinstance(error, APIConnectionError): logger.error( - "Unable to connect to Llama Stack for request %s: %s", + "Unable to connect to OGX for request %s: %s", request_id, type(error).__name__, ) error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause="Unable to connect to the inference backend", ) return HTTPException(**error_response.model_dump()) @@ -686,13 +709,12 @@ async def infer_endpoint( # pylint: disable=R0914,R0915 # Uses all configured shields; no-op when no shields are registered. # Runs before model/tool discovery so blocked requests short-circuit # without incurring external I/O. - blocked_response = await _check_shield_moderation( + blocked_response, moderated_input = await _check_shield_moderation( input_source, request_id, background_tasks, infer_request, request, - endpoint_path, ) if blocked_response is not None: return blocked_response @@ -727,7 +749,7 @@ async def infer_endpoint( # pylint: disable=R0914,R0915 logger.info("Building instructions for rlsapi v1 request %s", request_id) instructions = _build_instructions(infer_request.context.systeminfo) response = await _call_llm( - input_source, + moderated_input, instructions, tools=cast(list[Any], mcp_tools), model_id=model_id, diff --git a/src/app/endpoints/shields.py b/src/app/endpoints/shields.py index 641064234..4c257cc34 100644 --- a/src/app/endpoints/shields.py +++ b/src/app/endpoints/shields.py @@ -2,24 +2,22 @@ from typing import Annotated, Any -from fastapi import APIRouter, HTTPException, Request +from fastapi import APIRouter, Request from fastapi.params import Depends -from llama_stack_client import APIConnectionError from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder from configuration import configuration from log import get_logger from models.api.responses.constants import UNAUTHORIZED_OPENAPI_EXAMPLES from models.api.responses.error import ( ForbiddenResponse, InternalServerErrorResponse, - ServiceUnavailableResponse, UnauthorizedResponse, ) from models.api.responses.successful import ShieldsResponse +from models.common.shields import CatalogShield from models.config import Action from utils.endpoints import check_configuration_loaded @@ -32,9 +30,6 @@ 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), - 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] - ), } @@ -48,7 +43,7 @@ async def shields_endpoint_handler( Handle requests to the /shields endpoint. Process GET requests to the /shields endpoint, returning a list of available - shields from the Llama Stack service. + shields from Lightspeed Core Stack configuration. ### Parameters: - request: The incoming HTTP request (used by middleware). @@ -59,8 +54,6 @@ async def shields_endpoint_handler( - HTTPException: with status 403 if permission is denied. - HTTPException: with status 500 and a detail object containing `response` and `cause` when service configuration is wrong or incomplete. - - HTTPException: with status 503 and a detail object containing `response` - and `cause` when unable to connect to Llama Stack. ### Returns: - ShieldsResponse: An object containing the list of available shields. @@ -73,19 +66,9 @@ async def shields_endpoint_handler( check_configuration_loaded(configuration) - llama_stack_configuration = configuration.llama_stack_configuration - logger.info("Llama stack config: %s", llama_stack_configuration) - - try: - # try to get Llama Stack client - client = AsyncLlamaStackClientHolder().get_client() - # retrieve shields - shields = await client.shields.list() - s = [dict(s) for s in shields] - return ShieldsResponse(shields=s) - - # connection to Llama Stack server - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) - raise HTTPException(**response.model_dump()) from e + shields = [ + CatalogShield.model_validate(shield.model_dump()) + for shield in configuration.shields + ] + logger.info("Returning %d configured shield(s)", len(shields)) + return ShieldsResponse(shields=shields) diff --git a/src/app/endpoints/streaming_query.py b/src/app/endpoints/streaming_query.py index 40f462ebc..b83ae47af 100644 --- a/src/app/endpoints/streaming_query.py +++ b/src/app/endpoints/streaming_query.py @@ -3,54 +3,27 @@ import asyncio import datetime from collections.abc import AsyncIterator -from typing import Annotated, Any, Optional, cast +from typing import Annotated, Any, Optional from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse -from llama_stack_api import ( - OpenAIResponseObject, - OpenAIResponseObjectStream, -) -from llama_stack_api import ( - OpenAIResponseObjectStreamResponseMcpCallArgumentsDone as MCPArgsDoneChunk, -) -from llama_stack_api import ( - OpenAIResponseObjectStreamResponseOutputItemAdded as OutputItemAddedChunk, -) -from llama_stack_api import ( - OpenAIResponseObjectStreamResponseOutputItemDone as OutputItemDoneChunk, -) -from llama_stack_api import ( - OpenAIResponseObjectStreamResponseOutputTextDelta as TextDeltaChunk, -) -from llama_stack_api import ( - OpenAIResponseObjectStreamResponseOutputTextDone as TextDoneChunk, -) -from llama_stack_api import ( - OpenAIResponseOutputMessageMCPCall as MCPCall, -) -from llama_stack_client import ( +from ogx_client import ( APIConnectionError, ) -from llama_stack_client import ( +from ogx_client import ( APIStatusError as LLSApiStatusError, ) from openai._exceptions import APIStatusError as OpenAIAPIStatusError -from typing_extensions import deprecated from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.azure_token_manager import AzureEntraIDManager from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from constants import ( ENDPOINT_PATH_STREAMING_QUERY, IMAGE_CONTENT_TYPES, - LLM_TOKEN_EVENT, - LLM_TOOL_CALL_EVENT, - LLM_TOOL_RESULT_EVENT, - LLM_TURN_COMPLETE_EVENT, MEDIA_TYPE_EVENT_STREAM, MEDIA_TYPE_JSON, MEDIA_TYPE_TEXT, @@ -74,7 +47,6 @@ from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams from models.common.responses.types import ResponseInput -from models.common.turn_summary import TurnSummary from models.config import Action from utils.agents.streaming import ( generate_agent_response, @@ -86,9 +58,7 @@ apply_compaction, configured_conversation_cache, needs_compaction_path, - store_compacted_turn, ) -from utils.conversations import append_turn_items_to_conversation from utils.endpoints import ( check_configuration_loaded, validate_and_retrieve_conversation, @@ -96,46 +66,27 @@ from utils.mcp_headers import McpHeaders, mcp_headers_dependency from utils.mcp_oauth_probe import check_mcp_auth from utils.query import ( - consume_query_tokens, extract_provider_and_model_from_model_id, handle_known_apistatus_errors, is_context_length_error, prepare_input, - store_query_results, validate_attachments_metadata, validate_model_provider_override, ) -from utils.quota_utils import check_tokens_available, get_available_quotas +from utils.quota_utils import check_tokens_available from utils.responses import ( - build_mcp_tool_call_from_arguments_done, - build_tool_call_summary, - build_tool_result_from_mcp_output_item_done, deduplicate_referenced_documents, - extract_token_usage, extract_vector_store_ids_from_tools, - get_topic_summary, - parse_rag_chunks, - parse_referenced_documents, prepare_responses_params, ) from utils.shields import ( run_shield_moderation, validate_shield_ids_override, ) -from utils.stream_interrupts import ( - build_interrupted_response, - deregister_stream, - persist_interrupted_turn, - register_interrupt_callback, -) from utils.streaming_sse import ( http_exception_stream_event, - shield_violation_generator, stream_compaction_event, - stream_end_event, - stream_event, stream_http_error_event, - stream_interrupted_event, stream_start_event, ) from utils.suid import get_suid, normalize_conversation_id @@ -163,7 +114,7 @@ 429: QuotaExceededResponse.openapi_response(), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -205,7 +156,7 @@ async def streaming_query_endpoint_handler( # pylint: disable=too-many-locals - 422: Unprocessable Entity - Request validation failed - 429: Quota limit exceeded - The token quota for model or user has been exceeded - 500: Internal Server Error - Configuration not loaded or other server errors - - 503: Service Unavailable - Unable to connect to Llama Stack backend + - 503: Service Unavailable - Unable to connect to OGX backend """ check_configuration_loaded(configuration) @@ -244,7 +195,7 @@ async def streaming_query_endpoint_handler( # pylint: disable=too-many-locals in request.state.authorized_actions, ) - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Moderation input is the raw user content (query + attachments) without injected RAG # context, to avoid false positives from retrieved document content. @@ -283,7 +234,7 @@ async def streaming_query_endpoint_handler( # pylint: disable=too-many-locals and AzureEntraIDManager().is_token_expired and AzureEntraIDManager().refresh_token() ): - client = await AsyncLlamaStackClientHolder().update_azure_token() + client = await AsyncOgxClientHolder().update_azure_token() request_id = get_suid() @@ -370,86 +321,6 @@ async def streaming_query_endpoint_handler( # pylint: disable=too-many-locals ) -@deprecated( - "Deprecated in favor of utils.agents.streaming.retrieve_agent_response_generator.", - stacklevel=2, -) -async def retrieve_response_generator( - responses_params: ResponsesApiParams, - context: ResponseGeneratorContext, - endpoint_path: str, -) -> tuple[AsyncIterator[str], TurnSummary]: - """ - Retrieve the appropriate response generator. - - Handles shield moderation check and retrieves response. - Returns the generator (shield violation or response generator) and turn_summary. - Fills turn_summary attributes for token usage, referenced documents, and tool calls. - - Args: - responses_params: The Responses API parameters - context: The response generator context - endpoint_path: API endpoint path used for metric labeling. - Returns: - tuple[AsyncIterator[str], TurnSummary]: The response generator and turn summary - - """ - turn_summary = TurnSummary() - try: - if context.moderation_result.decision == "blocked": - turn_summary.llm_response = context.moderation_result.message - turn_summary.id = context.moderation_result.moderation_id - turn_summary.output_items = [context.moderation_result.refusal_response] - # In compacted mode the conversation parameter was omitted, so the - # refusal turn (with the original input) is persisted by - # generate_response; storing it here too would duplicate it. - if not responses_params.omit_conversation: - await append_turn_items_to_conversation( - context.client, - responses_params.conversation, - responses_params.input, - [context.moderation_result.refusal_response], - ) - media_type = context.query_request.media_type or MEDIA_TYPE_JSON - return ( - shield_violation_generator( - context.moderation_result.message, - media_type, - ), - turn_summary, - ) - # Retrieve response stream (may raise exceptions) - response = await context.client.responses.create( - **responses_params.model_dump(exclude_none=True) - ) - # Store pre-RAG documents for later merging with tool-based RAG - return ( - response_generator( - response, - context, - turn_summary, - endpoint_path, - ), - turn_summary, - ) - # Handle know LLS client errors only at stream creation time and shield execution - except RuntimeError as e: # library mode wraps 413 into runtime error - if is_context_length_error(str(e)): - error_response = PromptTooLongResponse(model=responses_params.model) - raise HTTPException(**error_response.model_dump()) from e - raise e - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - - except (LLSApiStatusError, OpenAIAPIStatusError) as e: - error_response = handle_known_apistatus_errors(e, responses_params.model) - raise HTTPException(**error_response.model_dump()) from e - - async def shutdown_background_topic_summary_tasks() -> None: """Cancel and await outstanding background topic summary tasks on shutdown. @@ -535,7 +406,7 @@ async def generate_response_with_compaction( return except APIConnectionError as e: yield stream_http_error_event( - ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)), + ServiceUnavailableResponse(backend_name="OGX", cause=str(e)), media_type, ) return @@ -564,404 +435,3 @@ async def generate_response_with_compaction( original_input=compacted_original_input, ): yield event - - -@deprecated( - "Deprecated in favor of utils.agents.streaming.generate_agent_response.", - stacklevel=2, -) -async def generate_response( # pylint: disable=too-many-arguments,too-many-positional-arguments,too-many-locals,too-many-branches,too-many-statements - generator: AsyncIterator[str], - context: ResponseGeneratorContext, - responses_params: ResponsesApiParams, - turn_summary: TurnSummary, - emit_start: bool = True, - compacted: bool = False, - original_input: Optional[ResponseInput] = None, -) -> AsyncIterator[str]: - """Wrap a generator with cleanup logic. - - Re-yields events from the generator, handles errors, and ensures - persistence and token consumption after completion. When the - stream is interrupted via ``CancelledError``, the user query and - an interrupted response are persisted to the conversation, but - token consumption is skipped (no usage data is available). - - Args: - generator: The base generator to wrap - context: The response generator context - responses_params: The Responses API parameters - turn_summary: TurnSummary populated during streaming - emit_start: Whether to emit the SSE start event. False when the caller - (the compaction-aware wrapper) has already emitted it. - compacted: Whether the conversation is in compacted mode. When True the - conversation parameter was not sent to Llama Stack, so the completed - turn is appended to the conversation here rather than being stored - automatically. - original_input: In compacted mode, the original user input before the - explicit-input rewrite. Used to persist the completed turn with its - structured input (preserving attachments); ``None`` otherwise. - - Yields: - SSE-formatted strings from the wrapped generator - """ - persist_guard = register_interrupt_callback( - context, - responses_params, - turn_summary, - _background_topic_summary_tasks, - original_input, - ) - - stream_completed = False - try: - if emit_start: - yield stream_start_event( - conversation_id=context.conversation_id, - request_id=context.request_id, - ) - - # Re-yield all events from the generator - async for event in generator: - yield event - - stream_completed = True - - # Handle known LLS client errors during response generation time - except RuntimeError as e: # library mode wraps 413 into runtime error - error_response = ( - PromptTooLongResponse(model=responses_params.model) - if is_context_length_error(str(e)) - else InternalServerErrorResponse.generic() - ) - yield stream_http_error_event(error_response, context.query_request.media_type) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(e), - ) - yield stream_http_error_event(error_response, context.query_request.media_type) - except (LLSApiStatusError, OpenAIAPIStatusError) as e: - error_response = handle_known_apistatus_errors(e, responses_params.model) - yield stream_http_error_event(error_response, context.query_request.media_type) - except asyncio.CancelledError: - logger.info("Streaming request %s interrupted by user", context.request_id) - current_task = asyncio.current_task() - if current_task is not None: - current_task.uncancel() - full_text, suffix = build_interrupted_response(turn_summary.partial_tokens) - if not persist_guard[0]: - persist_guard[0] = True - turn_summary.llm_response = full_text - await persist_interrupted_turn( - context, - responses_params, - turn_summary, - _background_topic_summary_tasks, - original_input, - ) - yield stream_event( - {"id": turn_summary.next_chunk_id, "token": suffix}, - LLM_TOKEN_EVENT, - context.query_request.media_type or MEDIA_TYPE_JSON, - ) - yield stream_interrupted_event(context.request_id) - finally: - deregister_stream(context.request_id) - - if not stream_completed: - return - - # Post-stream side effects: only run when streaming finished successfully - - # Get topic summary for new conversations if needed - topic_summary = None - if not context.query_request.conversation_id: - should_generate = context.query_request.generate_topic_summary - if should_generate: - logger.debug("Generating topic summary for new conversation") - topic_summary = await get_topic_summary( - context.query_request.query, - context.client, - responses_params.model, - ) - - # Consume tokens - logger.info("Consuming tokens") - consume_query_tokens( - user_id=context.user_id, - model_id=responses_params.model, - token_usage=turn_summary.token_usage, - ) - # Get available quotas - logger.info("Getting available quotas") - available_quotas = get_available_quotas( - quota_limiters=configuration.quota_limiters, user_id=context.user_id - ) - - yield stream_end_event( - turn_summary.token_usage, - available_quotas, - turn_summary.referenced_documents, - context.query_request.media_type or MEDIA_TYPE_JSON, - ) - completed_at = datetime.datetime.now(datetime.UTC).strftime("%Y-%m-%dT%H:%M:%SZ") - - # In compacted mode the conversation parameter was not sent, so Llama Stack - # did not persist this turn. Append it ourselves to keep the recent-turn - # buffer and audit history intact for the next request. - if compacted: - try: - await store_compacted_turn( - context.client, - responses_params.conversation, - ( - original_input - if original_input is not None - else context.query_request.query - ), - turn_summary.output_items, - ) - except Exception: # pylint: disable=broad-except - logger.exception( - "Failed to append compacted turn to conversation for request %s", - context.request_id, - ) - - # Store query results (transcript, conversation details, cache) - logger.info("Storing query results") - store_query_results( - user_id=context.user_id, - conversation_id=context.conversation_id, - model=responses_params.model, - completed_at=completed_at, - started_at=context.started_at, - summary=turn_summary, - query=context.query_request.query, - attachments=context.query_request.attachments, - skip_userid_check=context.skip_userid_check, - topic_summary=topic_summary, - ) - - -@deprecated( - "Deprecated in favor of utils.agents.streaming.agent_response_generator.", - stacklevel=2, -) -async def response_generator( # pylint: disable=too-many-branches,too-many-statements,too-many-locals - turn_response: AsyncIterator[OpenAIResponseObjectStream], - context: ResponseGeneratorContext, - turn_summary: TurnSummary, - endpoint_path: str, -) -> AsyncIterator[str]: - """Generate SSE formatted streaming response. - - Processes streaming chunks from Llama Stack and converts them to - Server-Sent Events (SSE) format. Uses handler functions to process - different event types and populate turn_summary during streaming. - - Args: - turn_response: The streaming response from Llama Stack - context: The response generator context - turn_summary: TurnSummary to populate during streaming - endpoint_path: API endpoint path used for metric labeling. - - Yields: - SSE-formatted strings for tokens, tool calls, tool results, - turn completion, and error events. - """ - chunk_id = 0 - media_type = context.query_request.media_type or MEDIA_TYPE_JSON - text_parts: list[str] = [] - mcp_calls: dict[int, tuple[str, str]] = ( - {} - ) # output_index -> (mcp_call_id, mcp_call_name) - latest_response_object: Optional[OpenAIResponseObject] = None - - logger.debug("Starting streaming response (Responses API) processing") - - async for chunk in turn_response: - event_type = getattr(chunk, "type", None) - logger.debug("Processing chunk %d, type: %s", chunk_id, event_type) - - # Content part started - emit an empty token to kick off UI streaming - if event_type == "response.content_part.added": - event_id = chunk_id - chunk_id += 1 - turn_summary.next_chunk_id = chunk_id - yield stream_event( - { - "id": event_id, - "token": "", - }, - LLM_TOKEN_EVENT, - media_type, - ) - - # Store MCP call item info for later lookup when arguments.done event occurs - elif event_type == "response.output_item.added": - item_added_chunk = cast(OutputItemAddedChunk, chunk) - if item_added_chunk.item.type == "mcp_call": - mcp_call_item = cast(MCPCall, item_added_chunk.item) - mcp_calls[item_added_chunk.output_index] = ( - mcp_call_item.id, - mcp_call_item.name, - ) - - # Text streaming - emit token delta - elif event_type == "response.output_text.delta": - delta_chunk = cast(TextDeltaChunk, chunk) - text_parts.append(delta_chunk.delta) - turn_summary.partial_tokens.append(delta_chunk.delta) - event_id = chunk_id - chunk_id += 1 - turn_summary.next_chunk_id = chunk_id - yield stream_event( - { - "id": event_id, - "token": delta_chunk.delta, - }, - LLM_TOKEN_EVENT, - media_type, - ) - - # Final text of the output (capture, but emit at response.completed) - elif event_type == "response.output_text.done": - text_done_chunk = cast(TextDoneChunk, chunk) - turn_summary.llm_response = text_done_chunk.text - - # Emit tool call when MCP call arguments are done - elif event_type == "response.mcp_call.arguments.done": - mcp_arguments_done_chunk = cast(MCPArgsDoneChunk, chunk) - tool_call = build_mcp_tool_call_from_arguments_done( - mcp_arguments_done_chunk.output_index, - mcp_arguments_done_chunk.arguments, - mcp_calls, - ) - if tool_call: - turn_summary.tool_calls.append(tool_call) - yield stream_event( - tool_call.model_dump(), - LLM_TOOL_CALL_EVENT, - media_type, - ) - - # Process tool calls and results when output items are done - # For mcp_call, only emit result (call was already emitted when arguments.done) - # For other types, emit both call and result - elif event_type == "response.output_item.done": - output_item_done_chunk = cast(OutputItemDoneChunk, chunk) - item_type = output_item_done_chunk.item.type - # Skip message items as they are parsed separately - if item_type == "message": - continue - - output_index = output_item_done_chunk.output_index - - # For mcp_call, only emit result if call was already emitted when arguments.done - # (indicated by output_index not being in mcp_calls dict) - # If output_index is in dict, process in else branch (emit both call and result) - if item_type == "mcp_call" and output_index not in mcp_calls: - # Call was already emitted during arguments.done, only emit result - mcp_call_item = cast(MCPCall, output_item_done_chunk.item) - tool_result = build_tool_result_from_mcp_output_item_done(mcp_call_item) - turn_summary.tool_results.append(tool_result) - yield stream_event( - tool_result.model_dump(), - LLM_TOOL_RESULT_EVENT, - media_type, - ) - else: - # For all other types (and mcp_call when arguments.done didn't happen), - # emit both call and result together - tool_call, tool_result = build_tool_call_summary( - output_item_done_chunk.item - ) - if tool_call: - turn_summary.tool_calls.append(tool_call) - yield stream_event( - tool_call.model_dump(), - LLM_TOOL_CALL_EVENT, - media_type, - ) - if tool_result: - turn_summary.tool_results.append(tool_result) - yield stream_event( - tool_result.model_dump(), - LLM_TOOL_RESULT_EVENT, - media_type, - ) - - # Completed response - capture final text and response object - elif event_type == "response.completed": - latest_response_object = cast( - OpenAIResponseObject, - getattr(chunk, "response"), # noqa: B009 - ) - turn_summary.llm_response = turn_summary.llm_response or "".join(text_parts) - # Capture structured output items for compacted-mode turn storage - # (LCORE-1572), so the persisted turn keeps non-text output items - # rather than being flattened to the response text. - turn_summary.output_items = list(latest_response_object.output or []) - event_id = chunk_id - chunk_id += 1 - turn_summary.next_chunk_id = chunk_id - yield stream_event( - { - "id": event_id, - "token": turn_summary.llm_response, - }, - LLM_TURN_COMPLETE_EVENT, - media_type, - ) - - # Incomplete or failed response - emit error - elif event_type in ("response.incomplete", "response.failed"): - latest_response_object = cast( - OpenAIResponseObject, - getattr(chunk, "response"), # noqa: B009 - ) - # Capture any partial output items so a compacted-mode turn is not - # persisted with empty output on these terminals (LCORE-1572). - turn_summary.output_items = list(latest_response_object.output or []) - error_message = ( - latest_response_object.error.message - if latest_response_object.error - else "An unexpected error occurred while processing the request." - ) - error_response = ( - PromptTooLongResponse(model=context.model_id) - if is_context_length_error(error_message) - else InternalServerErrorResponse.query_failed(error_message) - ) - yield stream_http_error_event(error_response, media_type) - - logger.debug( - "Streaming complete - Tool calls: %d, Response chars: %d", - len(turn_summary.tool_calls), - len(turn_summary.llm_response), - ) - - # Extract token usage and referenced documents from the final response object - if not latest_response_object: - return - - turn_summary.token_usage = extract_token_usage( - latest_response_object.usage, context.model_id, endpoint_path - ) - # Parse tool-based referenced documents from the final response object - tool_rag_docs = parse_referenced_documents( - latest_response_object, - vector_store_ids=context.vector_store_ids, - rag_id_mapping=context.rag_id_mapping, - ) - # Combine inline RAG results (BYOK + Solr) with tool-based results - turn_summary.referenced_documents = deduplicate_referenced_documents( - context.inline_rag_context.referenced_documents + tool_rag_docs - ) - tool_rag_chunks = parse_rag_chunks( - latest_response_object, - vector_store_ids=context.vector_store_ids, - rag_id_mapping=context.rag_id_mapping, - ) - turn_summary.rag_chunks = context.inline_rag_context.rag_chunks + tool_rag_chunks diff --git a/src/app/endpoints/tools.py b/src/app/endpoints/tools.py index f0cf99511..07159030d 100644 --- a/src/app/endpoints/tools.py +++ b/src/app/endpoints/tools.py @@ -1,14 +1,13 @@ """Handler for REST API call to list available tools from MCP servers.""" -from typing import Annotated, Any, Optional +from typing import Annotated, Any -from fastapi import APIRouter, Depends, HTTPException, Request -from llama_stack_client import APIConnectionError, BadRequestError +from fastapi import APIRouter, Depends, Request from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from models.api.responses.constants import UNAUTHORIZED_OPENAPI_EXAMPLES @@ -19,7 +18,9 @@ UnauthorizedResponse, ) from models.api.responses.successful import ToolsResponse -from models.config import Action +from models.common.tools import CatalogTool +from models.config import Action, ModelContextProtocolServer +from utils.builtin_tools import get_file_search_tools from utils.endpoints import check_configuration_loaded from utils.mcp_headers import ( McpHeaders, @@ -28,84 +29,28 @@ mcp_headers_dependency, ) from utils.mcp_oauth_probe import check_mcp_auth +from utils.mcp_tools import list_mcp_tools from utils.pydantic_ai_helpers import get_agent_capability_tools -from utils.tool_formatter import format_tools_list +from utils.tool_formatter import build_catalog_tool logger = get_logger(__name__) router = APIRouter(tags=["tools"]) -def _input_schema_to_parameters( - schema: Optional[dict[str, Any]], -) -> list[dict[str, Any]]: - """Convert a JSON Schema input_schema to a flat list of parameter dicts. - - The Llama Stack SDK returns tool parameters as a JSON Schema object - (``input_schema``). This function converts that representation into - the flat parameter list format used by the tools endpoint response. - - Parameters: - ---------- - schema: JSON Schema dict with ``properties`` and ``required`` keys, - or ``None`` if the tool has no parameters. - - Returns: - ------- - A list of parameter dicts, each containing ``name``, ``description``, - ``parameter_type``, ``required``, and ``default`` keys. - """ - if not schema or "properties" not in schema: - return [] - - required_params = set(schema.get("required", [])) - return [ - { - "name": name, - "description": prop.get("description", ""), - "parameter_type": prop.get("type", "string"), - "required": name in required_params, - "default": prop.get("default"), - } - for name, prop in schema["properties"].items() - ] - - -def _normalize_tool_dict(tool_dict: dict[str, Any], toolgroup: Any) -> None: - """Normalize a ToolDef dict to the endpoint's response format. - - Remaps field names (``name`` -> ``identifier``, ``input_schema`` -> - ``parameters``) and propagates ``provider_id``/``type`` from the - parent toolgroup. Handles both missing keys and empty legacy - placeholders. - """ - if "name" in tool_dict and not tool_dict.get("identifier"): - tool_dict["identifier"] = tool_dict["name"] - tool_dict.pop("name", None) - - if "input_schema" in tool_dict and not tool_dict.get("parameters"): - tool_dict["parameters"] = _input_schema_to_parameters(tool_dict["input_schema"]) - tool_dict.pop("input_schema", None) - - if not tool_dict.get("provider_id"): - tool_dict["provider_id"] = toolgroup.provider_id - if not tool_dict.get("type"): - tool_dict["type"] = getattr(toolgroup, "type", None) or "tool" - - tools_responses: dict[int | str, dict[str, Any]] = { 200: ToolsResponse.openapi_response(), 401: UnauthorizedResponse.openapi_response(examples=UNAUTHORIZED_OPENAPI_EXAMPLES), 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @router.get("/tools", responses=tools_responses) @authorize(Action.GET_TOOLS) -async def tools_endpoint_handler( # pylint: disable=too-many-locals,too-many-statements +async def tools_endpoint_handler( # pylint: disable=too-many-locals request: Request, auth: Annotated[AuthTuple, Depends(get_auth_dependency())], mcp_headers: McpHeaders = Depends(mcp_headers_dependency), @@ -147,120 +92,85 @@ async def tools_endpoint_handler( # pylint: disable=too-many-locals,too-many-st configuration, mcp_headers, request.headers, token ) - # Check MCP Auth + # Check MCP auth await check_mcp_auth(configuration, mcp_headers, token, request.headers) - toolgroups_response = [] - try: - client = AsyncLlamaStackClientHolder().get_client() - logger.debug("Retrieving tools from all toolgroups") - toolgroups_response = await client.toolgroups.list() - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) - raise HTTPException(**response.model_dump()) from e + client = AsyncOgxClientHolder().get_client() + consolidated_tools: list[CatalogTool] = list(await get_file_search_tools(client)) - consolidated_tools = [] - mcp_server_names = ( - {mcp_server.name for mcp_server in configuration.mcp_servers} - if configuration.mcp_servers - else set() - ) - - for toolgroup in toolgroups_response: - mcp_server = None - if toolgroup.identifier in mcp_server_names: - mcp_server = next( - ( - s - for s in configuration.mcp_servers - if s.name == toolgroup.identifier - ), - None, - ) - - headers = complete_mcp_headers.get(toolgroup.identifier, {}) - if mcp_server is not None: - unresolved = find_unresolved_auth_headers( - mcp_server.authorization_headers, headers - ) - if unresolved: - logger.warning( - "Skipping MCP server %s: required %d auth headers " - "but only resolved %d", - mcp_server.name, - len(mcp_server.authorization_headers), - len(mcp_server.authorization_headers) - len(unresolved), - ) - continue - - try: - authorization = headers.pop("Authorization", None) - - tools_response = await client.tools.list( - toolgroup_id=toolgroup.identifier, - extra_headers=headers, - extra_query={"authorization": authorization}, - ) - except BadRequestError: - logger.error("Toolgroup %s is not found", toolgroup.identifier) - continue - except APIConnectionError as e: - logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse( - backend_name="Llama Stack", cause=str(e) + for mcp_server in configuration.mcp_servers: + consolidated_tools.extend( + await _list_tools_for_mcp_server( + mcp_server, + complete_mcp_headers.get(mcp_server.name, {}), ) - raise HTTPException(**response.model_dump()) from e - - # Convert tools to dict format - tools_count = 0 - server_source = "unknown" - - for tool in tools_response: - tool_dict = dict(tool) - - _normalize_tool_dict(tool_dict, toolgroup) - - # Determine server source based on toolgroup type - if mcp_server: - tool_dict["server_source"] = mcp_server.url or toolgroup.identifier - else: - # This is a built-in toolgroup - tool_dict["server_source"] = "builtin" - - consolidated_tools.append(tool_dict) - tools_count += 1 - server_source = tool_dict["server_source"] - - logger.debug( - "Retrieved %d tools from toolgroup %s (source: %s)", - tools_count, - toolgroup.identifier, - server_source, ) existing_tool_ids = { - tool.get("identifier") for tool in consolidated_tools if tool.get("identifier") + tool.identifier for tool in consolidated_tools if tool.identifier } - capability_tools = get_agent_capability_tools(configuration.skills) - for tool_dict in capability_tools: - identifier = tool_dict.get("identifier") - if identifier and identifier not in existing_tool_ids: - consolidated_tools.append(tool_dict) - existing_tool_ids.add(identifier) + for tool in get_agent_capability_tools(configuration.skills): + if tool.identifier not in existing_tool_ids: + consolidated_tools.append(tool) + existing_tool_ids.add(tool.identifier) builtin_tool_count = len( - [t for t in consolidated_tools if t.get("server_source") == "builtin"] + [tool for tool in consolidated_tools if tool.server_source == "builtin"] ) mcp_tool_count = len(consolidated_tools) - builtin_tool_count logger.info( - "Retrieved total of %d tools (%d from built-in toolgroups, %d from MCP servers)", + "Retrieved total of %d tools (%d builtin, %d from MCP servers)", len(consolidated_tools), builtin_tool_count, mcp_tool_count, ) - # Format tools with structured description parsing - formatted_tools = format_tools_list(consolidated_tools) + return ToolsResponse(tools=consolidated_tools) + + +async def _list_tools_for_mcp_server( + mcp_server: ModelContextProtocolServer, + headers: dict[str, str], +) -> list[CatalogTool]: + """Discover tools from a single configured MCP server. + + ### Parameters: + - mcp_server: MCP server configuration entry. + - headers: Resolved request headers for the server. + + ### Returns: + - Catalog tools for the server, or an empty list when skipped or failing. + """ + unresolved = find_unresolved_auth_headers( + mcp_server.authorization_headers, + headers, + ) + if unresolved: + logger.warning( + "Skipping MCP server %s: required %d auth headers but only resolved %d", + mcp_server.name, + len(mcp_server.authorization_headers), + len(mcp_server.authorization_headers) - len(unresolved), + ) + return [] + + discovered_tools = await list_mcp_tools( + endpoint=mcp_server.url, + headers=headers, + ) + if not discovered_tools: + return [] - return ToolsResponse(tools=formatted_tools) + tools = [ + build_catalog_tool( + tool, mcp_server.provider_id, mcp_server.name, mcp_server.url + ) + for tool in discovered_tools + ] + logger.debug( + "Retrieved %d tools from MCP server %s (source: %s)", + len(tools), + mcp_server.name, + mcp_server.url, + ) + return tools diff --git a/src/app/endpoints/vector_stores.py b/src/app/endpoints/vector_stores.py index ee55bc00e..6174c5651 100644 --- a/src/app/endpoints/vector_stores.py +++ b/src/app/endpoints/vector_stores.py @@ -6,11 +6,11 @@ from typing import Annotated, Any, Optional from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status -from llama_stack_client import ( +from ogx_client import ( APIConnectionError, BadRequestError, ) -from llama_stack_client import ( +from ogx_client import ( APIStatusError as LLSApiStatusError, ) from openai._exceptions import APIStatusError as OpenAIAPIStatusError @@ -18,7 +18,7 @@ from authentication import get_auth_dependency from authentication.interface import AuthTuple from authorization.middleware import authorize -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from constants import DEFAULT_MAX_FILE_UPLOAD_SIZE from log import get_logger @@ -60,7 +60,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -71,7 +71,7 @@ 404: NotFoundResponse.openapi_response(examples=["vector store"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -82,7 +82,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -93,7 +93,7 @@ 404: NotFoundResponse.openapi_response(examples=["file"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -104,7 +104,7 @@ 404: NotFoundResponse.openapi_response(examples=["vector store"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -114,7 +114,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -124,7 +124,7 @@ 403: ForbiddenResponse.openapi_response(examples=["endpoint"]), 500: InternalServerErrorResponse.openapi_response(examples=["configuration"]), 503: ServiceUnavailableResponse.openapi_response( - examples=["llama stack", "kubernetes api"] + examples=["ogx", "kubernetes api"] ), } @@ -151,7 +151,7 @@ async def create_vector_store( - 401: Authentication failed - 403: Authorization failed - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -159,7 +159,7 @@ async def create_vector_store( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Extract provider_id for extra_body (not a direct client parameter) body_dict = body.model_dump(exclude_none=True) @@ -194,7 +194,7 @@ async def create_vector_store( ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e except (LLSApiStatusError, OpenAIAPIStatusError) as e: logger.error("API status error while creating vector store: %s", e) @@ -222,7 +222,7 @@ async def list_vector_stores( - 401: Authentication failed - 403: Authorization failed - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -230,7 +230,7 @@ async def list_vector_stores( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() vector_stores = await client.vector_stores.list() data = [ @@ -250,7 +250,7 @@ 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="Llama Stack", cause=str(e)) + response = ServiceUnavailableResponse(backend_name="OGX", cause=str(e)) raise HTTPException(**response.model_dump()) from e except (LLSApiStatusError, OpenAIAPIStatusError) as e: logger.error("API status error while listing vector stores: %s", e) @@ -281,7 +281,7 @@ async def get_vector_store( - 403: Authorization failed - 404: Vector store not found - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -289,7 +289,7 @@ async def get_vector_store( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() vector_store = await client.vector_stores.retrieve(vector_store_id) return VectorStoreResponse( @@ -304,7 +304,7 @@ async def get_vector_store( ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) @@ -343,7 +343,7 @@ async def update_vector_store( - 403: Authorization failed - 404: Vector store not found - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -351,7 +351,7 @@ async def update_vector_store( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() vector_store = await client.vector_stores.update( vector_store_id, **body.model_dump(exclude_none=True) ) @@ -368,7 +368,7 @@ async def update_vector_store( ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) @@ -404,7 +404,7 @@ async def delete_vector_store( - 401: Authentication failed - 403: Authorization failed - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX Returns: VectorStoreDeleteResponse: Delete outcome for the requested vector store. @@ -415,12 +415,12 @@ async def delete_vector_store( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + 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="Llama Stack", cause=str(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) @@ -454,7 +454,7 @@ async def create_file( # pylint: disable=too-many-branches,too-many-statements - 401: Authentication failed - 403: Authorization failed - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth @@ -485,7 +485,7 @@ async def create_file( # pylint: disable=too-many-branches,too-many-statements raise HTTPException(**response.model_dump()) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Read file content once content = await file.read() @@ -529,7 +529,7 @@ async def create_file( # pylint: disable=too-many-branches,too-many-statements ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) @@ -578,7 +578,7 @@ async def add_file_to_vector_store( # pylint: disable=too-many-locals,too-many- - 403: Authorization failed - 404: Vector store or file not found - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -586,7 +586,7 @@ async def add_file_to_vector_store( # pylint: disable=too-many-locals,too-many- check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Retry logic for database lock errors max_retries = 3 @@ -653,7 +653,7 @@ async def add_file_to_vector_store( # pylint: disable=too-many-locals,too-many- ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) @@ -695,7 +695,7 @@ async def list_vector_store_files( - 403: Authorization failed - 404: Vector store not found - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -703,7 +703,7 @@ async def list_vector_store_files( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() files = await client.vector_stores.files.list(vector_store_id=vector_store_id) data = [ @@ -724,7 +724,7 @@ async def list_vector_store_files( return VectorStoreFilesListResponse(data=data) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) @@ -766,7 +766,7 @@ async def get_vector_store_file( - 403: Authorization failed - 404: File not found in vector store - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX """ _ = auth _ = request @@ -774,7 +774,7 @@ async def get_vector_store_file( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() vs_file = await client.vector_stores.files.retrieve( vector_store_id=vector_store_id, file_id=file_id, @@ -794,7 +794,7 @@ async def get_vector_store_file( ) except APIConnectionError as e: logger.error("Unable to connect to Llama Stack: %s", e) - response = ServiceUnavailableResponse(backend_name="Llama Stack", cause=str(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) @@ -830,7 +830,7 @@ async def delete_vector_store_file( - 401: Authentication failed - 403: Authorization failed - 500: Lightspeed Stack configuration not loaded - - 503: Unable to connect to Llama Stack + - 503: Unable to connect to OGX Returns: VectorStoreFileDeleteResponse: Delete outcome for the requested file. @@ -841,7 +841,7 @@ async def delete_vector_store_file( check_configuration_loaded(configuration) try: - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() await client.vector_stores.files.delete( vector_store_id=vector_store_id, file_id=file_id, @@ -849,7 +849,7 @@ async def delete_vector_store_file( 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="Llama Stack", cause=str(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) diff --git a/src/app/main.py b/src/app/main.py index 0cf61752c..330eb9789 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 llama_stack_client import APIConnectionError, AsyncLlamaStackClient +from ogx_client import APIConnectionError, AsyncOgxClient from starlette.routing import Mount, Route, WebSocketRoute from starlette.types import ASGIApp, Message, Receive, Scope, Send @@ -19,14 +19,13 @@ from app.database import create_tables, initialize_database from app.endpoints.streaming_query import shutdown_background_topic_summary_tasks from authorization.azure_token_manager import AzureEntraIDManager -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from metrics import recording from metrics.utils import setup_model_metrics from models.api.responses.error import InternalServerErrorResponse from sentry import initialize_sentry -from utils.common import register_mcp_servers_async from utils.degraded_mode import DegradedModeTracker from utils.llama_stack_version import check_llama_stack_version @@ -84,8 +83,8 @@ async def lifespan(_app: FastAPI) -> AsyncIterator[None]: initialize_sentry() llama_stack_config = configuration.configuration.llama_stack - await AsyncLlamaStackClientHolder().load(llama_stack_config) - client: AsyncLlamaStackClient = AsyncLlamaStackClientHolder().get_client() + await AsyncOgxClientHolder().load(llama_stack_config) + client: AsyncOgxClient = AsyncOgxClientHolder().get_client() logger.debug("Llama Stack client initialized, trying to connect to Llama Stack") # Check connectivity to Llama Stack and set degraded mode if unavailable degraded_tracker = DegradedModeTracker() @@ -120,10 +119,8 @@ async def lifespan(_app: FastAPI) -> AsyncIterator[None]: azure_entra_id_config = configuration.configuration.azure_entra_id if azure_entra_id_config is not None: AzureEntraIDManager().set_config(azure_entra_id_config) - azure_base_url = await AsyncLlamaStackClientHolder().get_azure_base_url() + azure_base_url = await AsyncOgxClientHolder().get_azure_base_url() AzureEntraIDManager().set_base_url(azure_base_url) - logger.info("Registering MCP servers") - await register_mcp_servers_async(logger, configuration.configuration) # Set up model metrics if in healthy mode if not degraded_tracker.is_degraded(): diff --git a/src/client.py b/src/client.py index 05e6fe250..7fd4a1e5c 100644 --- a/src/client.py +++ b/src/client.py @@ -7,8 +7,8 @@ import yaml from fastapi import HTTPException -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from ogx.core.library_client import AsyncOGXAsLibraryClient +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient import constants from authorization.azure_token_manager import AzureEntraIDManager @@ -23,21 +23,22 @@ from log import get_logger, setup_logging from models.api.responses.error import ServiceUnavailableResponse from models.config import LlamaStackConfiguration +from utils.model_list import parse_model_list_response from utils.types import Singleton logger = get_logger(__name__) -class AsyncLlamaStackClientHolder(metaclass=Singleton): - """Container for an initialised AsyncLlamaStackClient.""" +class AsyncOgxClientHolder(metaclass=Singleton): + """Container for an initialised AsyncOgxClient.""" - _lsc: Optional[AsyncLlamaStackClient] = None + _lsc: Optional[AsyncOgxClient] = None _config_path: Optional[str] = None @property def is_library_client(self) -> bool: """Check if using library mode client.""" - return isinstance(self._lsc, AsyncLlamaStackAsLibraryClient) + return isinstance(self._lsc, AsyncOGXAsLibraryClient) async def load(self, llama_stack_config: LlamaStackConfiguration) -> None: """Initialize the Llama Stack client based on configuration.""" @@ -62,7 +63,7 @@ async def _load_library_client(self, config: LlamaStackConfiguration) -> None: inference.providers (with no config block) correctly falls through to synthesis. Stores the final config path for use in reload. """ - logger.info("Using Llama stack as library client") + logger.info("Using Llama Stack as library client") # Configure logging before synthesis/enrichment so INFO lines from those # steps are not dropped. Without handlers, Python's lastResort only @@ -77,13 +78,13 @@ async def _load_library_client(self, config: LlamaStackConfiguration) -> None: else: self._config_path = self._synthesize_library_config() - client = AsyncLlamaStackAsLibraryClient(self._config_path) + client = AsyncOGXAsLibraryClient(self._config_path) await client.initialize() self._lsc = client # Re-apply logging configuration after ogx's setup_logging() is called. # This ensures the desired logging configuration is applied when - # using AsyncLlamaStackAsLibraryClient. + # using AsyncOGXAsLibraryClient. setup_logging() def _synthesize_library_config(self) -> str: @@ -100,7 +101,7 @@ def _synthesize_library_config(self) -> str: config_file = os.environ.get(constants.CONFIG_PATH_ENV_VAR) if not config_file: raise ValueError( - f"Cannot synthesize Llama Stack config: {constants.CONFIG_PATH_ENV_VAR} " + f"Cannot synthesize OGX config: {constants.CONFIG_PATH_ENV_VAR} " "is not set" ) @@ -119,14 +120,14 @@ def _synthesize_library_config(self) -> str: def _load_service_client(self, config: LlamaStackConfiguration) -> None: """Initialize client in service mode (remote HTTP).""" - logger.info("Using Llama stack running as a service") + logger.info("Using Llama Stack running as a service") logger.info( "Using timeout of %d seconds for Llama Stack requests", config.timeout ) api_key = config.api_key.get_secret_value() if config.api_key else None # Convert AnyHttpUrl to string for the client base_url = str(config.url) if config.url else None - self._lsc = AsyncLlamaStackClient( + self._lsc = AsyncOgxClient( base_url=base_url, api_key=api_key, timeout=config.timeout ) @@ -166,23 +167,23 @@ def _enrich_library_config(self, input_config_path: str) -> str: logger.warning("Failed to write enriched config: %s", e) return input_config_path - def get_client(self) -> AsyncLlamaStackClient: + def get_client(self) -> AsyncOgxClient: """ Get the initialized client held by this holder. Returns: - AsyncLlamaStackClient: The initialized client instance. + AsyncOgxClient: The initialized client instance. Raises: RuntimeError: If the client has not been initialized; call `load(...)` first. """ if not self._lsc: raise RuntimeError( - "AsyncLlamaStackClient has not been initialised. Ensure 'load(..)' has been called." + "AsyncOgxClient has not been initialised. Ensure 'load(..)' has been called." ) return self._lsc - async def reload_library_client(self) -> AsyncLlamaStackClient: + async def reload_library_client(self) -> AsyncOgxClient: """Reload library client to pick up env var changes. For use with library mode only. @@ -193,18 +194,18 @@ async def reload_library_client(self) -> AsyncLlamaStackClient: if not self._config_path: raise RuntimeError("Cannot reload: config path not set") try: - client = AsyncLlamaStackAsLibraryClient(self._config_path) + client = AsyncOGXAsLibraryClient(self._config_path) await client.initialize() except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e self._lsc = client # Re-apply logging configuration after ogx's setup_logging() is called. # This ensures the desired logging configuration is applied when - # using AsyncLlamaStackAsLibraryClient. + # using AsyncOGXAsLibraryClient. setup_logging() return client @@ -234,7 +235,7 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: """ try: client = self.get_client() - models = await client.models.list() + models = parse_model_list_response(await client.models.list()) except RuntimeError as e: logger.warning("Client not initialized, skipping model check: %s", e) return False, f"Client not initialized: {e!s}" @@ -242,7 +243,7 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: logger.error("Error checking model availability: %s", e) return False, f"Error checking model availability: {e!s}" - if any(m.id == model_id for m in models): + if any(m.identifier == model_id for m in models): return True, f"Model {model_id} is available" # Model not found - attempt self-healing reload for library clients. @@ -256,8 +257,8 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: try: await self.reload_library_client() client = self.get_client() - reloaded_models = await client.models.list() - if any(m.id == model_id for m in reloaded_models): + reloaded_models = parse_model_list_response(await client.models.list()) + if any(m.identifier == model_id for m in reloaded_models): logger.info( "Model %s found after client reload", model_id, @@ -271,7 +272,7 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: ) as err: logger.error("Client reload failed: %s", err) - registered_ids = [m.id for m in models] + registered_ids = [m.identifier for m in models] logger.error( "Model %s not found in registry. Registered models: %s", model_id, @@ -279,7 +280,7 @@ async def check_model_available(self, model_id: str) -> tuple[bool, str]: ) return False, f"Model {model_id} not found in model registry" - async def update_azure_token(self) -> AsyncLlamaStackClient: + async def update_azure_token(self) -> AsyncOgxClient: """Apply cached Azure credentials and replace the held client. Returns: @@ -295,17 +296,17 @@ async def update_azure_token(self) -> AsyncLlamaStackClient: return self.get_client() current_provider_data = dict( - cast(AsyncLlamaStackAsLibraryClient, self._lsc).provider_data or {} + cast(AsyncOGXAsLibraryClient, self._lsc).provider_data or {} ) current_provider_data.update(updates) - client = AsyncLlamaStackAsLibraryClient( + client = AsyncOGXAsLibraryClient( self._config_path, provider_data=current_provider_data ) await client.initialize() self._lsc = client # Re-apply logging configuration after ogx's setup_logging() is called. # This ensures the desired logging configuration is applied when - # using AsyncLlamaStackAsLibraryClient. + # using AsyncOGXAsLibraryClient. setup_logging() return client @@ -313,7 +314,7 @@ async def update_azure_token(self) -> AsyncLlamaStackClient: # Service client mode current_client = self.get_client() current_headers = current_client.default_headers or {} - provider_data_json = current_headers.get("X-LlamaStack-Provider-Data") + provider_data_json = current_headers.get("X-OGX-Provider-Data") try: provider_data = json.loads(provider_data_json) if provider_data_json else {} @@ -324,7 +325,7 @@ async def update_azure_token(self) -> AsyncLlamaStackClient: updated_headers = { **current_headers, - "X-LlamaStack-Provider-Data": json.dumps(provider_data), + "X-OGX-Provider-Data": json.dumps(provider_data), } updated_client = current_client.copy( diff --git a/src/configuration.py b/src/configuration.py index e95e89083..87553e668 100644 --- a/src/configuration.py +++ b/src/configuration.py @@ -6,7 +6,7 @@ # We want to support environment variable replacement in the configuration # similarly to how it is done in llama-stack, so we use their function directly -from llama_stack.core.stack import replace_env_vars +from ogx.core.stack import replace_env_vars import constants from cache.cache import Cache @@ -32,6 +32,7 @@ RerankerConfiguration, RlsapiV1Configuration, ServiceConfiguration, + ShieldConfiguration, SkillsConfiguration, SplunkConfiguration, UserDataCollection, @@ -173,10 +174,10 @@ def service_configuration(self) -> ServiceConfiguration: @property def llama_stack_configuration(self) -> LlamaStackConfiguration: - """Return Llama stack configuration. + """Return Llama Stack configuration. Returns: - LlamaStackConfiguration: The configured Llama stack settings. + LlamaStackConfiguration: The configured Llama Stack settings. Raises: LogicError: If the application configuration has not been loaded. @@ -553,6 +554,13 @@ def skills(self) -> Optional[SkillsConfiguration]: raise LogicError("logic error: configuration is not loaded") return self._configuration.skills + @property + def shields(self) -> list[ShieldConfiguration]: + """Return the list of configured guardrail shields.""" + if self._configuration is None: + raise LogicError("logic error: configuration is not loaded") + return self._configuration.shields + @property def rag_id_mapping(self) -> dict[str, str]: """Return mapping from vector_db_id to rag_id from BYOK and OKP RAG config. diff --git a/src/constants.py b/src/constants.py index 91aa3cb83..e32927e79 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] = "0.6.0" +MAXIMAL_SUPPORTED_LLAMA_STACK_VERSION: Final[str] = "1.0.2" # 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/data/default_run.yaml b/src/data/default_run.yaml index 206ee8718..71b38d06f 100644 --- a/src/data/default_run.yaml +++ b/src/data/default_run.yaml @@ -8,8 +8,8 @@ # # It is intentionally thinner than the repo-root run.yaml: it carries only the # APIs and providers needed to boot a minimal, queryable stack (inference, -# safety, vector_io, agents, tool_runtime, files) plus the storage backends -# those providers reference. tool_runtime includes rag-runtime (file_search / +# vector_io, responses, tool_runtime, files) plus the storage backends +# those providers reference. tool_runtime includes file-search (file_search / # RAG) and model-context-protocol (MCP) so those capabilities work when # operators enable them. Operators extend it via high-level sections or # `native_override`. @@ -20,14 +20,14 @@ version: 2 apis: -- agents +- responses +- conversations - files - inference -- safety - tool_runtime - vector_io -image_name: starter +distro_name: starter external_providers_dir: ${env.EXTERNAL_PROVIDERS_DIR:=~/.llama/providers.d} providers: @@ -47,18 +47,13 @@ providers: storage_dir: ${env.SQLITE_STORE_DIR:=~/.llama/storage/files} provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard tool_runtime: - - config: {} - provider_id: rag-runtime - provider_type: inline::rag-runtime - config: {} provider_id: model-context-protocol provider_type: remote::model-context-protocol + - config: {} + provider_id: file-search + provider_type: inline::file-search vector_io: - config: persistence: @@ -66,17 +61,14 @@ providers: backend: kv_default provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default provider_id: meta-reference - provider_type: inline::meta-reference + provider_type: inline::builtin storage: backends: @@ -99,29 +91,18 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: [] - shields: - # NOTE: this is a minimal placeholder shield for the zero-dependency default - # baseline (it mirrors the repo's root run.yaml). provider_shield_id points at - # a chat model, NOT a real Llama Guard checkpoint, so it does not perform real - # safety gating. A real guard model (e.g. meta-llama/Llama-Guard-3-8B) is - # configured per deployment via the high-level schema or native_override, and - # real shield gating is exercised by the e2e suite — deliberately not pinned - # here so the default baseline boots with only an OPENAI_API_KEY. - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime # REQUIRED for file_search tool calls to work. Without it, llama-stack's -# rag-runtime silently fails all file_search operations with no error logged. +# file-search runtime silently fails all file_search operations with no error logged. vector_stores: # LCORE-1498: Disables Llama Stack RAG annotation generation that causes # unwanted citation/file markers in model output. @@ -131,6 +112,3 @@ vector_stores: default_embedding_model: provider_id: sentence-transformers model_id: nomic-ai/nomic-embed-text-v1.5 - -safety: - default_shield_id: llama-guard diff --git a/src/lightspeed_stack.py b/src/lightspeed_stack.py index 4070a335a..31ea4eb81 100644 --- a/src/lightspeed_stack.py +++ b/src/lightspeed_stack.py @@ -186,7 +186,7 @@ def main() -> None: configuration.load_configuration(args.config_file) logger.info("Configuration: %s", configuration.configuration) logger.info( - "Llama stack configuration: %s", configuration.llama_stack_configuration + "Llama Stack configuration: %s", configuration.llama_stack_configuration ) # Deprecation schedule (Decision S2): the legacy two-file path keeps diff --git a/src/llama_stack_configuration.py b/src/llama_stack_configuration.py index 68f344b1a..766f04fb4 100644 --- a/src/llama_stack_configuration.py +++ b/src/llama_stack_configuration.py @@ -26,7 +26,7 @@ from urllib.parse import urljoin import yaml -from llama_stack.core.stack import replace_env_vars +from ogx.core.stack import replace_env_vars from pydantic import SecretStr import constants diff --git a/src/log.py b/src/log.py index 4fc20408a..39e3414ab 100644 --- a/src/log.py +++ b/src/log.py @@ -116,7 +116,7 @@ def build_logging_config() -> dict[t.Any, t.Any]: "level": log_level, "propagate": False, }, - "llama_stack_client": { + "ogx_client": { "handlers": [handler], "level": log_level, "propagate": False, diff --git a/src/metrics/utils.py b/src/metrics/utils.py index afb832d29..62b5fd97e 100644 --- a/src/metrics/utils.py +++ b/src/metrics/utils.py @@ -1,10 +1,11 @@ """Utility functions for metrics handling.""" import metrics -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import configuration from log import get_logger from utils.endpoints import check_configuration_loaded +from utils.model_list import parse_model_list_response logger = get_logger(__name__) @@ -17,13 +18,11 @@ async def setup_model_metrics() -> None: """ logger.info("Setting up model metrics") check_configuration_loaded(configuration) - model_list = await AsyncLlamaStackClientHolder().get_client().models.list() + model_list = parse_model_list_response( + await AsyncOgxClientHolder().get_client().models.list() + ) - models = [ - model - for model in model_list - if model.custom_metadata and model.custom_metadata.get("model_type") == "llm" - ] + models = [model for model in model_list if model.model_type == "llm"] default_model_label = ( configuration.inference.default_provider, @@ -31,12 +30,8 @@ async def setup_model_metrics() -> None: ) for model in models: - provider = ( - str(model.custom_metadata.get("provider_id", "")) - if model.custom_metadata - else "" - ) - model_name = model.id + provider = str(model.provider_id or "") + model_name = model.identifier if provider and model_name: # If the model/provider combination is the default, set the metric value to 1 # Otherwise, set it to 0 diff --git a/src/models/api/requests/query.py b/src/models/api/requests/query.py index e48bc8b4c..8b6e7b83d 100644 --- a/src/models/api/requests/query.py +++ b/src/models/api/requests/query.py @@ -23,7 +23,7 @@ class QueryRequest(BaseModel): generate_topic_summary: Whether to generate topic summary for new conversations. media_type: The optional media type for response format (application/json or text/plain). vector_store_ids: The optional list of specific vector store IDs to query for RAG. - shield_ids: The optional list of safety shield IDs to apply. + shield_ids: The optional list of configured shield names to apply. solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict. """ @@ -105,9 +105,9 @@ class QueryRequest(BaseModel): shield_ids: Optional[list[str]] = Field( None, - description="Optional list of safety shield IDs to apply. " - "If None, all configured shields are used. ", - examples=["llama-guard", "custom-shield"], + description="Optional list of configured shield names to apply. " + "If None, all configured shields are used.", + examples=["topic-guard", "pii-redaction"], ) solr: Optional[SolrVectorSearchRequest] = Field( diff --git a/src/models/api/requests/responses_openai.py b/src/models/api/requests/responses_openai.py index dde597077..809b6963f 100644 --- a/src/models/api/requests/responses_openai.py +++ b/src/models/api/requests/responses_openai.py @@ -3,16 +3,16 @@ import json from typing import Any, Optional, Self -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoice as ToolChoice, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponsePrompt as Prompt, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseReasoning as Reasoning, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseText as Text, ) from pydantic import BaseModel, field_validator, model_validator @@ -57,8 +57,8 @@ class ResponsesRequest(BaseModel): calls, MCP tools). Defaults to all tools available to the model. generate_topic_summary: LCORE-specific flag indicating whether to generate a topic summary for new conversations. Defaults to True. - shield_ids: LCORE-specific list of safety shield IDs to apply. If None, all - configured shields are used. + shield_ids: LCORE-specific list of configured shield names to apply. + If None, all configured shields are used. solr: Optional Solr inline RAG options (mode, filters) or legacy filter-only dict. """ diff --git a/src/models/api/responses/error/service_unavailable.py b/src/models/api/responses/error/service_unavailable.py index 26dde8463..39744ea10 100644 --- a/src/models/api/responses/error/service_unavailable.py +++ b/src/models/api/responses/error/service_unavailable.py @@ -16,9 +16,9 @@ class ServiceUnavailableResponse(AbstractErrorResponse): "json_schema_extra": { "examples": [ { - "label": "llama stack", + "label": "ogx", "detail": { - "response": "Unable to connect to Llama Stack", + "response": "Unable to connect to OGX", "cause": "Connection error while trying to reach backend service.", }, }, diff --git a/src/models/api/responses/successful/catalog.py b/src/models/api/responses/successful/catalog.py index 3d357a724..54ade5d84 100644 --- a/src/models/api/responses/successful/catalog.py +++ b/src/models/api/responses/successful/catalog.py @@ -5,12 +5,14 @@ from pydantic import Field from models.api.responses.successful.bases import AbstractSuccessfulResponse +from models.common import CatalogModel, CatalogShield +from models.common.tools import CatalogTool class ModelsResponse(AbstractSuccessfulResponse): """Model representing a response to models request.""" - models: list[dict[str, Any]] = Field( + models: list[CatalogModel] = Field( ..., description="List of models available", ) @@ -39,7 +41,7 @@ class ModelsResponse(AbstractSuccessfulResponse): class ToolsResponse(AbstractSuccessfulResponse): """Model representing a response to tools request.""" - tools: list[dict[str, Any]] = Field( + tools: list[CatalogTool] = Field( description=( "List of tools available from all configured MCP servers and built-in toolgroups" ), @@ -77,9 +79,9 @@ class ToolsResponse(AbstractSuccessfulResponse): class ShieldsResponse(AbstractSuccessfulResponse): """Model representing a response to shields request.""" - shields: list[dict[str, Any]] = Field( + shields: list[CatalogShield] = Field( ..., - description="List of shields available", + description="List of shields configured in Lightspeed Core Stack", ) model_config = { @@ -88,12 +90,32 @@ class ShieldsResponse(AbstractSuccessfulResponse): { "shields": [ { - "identifier": "lightspeed_question_validity-shield", - "provider_resource_id": "lightspeed_question_validity-shield", - "provider_id": "lightspeed_question_validity", + "name": "question-validity", + "provider_id": "question_validity", "type": "shield", - "params": {}, - } + "config": { + "model_id": "openai/gpt-4o-mini", + "model_prompt": "Is this question valid?", + "invalid_question_response": ( + "I can only answer questions about the product." + ), + }, + }, + { + "name": "pii-redaction", + "provider_id": "redaction", + "type": "shield", + "config": { + "rules": [ + { + "pattern": r"\b\d{3}-\d{2}-\d{4}\b", + "replacement": "[REDACTED]", + "case_sensitive": None, + } + ], + "case_sensitive": False, + }, + }, ], } ] diff --git a/src/models/api/responses/successful/responses_openai.py b/src/models/api/responses/successful/responses_openai.py index 30ed13fb0..769e74a81 100644 --- a/src/models/api/responses/successful/responses_openai.py +++ b/src/models/api/responses/successful/responses_openai.py @@ -2,28 +2,28 @@ from typing import Any, Literal, Optional, cast -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseError as Error, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoice as ToolChoice, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutput as Output, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponsePrompt as Prompt, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseReasoning as Reasoning, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseText as Text, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseTool as OutputTool, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseUsage as Usage, ) diff --git a/src/models/common/__init__.py b/src/models/common/__init__.py index f3f599087..895c345f8 100644 --- a/src/models/common/__init__.py +++ b/src/models/common/__init__.py @@ -12,12 +12,14 @@ ProviderHealthStatus, ) from models.common.mcp import MCPServerAuthInfo, MCPServerInfo +from models.common.models import CatalogModel from models.common.moderation import ( ShieldModerationBlocked, ShieldModerationPassed, ShieldModerationResult, ) from models.common.query import Attachment, SolrVectorSearchRequest +from models.common.shields import CatalogShield from models.common.transcripts import Transcript, TranscriptMetadata from models.common.turn_summary import ( MCPListToolsSummary, @@ -32,6 +34,8 @@ __all__ = [ "Attachment", + "CatalogModel", + "CatalogShield", "ConversationData", "ConversationDetails", "ConversationTurn", diff --git a/src/models/common/models.py b/src/models/common/models.py new file mode 100644 index 000000000..12cea2a52 --- /dev/null +++ b/src/models/common/models.py @@ -0,0 +1,31 @@ +"""Backend-agnostic model catalog types.""" + +from __future__ import annotations + +from typing import Any + +from pydantic import BaseModel, Field + + +class CatalogModel(BaseModel): + """Normalized model entry used by ``/models`` and internal model resolution. + + Unifies OpenAI-style, Anthropic, and Google ``models.list()`` payloads into + one catalog shape. + """ + + identifier: str = Field(description="Model identifier") + metadata: dict[str, Any] = Field( + default_factory=dict, + description="Provider-specific metadata excluding core catalog fields", + ) + api_model_type: str = Field( + description="API model type (typically mirrors model_type)" + ) + provider_id: str = Field(description="Provider identifier") + type: str = Field(default="model", description="Object type, always 'model'") + provider_resource_id: str = Field( + default="", + description="Provider-native resource identifier for the model", + ) + model_type: str = Field(description="Model type such as 'llm' or 'embedding'") diff --git a/src/models/common/moderation.py b/src/models/common/moderation.py index 1e4f16368..575d6e535 100644 --- a/src/models/common/moderation.py +++ b/src/models/common/moderation.py @@ -2,7 +2,7 @@ from typing import Annotated, Literal -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMessage as ResponseMessage, ) from pydantic import BaseModel, Field @@ -20,7 +20,14 @@ class ShieldModerationBlocked(BaseModel): decision: Literal["blocked"] = "blocked" message: str moderation_id: str - refusal_response: ResponseMessage + + @property + def refusal_response(self) -> ResponseMessage: + """Build a ResponseMessage carrying the shield's refusal text.""" + return ResponseMessage( + role="assistant", + content=self.message, + ) ShieldModerationResult = Annotated[ diff --git a/src/models/common/responses/contexts.py b/src/models/common/responses/contexts.py index b2497b58e..042520e2f 100644 --- a/src/models/common/responses/contexts.py +++ b/src/models/common/responses/contexts.py @@ -5,7 +5,7 @@ from typing import Optional from fastapi import BackgroundTasks -from llama_stack_client import AsyncLlamaStackClient +from ogx_client import AsyncOgxClient from pydantic import BaseModel, ConfigDict, Field from models.api.requests import QueryRequest @@ -20,7 +20,7 @@ class ResponsesContext(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) - client: AsyncLlamaStackClient = Field(description="The Llama Stack client") + client: AsyncOgxClient = Field(description="The Llama Stack client") auth: tuple[str, str, bool, str] = Field( description="Authentication tuple (user_id, username, skip_userid_check, token)", ) @@ -103,7 +103,7 @@ class ResponseGeneratorContext: # pylint: disable=too-many-instance-attributes started_at: str # Dependencies & State - client: AsyncLlamaStackClient + client: AsyncOgxClient moderation_result: ShieldModerationResult # RAG index identification diff --git a/src/models/common/responses/responses_api_params.py b/src/models/common/responses/responses_api_params.py index fb9e064b9..1a392fbd5 100644 --- a/src/models/common/responses/responses_api_params.py +++ b/src/models/common/responses/responses_api_params.py @@ -3,19 +3,19 @@ from collections.abc import Mapping from typing import Any, Final, Optional -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoice as ToolChoice, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponsePrompt as Prompt, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseReasoning as Reasoning, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseText as Text, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseToolMCP as OutputToolMCP, ) from pydantic import BaseModel, Field diff --git a/src/models/common/responses/types.py b/src/models/common/responses/types.py index 33e0e91ab..07fe4e74f 100644 --- a/src/models/common/responses/types.py +++ b/src/models/common/responses/types.py @@ -2,41 +2,44 @@ from typing import Annotated, Literal, Optional -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputFunctionToolCallOutput as FunctionToolCallOutput, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolFileSearch as InputToolFileSearch, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolFunction as InputToolFunction, ) -from llama_stack_api.openai_responses import OpenAIResponseInputToolMCP -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import OpenAIResponseInputToolMCP +from ogx_api.openai_responses import ( OpenAIResponseInputToolWebSearch as InputToolWebSearch, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMCPApprovalRequest as McpApprovalRequest, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMCPApprovalResponse as McpApprovalResponse, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMessage as ResponseMessage, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFileSearchToolCall as FileSearchToolCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFunctionToolCall as FunctionToolCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPCall as McpCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPListTools as McpListTools, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( + OpenAIResponseOutputMessageReasoningItem as ReasoningItem, +) +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageWebSearchToolCall as WebSearchToolCall, ) from pydantic import Field @@ -60,7 +63,6 @@ class InputToolMCP(OpenAIResponseInputToolMCP): "file_search_call.results", "message.input_image.image_url", "message.output_text.logprobs", - "reasoning.encrypted_content", ] type ResponseItem = ( @@ -73,6 +75,7 @@ class InputToolMCP(OpenAIResponseInputToolMCP): | McpApprovalRequest | FunctionToolCall | McpApprovalResponse + | ReasoningItem ) type ResponseInput = str | list[ResponseItem] diff --git a/src/models/common/shields.py b/src/models/common/shields.py new file mode 100644 index 000000000..29b88092d --- /dev/null +++ b/src/models/common/shields.py @@ -0,0 +1,30 @@ +"""Catalog models for the ``/shields`` endpoint.""" + +from __future__ import annotations + +from typing import Any, Literal + +from pydantic import BaseModel, Field + + +class CatalogShield(BaseModel): + """Shield entry in the ``/shields`` catalog response. + + Attributes: + name: Unique, user-facing name identifying this shield instance. + provider_id: Shield provider / type discriminator. + type: Catalog entry type; always shield. + config: Type-specific shield configuration. + """ + + name: str = Field(description="Unique, user-facing name of the shield instance") + provider_id: Literal["question_validity", "redaction"] = Field( + description="Shield provider / type discriminator", + ) + type: Literal["shield"] = Field( + default="shield", + description="Catalog entry type; always shield", + ) + config: dict[str, Any] = Field( + description="Type-specific shield configuration", + ) diff --git a/src/models/common/tools.py b/src/models/common/tools.py new file mode 100644 index 000000000..70f7c49e2 --- /dev/null +++ b/src/models/common/tools.py @@ -0,0 +1,37 @@ +"""Backend-agnostic tool listing models.""" + +from __future__ import annotations + +from typing import Any, Optional + +from pydantic import BaseModel + + +class ListedMcpTool(BaseModel): + """Tool metadata returned from an MCP ``tools/list`` call.""" + + name: str + description: Optional[str] = None + input_schema: Optional[dict[str, Any]] = None + + +class CatalogToolParameter(BaseModel): + """Parameter entry for a tool in the ``/tools`` catalog response.""" + + name: str + description: str + parameter_type: str + required: bool = False + default: Optional[Any] = None + + +class CatalogTool(BaseModel): + """Tool entry in the ``/tools`` catalog response.""" + + identifier: str + description: str + parameters: list[CatalogToolParameter] + provider_id: str + toolgroup_id: str + server_source: str + type: str = "tool" diff --git a/src/models/common/turn_summary.py b/src/models/common/turn_summary.py index 2b342b758..37a4a8f47 100644 --- a/src/models/common/turn_summary.py +++ b/src/models/common/turn_summary.py @@ -5,7 +5,7 @@ from typing import Any, Optional -from llama_stack_api import OpenAIResponseOutput +from ogx_api import OpenAIResponseOutput from pydantic import AnyUrl, BaseModel, Field from utils.token_counter import TokenCounter diff --git a/src/models/config.py b/src/models/config.py index 0449a233d..941507ce1 100644 --- a/src/models/config.py +++ b/src/models/config.py @@ -921,11 +921,11 @@ def check_llama_stack_model(self) -> Self: # it means that use_as_library_client attribute must be set to True if self.use_as_library_client is None: raise ValueError( - "Llama stack URL is not specified and library client mode is not specified" + "Llama Stack URL is not specified and library client mode is not specified" ) if self.use_as_library_client is False: raise ValueError( - "Llama stack URL is not specified and library client mode is not enabled" + "Llama Stack URL is not specified and library client mode is not enabled" ) # None -> False conversion @@ -2885,6 +2885,74 @@ def from_environment(cls) -> "ObservabilityConfiguration": return cls(otel=otel_vars) +class QuestionValidityShieldConfiguration(ConfigurationBase): + """Configuration for a named question-validity guardrail shield. + + Attributes: + name: Unique, user-facing name identifying this shield instance. + provider_id: Discriminator identifying this as a question-validity shield. + config: Question-validity-specific configuration. + """ + + name: str = Field( + ..., + title="Shield name", + description="Unique, user-facing name identifying this shield instance.", + ) + + provider_id: Literal["question_validity"] = Field( + ..., + title="Shield provider id", + description="Discriminator identifying this as a question-validity shield.", + ) + + config: QuestionValidityConfig = Field( + ..., + title="Shield configuration", + description="Question-validity-specific configuration for this shield.", + ) + + +class RedactionShieldConfiguration(ConfigurationBase): + """Configuration for a named PII-redaction guardrail shield. + + Attributes: + name: Unique, user-facing name identifying this shield instance. + provider_id: Discriminator identifying this as a redaction shield. + config: Redaction-specific configuration. + """ + + name: str = Field( + ..., + title="Shield name", + description="Unique, user-facing name identifying this shield instance.", + ) + + provider_id: Literal["redaction"] = Field( + ..., + title="Shield provider id", + description="Discriminator identifying this as a redaction shield.", + ) + + config: RedactionConfig = Field( + ..., + title="Shield configuration", + description="Redaction-specific configuration for this shield.", + ) + + +ShieldConfiguration = Annotated[ + QuestionValidityShieldConfiguration | RedactionShieldConfiguration, + Field(discriminator="provider_id"), +] +"""Configuration for a single named guardrail shield (question validity or redaction). + +A discriminated union on ``provider_id``: Pydantic selects +``QuestionValidityShieldConfiguration`` or ``RedactionShieldConfiguration`` +and validates ``config`` against the matching model. +""" + + class Configuration(ConfigurationBase): """Global service configuration.""" @@ -3092,6 +3160,33 @@ class Configuration(ConfigurationBase): "maximum prompts per user, display name length, and content length.", ) + shields: list[ShieldConfiguration] = Field( + default_factory=list, + title="Shields configuration", + description="List of pydantic-ai-lightspeed agent guardrail shields " + "(question validity and PII redaction). Each entry has a unique " + "'name', a 'provider_id' ('question_validity' or 'redaction'), " + "and a type-specific 'config'.", + ) + + @model_validator(mode="after") + def validate_shield_names_unique(self) -> Self: + """Reject shields lists containing duplicate names. + + Returns: + Self: The model instance after validation. + + Raises: + ValueError: If two or more shields share the same name. + """ + names = [shield.name for shield in self.shields] + duplicates = {name for name in names if names.count(name) > 1} + if duplicates: + raise ValueError( + f"Shield names must be unique, found duplicates: {sorted(duplicates)}" + ) + return self + @model_validator(mode="after") def validate_mcp_auth_headers(self) -> Self: """ diff --git a/src/pydantic_ai_lightspeed/capabilities/base.py b/src/pydantic_ai_lightspeed/capabilities/base.py new file mode 100644 index 000000000..25b5dc8b4 --- /dev/null +++ b/src/pydantic_ai_lightspeed/capabilities/base.py @@ -0,0 +1,18 @@ +"""Abstract base for safety capabilities with a standalone run interface.""" + +from abc import abstractmethod + +from pydantic_ai.capabilities import AbstractCapability +from typing_extensions import TypeVar + +from models.common.moderation import ShieldModerationResult + +T = TypeVar("T", default=object) + + +class AbstractSafetyCapability(AbstractCapability[T]): + """Interface for safety/moderation that can be called directly.""" + + @abstractmethod + async def run(self, input_text: str) -> ShieldModerationResult: + """Run moderation on input text.""" diff --git a/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py b/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py index d07a8a764..c8097d1ad 100644 --- a/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py +++ b/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py @@ -13,21 +13,35 @@ from dataclasses import dataclass, field from string import Template from typing import Optional +from uuid import uuid4 from pydantic_ai import AgentRunResult, RunContext from pydantic_ai._agent_graph import GraphAgentState -from pydantic_ai.capabilities import AbstractCapability, WrapRunHandler +from pydantic_ai.capabilities import WrapRunHandler from pydantic_ai.direct import model_request -from pydantic_ai.messages import ModelRequest, TextContent, UserContent +from pydantic_ai.messages import ( + ModelRequest, + ModelResponse, + TextContent, + TextPart, + UserContent, +) from pydantic_ai.models import Model from pydantic_ai.models.openai import OpenAIResponsesModelSettings -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from log import get_logger +from models.common.moderation import ( + ShieldModerationBlocked, + ShieldModerationPassed, + ShieldModerationResult, +) from models.config import ( QuestionValidityConfig, ) -from pydantic_ai_lightspeed.llamastack import LlamaStackResponsesModel +from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability +from pydantic_ai_lightspeed.llamastack import OgxResponsesModel +from utils.conversations import append_turn_to_conversation logger = get_logger(__name__) @@ -55,8 +69,49 @@ def _extract_message_str_from_user_content(user_content: Sequence[UserContent]) return "\n".join(str_arr) +def _message_to_str(message: Optional[str | Sequence[UserContent]]) -> str: + """Convert a user message (string, content sequence, or None) to plain text. + + Parameters: + message: The user input as a string, sequence of user content, or None. + + Returns: + A plain-text representation of the message, or an empty string for None. + """ + match message: + case str() as s: + return s + case Sequence() as seq: + return _extract_message_str_from_user_content(seq) + case None: + return "" + + +def _extract_conversation_id(model: Model) -> Optional[str]: + """Extract the Llama Stack conversation ID from the agent's model settings. + + The main agent's model is built with ``conversation`` in its + ``extra_body`` model settings (see ``OgxResponsesModel.from_ogx_client``). + This pulls it back out so the capability can persist the rejected turn + to the same conversation. + + Parameters: + model: The model bound to the current agent run (``ctx.model``). + + Returns: + The conversation ID, or None if the model has no such setting + (e.g. when used outside a Llama Stack-backed agent). + """ + extra_body = (model.settings or {}).get("extra_body") + if not isinstance(extra_body, dict): + return None + + conversation_id = extra_body.get("conversation") + return conversation_id if isinstance(conversation_id, str) else None + + @dataclass -class QuestionValidity(AbstractCapability[None]): +class QuestionValidity(AbstractSafetyCapability): """Block or modify user input based on a guardrail check. The guard function receives the user prompt and returns True if safe. @@ -76,11 +131,11 @@ class QuestionValidity(AbstractCapability[None]): def __post_init__(self) -> None: """Initialize the model instance from the configured model ID.""" - llama_stack_client = AsyncLlamaStackClientHolder().get_client() + ogx_client = AsyncOgxClientHolder().get_client() - self._model = LlamaStackResponsesModel.from_llama_stack_client( + self._model = OgxResponsesModel.from_ogx_client( self.config.model_id, - llama_stack_client, + ogx_client, model_settings=OpenAIResponsesModelSettings(openai_store=False), ) @@ -93,16 +148,10 @@ def _build_prompt(self, message: Optional[str | Sequence[UserContent]]) -> str: Returns: The rendered prompt string ready to send to the validity model. """ - match message: - case str() as s: - _message = s - case Sequence() as seq: - _message = _extract_message_str_from_user_content(seq) - case None: - _message = "" - return Template(self.config.model_prompt).substitute( - message=_message, allowed=SUBJECT_ALLOWED, rejected=SUBJECT_REJECTED + message=_message_to_str(message), + allowed=SUBJECT_ALLOWED, + rejected=SUBJECT_REJECTED, ) async def wrap_run( @@ -136,7 +185,46 @@ async def wrap_run( return await handler() # proceed with the real run # short-circuit: return the rejection message with shield usage tracked - state = GraphAgentState(usage=ctx.usage) + user_message = _message_to_str(ctx.prompt) + state = GraphAgentState( + usage=ctx.usage, + message_history=[ + ModelRequest.user_text_prompt(user_message), + ModelResponse( + [TextPart(self.config.invalid_question_response)], + finish_reason="stop", + ), + ], + ) + + conversation_id = _extract_conversation_id(ctx.model) + if conversation_id is not None: + await append_turn_to_conversation( + AsyncOgxClientHolder().get_client(), + conversation_id, + user_message, + self.config.invalid_question_response, + ) + else: + logger.warning( + "Unable to determine conversation ID from model settings; " + "skipping v1/conversation persistence for rejected question." + ) + return AgentRunResult( output=self.config.invalid_question_response, _state=state ) + + async def run(self, input_text: str) -> ShieldModerationResult: + """Run question-validity check and return a moderation result.""" + result = await model_request( + model=self._model, messages=[ModelRequest.user_text_prompt(input_text)] + ) + + if result.text is not None and result.text.strip() == SUBJECT_ALLOWED: + return ShieldModerationPassed() + + return ShieldModerationBlocked( + message=self.config.invalid_question_response, + moderation_id=f"modr-{uuid4()}", + ) diff --git a/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py b/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py index 0bdedb174..29c3ef910 100644 --- a/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py +++ b/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py @@ -3,9 +3,9 @@ from collections.abc import Sequence from dataclasses import dataclass, replace from typing import Any, Optional +from uuid import uuid4 from pydantic_ai import RunContext -from pydantic_ai.capabilities import AbstractCapability from pydantic_ai.messages import ( ModelMessage, ModelRequest, @@ -19,7 +19,13 @@ ) from pydantic_ai.models import ModelRequestContext +from models.common.moderation import ( + ShieldModerationBlocked, + ShieldModerationPassed, + ShieldModerationResult, +) from models.config import RedactionConfig +from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability from pydantic_ai_lightspeed.capabilities.redaction.core import ( CompiledPatterns, redact_text, @@ -257,7 +263,7 @@ def _redact_response( @dataclass -class PiiRedactionCapability(AbstractCapability[Any]): +class PiiRedactionCapability(AbstractSafetyCapability): """Pydantic AI capability that redacts PII from agent messages. Applies configurable regex-based redaction rules to user prompt @@ -321,3 +327,14 @@ async def after_model_request( ) return new_response + + async def run(self, input_text: str) -> ShieldModerationResult: + """Run PII redaction on input text and return a moderation result.""" + result = redact_text(input_text, self.config.compiled_patterns) + + if result.redacted: + return ShieldModerationBlocked( + message="Sensitive content detected.", moderation_id=f"modr-{uuid4()}" + ) + + return ShieldModerationPassed() diff --git a/src/pydantic_ai_lightspeed/llamastack/__init__.py b/src/pydantic_ai_lightspeed/llamastack/__init__.py index fac9ee826..ed11a43c2 100644 --- a/src/pydantic_ai_lightspeed/llamastack/__init__.py +++ b/src/pydantic_ai_lightspeed/llamastack/__init__.py @@ -1,6 +1,6 @@ """Pydantic AI provider for Llama Stack.""" -from pydantic_ai_lightspeed.llamastack._model import LlamaStackResponsesModel -from pydantic_ai_lightspeed.llamastack._provider import LlamaStackProvider +from pydantic_ai_lightspeed.llamastack._model import OgxResponsesModel +from pydantic_ai_lightspeed.llamastack._provider import OgxProvider -__all__ = ["LlamaStackProvider", "LlamaStackResponsesModel"] +__all__ = ["OgxProvider", "OgxResponsesModel"] diff --git a/src/pydantic_ai_lightspeed/llamastack/_model.py b/src/pydantic_ai_lightspeed/llamastack/_model.py index 422e7e19d..9a10856ff 100644 --- a/src/pydantic_ai_lightspeed/llamastack/_model.py +++ b/src/pydantic_ai_lightspeed/llamastack/_model.py @@ -10,8 +10,11 @@ deltas must be replayed with the matching suffix so pydantic_ai can append the streamed ``tool_args`` content to the correct part. -This module provides ``LlamaStackResponsesModel`` which wraps the event stream to +This module provides ``OgxResponsesModel`` which wraps the event stream to buffer those early delta events and replay them correctly once the item is announced. + +Additionally overrides ``_responses_create`` to filter out ``reasoning.encrypted_content`` +from the include parameter, which llama-stack / OGX doesn't support. """ from __future__ import annotations as _annotations @@ -21,8 +24,8 @@ from contextlib import asynccontextmanager from typing import Any, Final, Optional, cast -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack_client import AsyncLlamaStackClient +from ogx.core.library_client import AsyncOGXAsLibraryClient +from ogx_client import AsyncOgxClient from openai import AsyncStream from openai.types import responses from pydantic_ai import UnexpectedModelBehavior @@ -45,7 +48,7 @@ from log import get_logger from models.common.responses.responses_api_params import ResponsesApiParams -from pydantic_ai_lightspeed.llamastack._provider import LlamaStackProvider +from pydantic_ai_lightspeed.llamastack._provider import OgxProvider logger = get_logger(__name__) @@ -220,14 +223,56 @@ def _replay_mcp_buffered_deltas( ] -class LlamaStackResponsesModel(OpenAIResponsesModel): +class OgxResponsesModel(OpenAIResponsesModel): """OpenAI Responses model with Llama Stack streaming compatibility fixes. Overrides the streaming response processing to buffer and replay ``ResponseFunctionCallArgumentsDeltaEvent`` events that Llama Stack emits before the corresponding ``McpCall`` or ``ResponseFunctionToolCall`` item. + + Also filters ``reasoning.encrypted_content`` from the include parameter since + OGX doesn't support it. """ + async def _responses_create( + self, + messages: list[ModelMessage], + stream: bool, + model_settings: OpenAIResponsesModelSettings, + model_request_parameters: ModelRequestParameters, + ) -> Any: + """Call parent's ``_responses_create``, filtering encrypted reasoning include. + + OGX doesn't support ``reasoning.encrypted_content`` in the include + parameter. pydantic-ai adds it automatically based on the model profile, so we + disable that profile flag before sending. + + Args: + messages: Model messages for the request. + stream: Whether this is a streaming request. + model_settings: Model-specific settings. + model_request_parameters: Request parameters for the model. + + Returns: + Response from the Responses API. + """ + # Parent gates include on this profile flag; disable it for OGX. + self.profile["openai_supports_encrypted_reasoning_content"] = False + # Branch on stream so mypy matches OpenAIResponsesModel overloads + if stream: + return await super()._responses_create( + messages, + True, + model_settings, + model_request_parameters, + ) + return await super()._responses_create( + messages, + False, + model_settings, + model_request_parameters, + ) + async def request( # pylint: disable=unused-argument self, messages: list[ModelMessage], @@ -356,15 +401,15 @@ async def request_stream( # pylint: disable=unused-argument ) @staticmethod - def from_llama_stack_client( + def from_ogx_client( model_name: str, - client: AsyncLlamaStackClient | AsyncLlamaStackAsLibraryClient, + client: AsyncOgxClient | AsyncOGXAsLibraryClient, *, responses_params: Optional[ResponsesApiParams] = None, model_settings: Optional[ModelSettings] = None, profile: Optional[ModelProfileSpec] = None, - ) -> LlamaStackResponsesModel: - """Create a ``LlamaStackResponsesModel`` from a Llama Stack client. + ) -> OgxResponsesModel: + """Create a ``OgxResponsesModel`` from a Llama Stack client. Mirrors ``OpenAIResponsesModel.__init__`` parameters, but accepts a Llama Stack client instead of a provider. Exactly one of @@ -385,9 +430,9 @@ def from_llama_stack_client( are provided. Returns: - Configured ``LlamaStackResponsesModel`` instance. + Configured ``OgxResponsesModel`` instance. """ - provider = LlamaStackProvider.from_llama_stack_client(client) + provider = OgxProvider.from_ogx_client(client) if responses_params is not None and model_settings is not None: raise ValueError( @@ -401,6 +446,6 @@ def from_llama_stack_client( elif model_settings is not None: _settings = model_settings - return LlamaStackResponsesModel( + return OgxResponsesModel( model_name, provider=provider, profile=profile, settings=_settings ) diff --git a/src/pydantic_ai_lightspeed/llamastack/_provider.py b/src/pydantic_ai_lightspeed/llamastack/_provider.py index e710da8ec..3ead58d6a 100644 --- a/src/pydantic_ai_lightspeed/llamastack/_provider.py +++ b/src/pydantic_ai_lightspeed/llamastack/_provider.py @@ -5,9 +5,9 @@ from typing import TYPE_CHECKING, Optional import httpx -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack.core.request_headers import parse_request_provider_data -from llama_stack_client import AsyncLlamaStackClient +from ogx.core.library_client import AsyncOGXAsLibraryClient +from ogx.core.request_headers import parse_request_provider_data +from ogx_client import AsyncOgxClient from openai import AsyncOpenAI from pydantic_ai import ModelProfile from pydantic_ai.models import create_async_http_client @@ -15,25 +15,25 @@ from pydantic_ai.providers import Provider from pydantic_ai_lightspeed.llamastack._transport import ( - LlamaStackLibraryTransport, + OgxLibraryTransport, wrap_http_client_with_provider_data, ) if TYPE_CHECKING: - from llama_stack.core.library_client import ( # pylint: disable=reimported - AsyncLlamaStackAsLibraryClient, + from ogx.core.library_client import ( # pylint: disable=reimported + AsyncOGXAsLibraryClient, ) DEFAULT_BASE_URL = "http://localhost:8321/v1" -class LlamaStackProvider(Provider[AsyncOpenAI]): +class OgxProvider(Provider[AsyncOpenAI]): """Provider for Llama Stack — connects to a Llama Stack server's OpenAI-compatible API. Supports two modes: 1. **Server mode** — connect to a running Llama Stack server via HTTP - 2. **Library mode** — run Llama Stack in-process via ``AsyncLlamaStackAsLibraryClient`` + 2. **Library mode** — run Llama Stack in-process via ``AsyncOGXAsLibraryClient`` """ @property @@ -57,23 +57,23 @@ def model_profile(model_name: str) -> Optional[ModelProfile]: return openai_model_profile(model_name) @staticmethod - def from_llama_stack_client( - client: AsyncLlamaStackClient | AsyncLlamaStackAsLibraryClient, - ) -> LlamaStackProvider: - """Create a ``LlamaStackProvider`` from a Llama Stack client. + def from_ogx_client( + client: AsyncOgxClient | AsyncOGXAsLibraryClient, + ) -> OgxProvider: + """Create a ``OgxProvider`` from a Llama Stack client. - For an ``AsyncLlamaStackAsLibraryClient``, delegates to library mode. - For an ``AsyncLlamaStackClient``, extracts the base URL, API key, and + For an ``AsyncOGXAsLibraryClient``, delegates to library mode. + For an ``AsyncOgxClient``, extracts the base URL, API key, and underlying HTTP client to create a server-mode provider. Args: client: A Llama Stack client (server or library variant). Returns: - Configured ``LlamaStackProvider`` instance. + Configured ``OgxProvider`` instance. """ - if isinstance(client, AsyncLlamaStackAsLibraryClient): - return LlamaStackProvider(library_client=client) + if isinstance(client, AsyncOGXAsLibraryClient): + return OgxProvider(library_client=client) api_key = client.api_key or "not-needed" base = str(client.base_url).rstrip("/") base_url = base if base.endswith("/v1") else f"{base}/v1" @@ -86,7 +86,7 @@ def from_llama_stack_client( provider_data = parse_request_provider_data(default_headers) http_client = client._client # pylint: disable=protected-access http_client = wrap_http_client_with_provider_data(http_client, provider_data) - return LlamaStackProvider( + return OgxProvider( base_url=base_url, api_key=api_key, http_client=http_client, @@ -97,7 +97,7 @@ def __init__( *, base_url: Optional[str] = None, api_key: Optional[str] = None, - library_client: Optional[AsyncLlamaStackAsLibraryClient] = None, + library_client: Optional[AsyncOGXAsLibraryClient] = None, http_client: Optional[httpx.AsyncClient] = None, ) -> None: """Create a new Llama Stack provider. @@ -109,7 +109,7 @@ def __init__( api_key: The API key for authentication. Defaults to ``'not-needed'`` since local Llama Stack servers typically don't require one. Must be ``None`` when ``library_client`` is provided. - library_client: An initialized ``AsyncLlamaStackAsLibraryClient`` for library mode. + library_client: An initialized ``AsyncOGXAsLibraryClient`` for library mode. When provided, requests are dispatched in-process (no server needed). Mutually exclusive with ``base_url``, ``api_key``, and ``http_client``. http_client: An existing ``httpx.AsyncClient`` to use for making HTTP requests. @@ -126,7 +126,7 @@ def __init__( ) self._library_client = library_client - transport = LlamaStackLibraryTransport(library_client) + transport = OgxLibraryTransport(library_client) lib_http_client = httpx.AsyncClient( transport=transport, base_url="http://llama-stack-library", @@ -149,7 +149,7 @@ def __init__( def __repr__(self) -> str: """Return a string representation of the provider.""" - return f"LlamaStackProvider(name={self.name!r}, base_url={self.base_url!r})" + return f"OgxProvider(name={self.name!r}, base_url={self.base_url!r})" def _set_http_client(self, http_client: httpx.AsyncClient) -> None: """Inject an httpx.AsyncClient into the underlying OpenAI client. diff --git a/src/pydantic_ai_lightspeed/llamastack/_transport.py b/src/pydantic_ai_lightspeed/llamastack/_transport.py index 4d2a179b2..d78ec27e0 100644 --- a/src/pydantic_ai_lightspeed/llamastack/_transport.py +++ b/src/pydantic_ai_lightspeed/llamastack/_transport.py @@ -7,20 +7,20 @@ from typing import Any, Optional import httpx -from llama_stack.core.library_client import ( - AsyncLlamaStackAsLibraryClient, +from ogx.core.library_client import ( + AsyncOGXAsLibraryClient, convert_pydantic_to_json_value, ) -from llama_stack.core.request_headers import ( +from ogx.core.request_headers import ( PROVIDER_DATA_VAR, request_provider_data_context, ) -from llama_stack.core.server.routes import find_matching_route -from llama_stack.core.utils.context import preserve_contexts_async_generator +from ogx.core.server.routes import find_matching_route +from ogx.core.utils.context import preserve_contexts_async_generator from starlette.responses import StreamingResponse _PROVIDER_DATA_HEADER_KEYS = ( - "X-LlamaStack-Provider-Data", + "X-OGX-Provider-Data", "x-llamastack-provider-data", ) @@ -46,7 +46,7 @@ def inject_provider_data_into_headers( headers: Mapping[str, str], provider_data: Optional[Mapping[str, Any]], ) -> dict[str, str]: - """Add ``X-LlamaStack-Provider-Data`` when provider data is configured. + """Add ``X-OGX-Provider-Data`` when provider data is configured. Args: headers: Existing request headers. @@ -60,7 +60,7 @@ def inject_provider_data_into_headers( if any(key in headers for key in _PROVIDER_DATA_HEADER_KEYS): return dict(headers) result = dict(headers) - result["X-LlamaStack-Provider-Data"] = json.dumps(provider_data) + result["X-OGX-Provider-Data"] = json.dumps(provider_data) return result @@ -109,7 +109,7 @@ def wrap_http_client_with_provider_data( if not provider_data: return http_client - transport = LlamaStackServerTransport( + transport = OgxServerTransport( http_client._transport, # pylint: disable=protected-access provider_data=provider_data, ) @@ -120,7 +120,7 @@ def wrap_http_client_with_provider_data( ) -class LlamaStackServerTransport(httpx.AsyncBaseTransport): +class OgxServerTransport(httpx.AsyncBaseTransport): """httpx transport that injects provider data headers before delegating over HTTP.""" def __init__( @@ -176,7 +176,7 @@ async def __aiter__(self) -> AsyncIterator[bytes]: yield chunk -class LlamaStackLibraryTransport(httpx.AsyncBaseTransport): +class OgxLibraryTransport(httpx.AsyncBaseTransport): """Custom httpx transport that dispatches requests through a Llama Stack library client. Instead of making real HTTP calls, this transport routes requests directly @@ -184,11 +184,11 @@ class LlamaStackLibraryTransport(httpx.AsyncBaseTransport): route matching and body conversion logic. """ - def __init__(self, client: AsyncLlamaStackAsLibraryClient) -> None: + def __init__(self, client: AsyncOGXAsLibraryClient) -> None: """Initialize the transport with a Llama Stack library client. Args: - client: An initialized ``AsyncLlamaStackAsLibraryClient`` whose route + client: An initialized ``AsyncOGXAsLibraryClient`` whose route handlers will receive dispatched requests. """ self._client = client diff --git a/src/utils/agents/error_handler.py b/src/utils/agents/error_handler.py new file mode 100644 index 000000000..aeeddec0c --- /dev/null +++ b/src/utils/agents/error_handler.py @@ -0,0 +1,103 @@ +"""Error mapping for agent inference failures to structured API error responses.""" + +from typing import TypeAlias + +from ogx_client import APIConnectionError, APIStatusError +from pydantic_ai.exceptions import ( + AgentRunError, + ContentFilterError, + IncompleteToolCall, + ModelAPIError, + ModelHTTPError, + UnexpectedModelBehavior, + UsageLimitExceeded, +) + +from log import get_logger +from models.api.responses.error import ( + AbstractErrorResponse, + InternalServerErrorResponse, + PromptTooLongResponse, + QuotaExceededResponse, + ServiceUnavailableResponse, +) +from utils.query import ( + handle_known_apistatus_errors, + is_context_length_error, +) + +AgentInferenceError: TypeAlias = ( + AgentRunError | APIStatusError | APIConnectionError | RuntimeError +) + +logger = get_logger(__name__) + + +def map_agent_inference_error( + exc: AgentInferenceError, + model_id: str, +) -> AbstractErrorResponse: + """Map agent run failures from pydantic-ai or Llama Stack to an LCS error response. + + Args: + exc: Agent, HTTP status, connection, or context-length runtime error. + model_id: Model identifier in provider/model format. + + Returns: + Structured error response for HTTP or SSE error events. + + Raises: + RuntimeError: Re-raised when ``exc`` is a non-agent ``RuntimeError`` that is + not a recognized context-length failure. + """ + match exc: + case AgentRunError() as agent_exc: + return map_pydantic_agent_run_error(agent_exc, model_id) + case APIStatusError() as status_exc: + return handle_known_apistatus_errors(status_exc, model_id) + case APIConnectionError() as connection_exc: + return ServiceUnavailableResponse( + backend_name="OGX", + cause=str(connection_exc), + ) + case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)): + return PromptTooLongResponse(model=model_id) + case _: + return InternalServerErrorResponse.generic() + + +def map_pydantic_agent_run_error( # pylint: disable=too-many-return-statements + exc: AgentRunError, model_id: str +) -> AbstractErrorResponse: + """Map pydantic-ai ``AgentRunError`` subclasses to LCS error responses. + + Args: + exc: Agent exception to map. + model_id: Model identifier in provider/model format. + + Returns: + Structured error response for HTTP or SSE error events. + """ + match exc: + case ContentFilterError() as filter_exc: + return InternalServerErrorResponse.query_failed(str(filter_exc)) + case IncompleteToolCall(): + return PromptTooLongResponse(model=model_id) + case UnexpectedModelBehavior(): + logger.error("Unexpected model behavior: %s", exc, exc_info=True) + return InternalServerErrorResponse.generic() + case UsageLimitExceeded(): + return QuotaExceededResponse.model(model_id) + case ModelHTTPError() as http_exc if is_context_length_error(str(http_exc)): + return PromptTooLongResponse(model=model_id) + case ModelHTTPError(status_code=429): + return QuotaExceededResponse.model(model_id) + case ModelHTTPError(): + return InternalServerErrorResponse.generic() + case ModelAPIError() as api_exc: + return ServiceUnavailableResponse( + backend_name="OGX", + cause=str(api_exc), + ) + case _: + return InternalServerErrorResponse.query_failed(str(exc)) diff --git a/src/utils/agents/query.py b/src/utils/agents/query.py index 189b1b989..044cf67cc 100644 --- a/src/utils/agents/query.py +++ b/src/utils/agents/query.py @@ -6,15 +6,9 @@ from typing import Optional, TypeAlias, cast from fastapi import HTTPException -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient from pydantic_ai.exceptions import ( AgentRunError, - ContentFilterError, - IncompleteToolCall, - ModelAPIError, - ModelHTTPError, - UnexpectedModelBehavior, - UsageLimitExceeded, ) from pydantic_ai.messages import ModelRequest, ModelResponse, ToolReturnPart from pydantic_ai.run import AgentRunResult @@ -27,8 +21,6 @@ AbstractErrorResponse, InternalServerErrorResponse, PromptTooLongResponse, - QuotaExceededResponse, - ServiceUnavailableResponse, ) from models.common.agents import AgentTurnAccumulator from models.common.moderation import ShieldModerationResult @@ -36,6 +28,7 @@ from models.common.responses.responses_api_params import ResponsesApiParams from models.common.responses.types import ResponseInput from models.common.turn_summary import TurnSummary +from utils.agents.error_handler import map_agent_inference_error from utils.agents.tool_processor import ( process_function_tool_call, process_function_tool_result, @@ -47,8 +40,6 @@ from utils.query import ( build_multimodal_input, extract_provider_and_model_from_model_id, - handle_known_apistatus_errors, - is_context_length_error, ) from utils.responses import extract_vector_store_ids_from_tools from utils.token_counter import TokenCounter @@ -70,76 +61,6 @@ class AgentFinishReason(str, Enum): ERROR = "error" -def map_agent_inference_error( - exc: AgentInferenceError, - model_id: str, -) -> AbstractErrorResponse: - """Map agent run failures from pydantic-ai or Llama Stack to an LCS error response. - - Args: - exc: Agent, HTTP status, connection, or context-length runtime error. - model_id: Model identifier in provider/model format. - - Returns: - Structured error response for HTTP or SSE error events. - - Raises: - RuntimeError: Re-raised when ``exc`` is a non-agent ``RuntimeError`` that is - not a recognized context-length failure. - """ - match exc: - case AgentRunError() as agent_exc: - return map_pydantic_agent_run_error(agent_exc, model_id) - case APIStatusError() as status_exc: - return handle_known_apistatus_errors(status_exc, model_id) - case APIConnectionError() as connection_exc: - return ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(connection_exc), - ) - case RuntimeError() as runtime_exc if is_context_length_error(str(runtime_exc)): - return PromptTooLongResponse(model=model_id) - case _: - return InternalServerErrorResponse.generic() - - -def map_pydantic_agent_run_error( - exc: AgentRunError, model_id: str -) -> AbstractErrorResponse: - """Map pydantic-ai ``AgentRunError`` subclasses to LCS error responses. - - Args: - exc: Agent exception to map. - model_id: Model identifier in provider/model format. - - Returns: - Structured error response for HTTP or SSE error events. - """ - match exc: - case ContentFilterError() as filter_exc: - return InternalServerErrorResponse.query_failed(str(filter_exc)) - case IncompleteToolCall(): - return PromptTooLongResponse(model=model_id) - case UnexpectedModelBehavior(): - logger.error("Unexpected model behavior: %s", exc, exc_info=True) - return InternalServerErrorResponse.generic() - case UsageLimitExceeded(): - return QuotaExceededResponse.model(model_id) - case ModelHTTPError() as http_exc if is_context_length_error(str(http_exc)): - return PromptTooLongResponse(model=model_id) - case ModelHTTPError(status_code=429): - return QuotaExceededResponse.model(model_id) - case ModelHTTPError(): - return InternalServerErrorResponse.generic() - case ModelAPIError() as api_exc: - return ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(api_exc), - ) - case _: - return InternalServerErrorResponse.query_failed(str(exc)) - - def get_agent_finish_reason(response: ModelResponse) -> AgentFinishReason: """Get the finish reason from a completed agent model response. @@ -283,18 +204,17 @@ def build_turn_summary_from_agent_run( async def retrieve_agent_response( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, responses_params: ResponsesApiParams, moderation_result: ShieldModerationResult, endpoint_path: str, _original_input: Optional[ResponseInput] = None, no_tools: bool = False, image_attachments: Optional[list[Attachment]] = None, + shield_ids: Optional[list[str]] = None, ) -> TurnSummary: """Retrieve a turn summary from a blocking agent run. - Mirrors :func:`app.endpoints.query.retrieve_response` for the agent path. - Args: client: Llama Stack client for conversation persistence on moderation block. responses_params: Prepared Responses API parameters. @@ -303,6 +223,8 @@ async def retrieve_agent_response( _original_input: Original user input before the explicit-input rewrite. no_tools: Whether to skip tool processing. image_attachments: Image attachments for multimodal prompt construction. + shield_ids: Optional list of shield names to run for this turn, mirroring + ``QueryRequest.shield_ids``. If ``None``, all configured shields run. Returns: Turn summary for the completed agent run. @@ -322,7 +244,11 @@ async def retrieve_agent_response( ) try: agent = build_agent( - client, responses_params, configuration.skills, no_tools=no_tools + client, + responses_params, + configuration, + shields=shield_ids, + no_tools=no_tools, ) logger.debug("Starting agent non-streaming response processing") if image_attachments: diff --git a/src/utils/agents/streaming.py b/src/utils/agents/streaming.py index da6a21c54..e03b2326b 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 llama_stack_client import APIConnectionError, APIStatusError +from ogx_client import APIConnectionError, APIStatusError from pydantic_ai import Agent, AgentRunError, AgentRunResultEvent, ToolReturnPart from pydantic_ai.messages import ( AgentStreamEvent, @@ -46,12 +46,12 @@ from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams from models.common.turn_summary import TurnSummary +from utils.agents.error_handler import map_agent_inference_error from utils.agents.query import ( AgentFinishReason, extract_agent_token_usage, get_agent_finish_reason, get_finish_reason_error, - map_agent_inference_error, ) from utils.agents.tool_processor import ( process_function_tool_call, @@ -130,7 +130,11 @@ async def retrieve_agent_response_generator( ) agent = build_agent( - context.client, responses_params, configuration.skills, no_tools=no_tools + context.client, + responses_params, + configuration, + shields=context.query_request.shield_ids, + no_tools=no_tools, ) return ( diff --git a/src/utils/builtin_tools.py b/src/utils/builtin_tools.py new file mode 100644 index 000000000..ef9da1f99 --- /dev/null +++ b/src/utils/builtin_tools.py @@ -0,0 +1,117 @@ +"""Discover builtin file-search tools when that provider is configured.""" + +from __future__ import annotations + +from typing import Final + +from fastapi import HTTPException +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client.types.shared.provider_info import ProviderInfo + +from log import get_logger +from models.api.responses.error import ServiceUnavailableResponse +from models.common.tools import CatalogTool, CatalogToolParameter + +logger = get_logger(__name__) + +TOOL_RUNTIME_API: Final[str] = "tool_runtime" +FILE_SEARCH_PROVIDER_TYPE: Final[str] = "inline::file-search" +FILE_SEARCH_PROVIDER_ID: Final[str] = "file-search" +FILE_SEARCH_TOOLGROUP_ID: Final[str] = "builtin::file_search" +BUILTIN_SERVER_SOURCE: Final[str] = "builtin" + +# OGX server-mode GET /v1/admin/tools is broken (admin deps / nested routers), +# so expose the known builtin::file_search catalog when the provider is present. +FILE_SEARCH_CATALOG_TOOLS: Final[list[CatalogTool]] = [ + CatalogTool( + identifier="insert_into_memory", + description="Insert documents into memory", + parameters=[], + provider_id=FILE_SEARCH_PROVIDER_ID, + toolgroup_id=FILE_SEARCH_TOOLGROUP_ID, + server_source=BUILTIN_SERVER_SOURCE, + type="tool", + ), + CatalogTool( + identifier="file_search", + description="Search files for relevant information", + parameters=[ + CatalogToolParameter( + name="query", + description=( + "The query to search for. Can be a natural language " + "sentence or keywords." + ), + parameter_type="string", + required=True, + default=None, + ) + ], + provider_id=FILE_SEARCH_PROVIDER_ID, + toolgroup_id=FILE_SEARCH_TOOLGROUP_ID, + server_source=BUILTIN_SERVER_SOURCE, + type="tool", + ), +] + + +def _is_file_search_provider(provider: ProviderInfo) -> bool: + """Return whether a provider entry is a file-search tool runtime. + + Parameters: + provider: Provider metadata from the backend. + + Returns: + True when the provider implements file-search tool runtime. + """ + return ( + provider.api == TOOL_RUNTIME_API + and provider.provider_type == FILE_SEARCH_PROVIDER_TYPE + ) + + +async def get_file_search_tools( + client: AsyncOgxClient, +) -> list[CatalogTool]: + """Return builtin file-search tools when that provider is configured. + + Provider presence is checked via ``providers.list()``. Tool definitions are + not fetched from ``/v1/admin/tools`` (broken in OGX server mode); the + known ``builtin::file_search`` catalog is returned instead. + + Parameters: + client: Initialized OGX client. + + Returns: + Catalog tools for the configured file-search runtime, or an empty + list when file search is not configured. + """ + 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: + logger.error("Unable to connect to OGX: %s", e) + response = ServiceUnavailableResponse( + backend_name="OGX", cause=str(e) + ).model_dump() + raise HTTPException(**response) from e + + file_search_provider = next( + (provider for provider in providers if _is_file_search_provider(provider)), + None, + ) + if file_search_provider is None: + logger.debug( + "No %s provider configured", + FILE_SEARCH_PROVIDER_TYPE, + ) + return [] + + logger.debug( + "Using static catalog for %d file-search tools (toolgroup %s)", + len(FILE_SEARCH_CATALOG_TOOLS), + FILE_SEARCH_TOOLGROUP_ID, + ) + return list(FILE_SEARCH_CATALOG_TOOLS) diff --git a/src/utils/common.py b/src/utils/common.py index ef07cb6df..e558e0e6f 100644 --- a/src/utils/common.py +++ b/src/utils/common.py @@ -3,105 +3,7 @@ import asyncio from collections.abc import Callable from functools import wraps -from logging import Logger -from typing import Any, cast - -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack_client import AsyncLlamaStackClient - -from client import AsyncLlamaStackClientHolder -from models.config import Configuration, ModelContextProtocolServer - - -async def register_mcp_servers_async( - logger: Logger, configuration: Configuration -) -> None: - """Register Model Context Protocol (MCP) servers with the LlamaStack client (async). - - If no MCP servers are present in the provided configuration this function returns immediately. - Selects between a library client (initializes it) and a service client based on - configuration.llama_stack.use_as_library_client, then registers any MCP servers not already - present in the client's toolgroups. - - Parameters: - ---------- - logger: Logger instance. - configuration: Configuration containing the `mcp_servers` list and - `llama_stack` client mode. - - Notes: - ----- - - The `logger` parameter is used for debug/info logging and is - intentionally undocumented as a common service. - - Exceptions from the LlamaStack client (network/errors during - initialization or registration) are not caught here and will - propagate to the caller. - """ - # Skip MCP registration if no MCP servers are configured - if not configuration.mcp_servers: - logger.debug("No MCP servers configured, skipping registration") - return - - if configuration.llama_stack.use_as_library_client: - # Library client - use async interface - client = cast( - AsyncLlamaStackAsLibraryClient, AsyncLlamaStackClientHolder().get_client() - ) - await client.initialize() - await _register_mcp_toolgroups_async(client, configuration.mcp_servers, logger) - else: - # Service client - also use async interface - client = AsyncLlamaStackClientHolder().get_client() - await _register_mcp_toolgroups_async(client, configuration.mcp_servers, logger) - - -async def _register_mcp_toolgroups_async( - client: AsyncLlamaStackClient, - mcp_servers: list[ModelContextProtocolServer], - logger: Logger, -) -> None: - """ - Register MCP (Model Context Protocol) toolgroups with a LlamaStack async client. - - Checks the client's existing toolgroups and registers any servers from `mcp_servers` - whose `name` is not present in the client's `provider_resource_id` list. For each - new server it calls the client's toolgroups.register with parameters: - `toolgroup_id`=`mcp.name`, `provider_id`=`mcp.provider_id`, and - `mcp_endpoint` containing the server `url`. - - This function performs network calls against the provided async client and does not - catch exceptions raised by those calls — any exceptions from the client (e.g., RPC - or HTTP errors) will propagate to the caller. - - Parameters: - ---------- - client (AsyncLlamaStackClient): The LlamaStack async client used to - query and register toolgroups. - mcp_servers (List[ModelContextProtocolServer]): MCP server descriptors - to ensure are registered. - logger (Logger): Logger used for debug messages about registration - progress. - """ - # Get registered tools - registered_toolgroups = await client.toolgroups.list() - registered_toolgroups_ids = [ - tool_group.provider_resource_id for tool_group in registered_toolgroups - ] - logger.debug("Registered toolgroups: %s", registered_toolgroups_ids) - - # Register toolgroups for MCP servers if not already registered - for mcp in mcp_servers: - if mcp.name not in registered_toolgroups_ids: - logger.debug("Registering MCP server: %s, %s", mcp.name, mcp.url) - - registration_params = { - "toolgroup_id": mcp.name, - "provider_id": mcp.provider_id, - "mcp_endpoint": {"uri": mcp.url}, - } - - await client.toolgroups.register(**registration_params) - logger.debug("MCP server %s registered successfully", mcp.name) +from typing import Any def run_once_async(func: Callable[..., Any]) -> Callable[..., Any]: diff --git a/src/utils/compaction.py b/src/utils/compaction.py index db2822378..12ee6b8d2 100644 --- a/src/utils/compaction.py +++ b/src/utils/compaction.py @@ -30,7 +30,7 @@ from datetime import UTC, datetime from typing import Any -from llama_stack_client import AsyncLlamaStackClient +from ogx_client import AsyncOgxClient from log import get_logger from models.compaction import ConversationSummary @@ -198,7 +198,7 @@ def _extract_response_text(response: Any) -> str: async def summarize_chunk( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, model: str, old_items: list[Any], summarized_through_turn: int, @@ -310,7 +310,7 @@ async def summarize_chunk( async def recursively_resummarize( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, model: str, summaries: list[ConversationSummary], encoding_name: str, diff --git a/src/utils/conversation_compaction.py b/src/utils/conversation_compaction.py index d71f57f7d..e45cefda4 100644 --- a/src/utils/conversation_compaction.py +++ b/src/utils/conversation_compaction.py @@ -47,9 +47,9 @@ from dataclasses import dataclass from typing import Any, Optional, cast -from llama_stack_api.openai_responses import OpenAIResponseMessage -from llama_stack_client import AsyncLlamaStackClient -from llama_stack_client.types.conversations.item_create_params import Item +from ogx_api.openai_responses import OpenAIResponseMessage +from ogx_client import AsyncOgxClient +from ogx_client.types.conversations.item_create_params import Item from cache.cache import Cache from cache.cache_error import CacheError @@ -271,7 +271,7 @@ def _build_explicit_input( async def _write_summary_marker( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, conversation_id: str, summary_text: str, ) -> None: @@ -421,7 +421,7 @@ def _estimate_total_tokens( async def _persist_new_summary_chunk( # pylint: disable=too-many-arguments,too-many-positional-arguments - client: AsyncLlamaStackClient, + client: AsyncOgxClient, conversation_id: str, summary: ConversationSummary, cache: Optional[Cache], @@ -434,7 +434,7 @@ async def _persist_new_summary_chunk( # pylint: disable=too-many-arguments,too- async def _maybe_persist_fold( # pylint: disable=too-many-arguments,too-many-positional-arguments - client: AsyncLlamaStackClient, + client: AsyncOgxClient, model: str, conversation_id: str, cache: Optional[Cache], @@ -491,7 +491,7 @@ def _compacted_result( async def apply_compaction( # pylint: disable=too-many-arguments,too-many-positional-arguments,too-many-locals - client: AsyncLlamaStackClient, + client: AsyncOgxClient, params: ResponsesApiParams, inference_config: InferenceConfiguration, compaction_config: CompactionConfiguration, @@ -616,7 +616,7 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit async def apply_compaction_blocking( # pylint: disable=too-many-arguments,too-many-positional-arguments - client: AsyncLlamaStackClient, + client: AsyncOgxClient, params: ResponsesApiParams, inference_config: InferenceConfiguration, compaction_config: CompactionConfiguration, @@ -651,7 +651,7 @@ async def apply_compaction_blocking( # pylint: disable=too-many-arguments,too-m async def needs_compaction_path( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, params: ResponsesApiParams, inference_config: InferenceConfiguration, compaction_config: CompactionConfiguration, @@ -693,7 +693,7 @@ async def needs_compaction_path( async def store_compacted_turn( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, conversation_id: str, original_input: ResponseInput, output_items: Sequence[Any], diff --git a/src/utils/conversations.py b/src/utils/conversations.py index ac2659688..9c98e1ed6 100644 --- a/src/utils/conversations.py +++ b/src/utils/conversations.py @@ -6,37 +6,40 @@ from typing import Any, Optional, cast from fastapi import HTTPException -from llama_stack_api import OpenAIResponseMessage, OpenAIResponseOutput -from llama_stack_api.openai_responses import ( +from ogx_api import OpenAIResponseMessage, OpenAIResponseOutput +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFileSearchToolCall as FileSearchCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFunctionToolCall as FunctionCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPCall as MCPCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPListTools as MCPListTools, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageWebSearchToolCall as WebSearchCall, ) -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient -from llama_stack_client.types.conversations.item_create_params import Item -from llama_stack_client.types.conversations.item_list_response import ( +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client.types.conversations.item_create_params import Item +from ogx_client.types.conversations.item_list_response import ( ItemListResponse, ) -from llama_stack_client.types.conversations.item_list_response import ( +from ogx_client.types.conversations.item_list_response import ( OpenAIResponseInputFunctionToolCallOutput as FunctionToolCallOutput, ) -from llama_stack_client.types.conversations.item_list_response import ( +from ogx_client.types.conversations.item_list_response import ( + OpenAIResponseInputFunctionToolCallOutputOutputListOpenAIResponseInputMessageContentTextOpenAIResponseInputMessageContentImageOpenAIResponseInputMessageContentFile as FunctionCallOutputContentPart, # pylint: disable=line-too-long +) +from ogx_client.types.conversations.item_list_response import ( OpenAIResponseMcpApprovalRequest as MCPApprovalRequest, ) -from llama_stack_client.types.conversations.item_list_response import ( +from ogx_client.types.conversations.item_list_response import ( OpenAIResponseMcpApprovalResponse as MCPApprovalResponse, ) -from llama_stack_client.types.conversations.item_list_response import ( +from ogx_client.types.conversations.item_list_response import ( OpenAIResponseMessageOutput as MessageOutput, ) @@ -89,6 +92,29 @@ def _extract_text_from_content(content: str | list[Any]) -> str: return "".join(text_fragments) +def _function_call_output_to_str( + output: str | list[FunctionCallOutputContentPart], +) -> str: + """Convert function call output content into a string summary. + + Parameters: + output: Raw function call output from the Conversations API. + + Returns: + Plain string content for ``ToolResultSummary``. + """ + if isinstance(output, str): + return output + + fragments: list[str] = [] + for part in output: + if part.type == "input_text": + fragments.append(part.text) + else: + fragments.append(part.model_dump_json(exclude_none=True)) + return "\n\n".join(fragments) + + def _parse_message_item(item: MessageOutput) -> Message: """Parse a message item into a Message object. @@ -259,7 +285,7 @@ def _build_tool_call_summary_from_item( # pylint: disable=too-many-return-state ToolResultSummary( id=function_output.call_id, status=function_output.status or "success", - content=function_output.output, + content=_function_call_output_to_str(function_output.output), type="function_call_output", round=1, ), @@ -451,7 +477,7 @@ def build_conversation_turns_from_items( async def append_turn_items_to_conversation( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, conversation_id: str, user_input: ResponseInput, llm_output: Sequence[OpenAIResponseOutput], @@ -484,7 +510,7 @@ async def append_turn_items_to_conversation( ) except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e @@ -494,7 +520,7 @@ async def append_turn_items_to_conversation( async def get_all_conversation_items( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, conversation_id_llama_stack: str, ) -> list[ItemListResponse]: """Fetch all items for a conversation (Conversations API), paginating as needed. @@ -520,7 +546,45 @@ async def get_all_conversation_items( return items except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + except APIStatusError as e: + error_response = InternalServerErrorResponse.generic() + raise HTTPException(**error_response.model_dump()) from e + + +async def append_turn_to_conversation( + client: AsyncOgxClient, + conversation_id: str, + user_message: str, + assistant_message: str, +) -> None: + """ + Append a user/assistant turn to a conversation. + + Used to record a conversation turn when a shield blocks the request, + storing both the user's original message and the violation response. + + Parameters: + ---------- + client: The Llama Stack client. + conversation_id: The Llama Stack conversation ID. + user_message: The user's input message. + assistant_message: The shield violation response message. + """ + try: + await client.conversations.items.create( + conversation_id, + items=[ + {"type": "message", "role": "user", "content": user_message}, + {"type": "message", "role": "assistant", "content": assistant_message}, + ], + ) + except APIConnectionError as e: + error_response = ServiceUnavailableResponse( + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e diff --git a/src/utils/endpoints.py b/src/utils/endpoints.py index 5ea928b51..737bdaa24 100644 --- a/src/utils/endpoints.py +++ b/src/utils/endpoints.py @@ -8,7 +8,7 @@ import constants from app.database import get_session -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import AppConfig, LogicError from log import get_logger from models.api.responses.error import ( @@ -218,7 +218,7 @@ async def resolve_response_context( HTTPException: 404 if previous_response_id is set but the turn does not exist; other HTTP exceptions from validate_and_retrieve_conversation. """ - client = AsyncLlamaStackClientHolder().get_client() + client = AsyncOgxClientHolder().get_client() # Context for the LLM passed by conversation if conversation_id: logger.info("Conversation ID specified in request: %s", conversation_id) diff --git a/src/utils/llama_stack_version.py b/src/utils/llama_stack_version.py index 7075a94ec..fb8598178 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 llama_stack_client import APIConnectionError, AsyncLlamaStackClient +from ogx_client import APIConnectionError, AsyncOgxClient from semver import Version from constants import ( @@ -23,7 +23,7 @@ class InvalidLlamaStackVersionException(Exception): async def check_llama_stack_version( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, max_retries: int = DEFAULT_MAX_RETRIES, retry_delay: int = DEFAULT_RETRY_DELAY, ) -> Optional[str]: diff --git a/src/utils/mcp_tools.py b/src/utils/mcp_tools.py new file mode 100644 index 000000000..0e575eb09 --- /dev/null +++ b/src/utils/mcp_tools.py @@ -0,0 +1,157 @@ +"""Utilities for discovering tools from remote MCP servers without Llama Stack.""" + +from __future__ import annotations + +from typing import Any, Optional + +import httpx +from mcp import ClientSession, McpError +from mcp.client.sse import sse_client +from mcp.client.streamable_http import streamable_http_client + +from log import get_logger +from models.common.tools import ListedMcpTool + +logger = get_logger(__name__) + + +# Match MCP SDK defaults previously provided by create_mcp_http_client. +_MCP_HTTP_TIMEOUT = httpx.Timeout(30.0, read=300.0) + +_TRANSPORT_ERRORS = ( + httpx.HTTPStatusError, + httpx.ConnectError, + httpx.TimeoutException, + httpx.RequestError, + McpError, +) + +# Streamable HTTP wraps transport failures in ExceptionGroup (BaseException subclass). +_LIST_MCP_ERRORS = (*_TRANSPORT_ERRORS, ExceptionGroup) + + +def _transport_failure(exc: BaseException) -> Optional[BaseException]: + """Return a transport failure if ``exc`` (or a nested ExceptionGroup) is one.""" + if isinstance(exc, _TRANSPORT_ERRORS): + return exc + if isinstance(exc, ExceptionGroup): + for nested in exc.exceptions: + found = _transport_failure(nested) + if found is not None: + return found + return None + + +async def _list_tools_from_session( + read_stream: Any, + write_stream: Any, +) -> list[ListedMcpTool]: + """Run ``tools/list`` on an initialized MCP client session.""" + async with ClientSession( + read_stream=read_stream, + write_stream=write_stream, + ) as session: + await session.initialize() + tools_result = await session.list_tools() + return [ + ListedMcpTool( + name=tool.name, + description=tool.description, + input_schema=tool.inputSchema, + ) + for tool in tools_result.tools + ] + + +async def _list_via_streamable_http( + endpoint: str, + headers: dict[str, str], +) -> list[ListedMcpTool]: + """List tools using the streamable HTTP MCP transport.""" + async with httpx.AsyncClient( + headers=headers, + timeout=_MCP_HTTP_TIMEOUT, + follow_redirects=True, + ) as http_client: + async with streamable_http_client( + endpoint, + http_client=http_client, + ) as (read_stream, write_stream, _): + return await _list_tools_from_session(read_stream, write_stream) + + +async def _list_via_sse( + endpoint: str, + headers: dict[str, str], +) -> list[ListedMcpTool]: + """List tools using the SSE MCP transport.""" + async with sse_client(endpoint, headers=headers) as (read_stream, write_stream): + return await _list_tools_from_session(read_stream, write_stream) + + +# Prefer streamable HTTP; fall back to SSE for servers that only support it. +_MCP_TRANSPORTS = ( + ("streamable HTTP", _list_via_streamable_http), + ("SSE", _list_via_sse), +) + + +def _prepare_mcp_request_headers(headers: dict[str, str]) -> dict[str, str]: + """Normalize headers for a direct MCP HTTP call. + + File-based secrets are stored as raw tokens. MCP servers expect + ``Authorization: Bearer ``. Query/Responses keep the raw value in + ``build_mcp_headers`` and hand it to Llama Stack separately; only this + direct client path needs the Bearer scheme. + """ + prepared = dict(headers) + for header_name, value in list(prepared.items()): + if ( + header_name.lower() == "authorization" + and value + and not value.startswith("Bearer ") + ): + prepared[header_name] = f"Bearer {value}" + return prepared + + +async def list_mcp_tools( + endpoint: str, + headers: dict[str, str], +) -> list[ListedMcpTool]: + """List tools exposed by a remote MCP server. + + Tries streamable HTTP first, then SSE. Discovery failures are logged and + result in an empty list so callers can skip unavailable servers. + + Parameters: + endpoint: MCP server URL. + headers: Headers to forward (already resolved by the caller). + + Returns: + Tool definitions discovered from the MCP server, or an empty list when + the server is unavailable or returns an error. + """ + request_headers = _prepare_mcp_request_headers(headers) + for index, (transport_name, list_via) in enumerate(_MCP_TRANSPORTS): + try: + return await list_via(endpoint, request_headers) + except _LIST_MCP_ERRORS as exc: + transport_exc = _transport_failure(exc) + if transport_exc is None: + raise + if index < len(_MCP_TRANSPORTS) - 1: + logger.warning( + "Failed to list tools from %s via %s, trying next transport: %s", + endpoint, + transport_name, + transport_exc, + ) + continue + logger.warning( + "Skipping MCP server at %s: unable to list tools via any transport: %s", + endpoint, + transport_exc, + ) + + return [] diff --git a/src/utils/model_list.py b/src/utils/model_list.py new file mode 100644 index 000000000..5de76b4d0 --- /dev/null +++ b/src/utils/model_list.py @@ -0,0 +1,120 @@ +"""Helpers for normalizing OGX ``models.list()`` union responses.""" + +from typing import Any + +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model +from ogx_client.types.model_list_response import ( + AnthropicListModelsResponse, + AnthropicListModelsResponseData, + GoogleListModelsResponse, + GoogleListModelsResponseModel, + ModelListResponse, +) + +from models.common.models import CatalogModel + + +def parse_openai_style_model(model: Model) -> CatalogModel: + """ + Parse an OpenAI-style OGX ``Model`` into a unified catalog model. + + Uses the OGX ``Model`` properties for identifier, model_type, provider + fields, and filtered metadata. + + Parameters: + model: Model object from ``ListModelsResponse.data``. + + Returns: + CatalogModel: Normalized catalog entry. + """ + model_type = model.model_type or "unknown" + + return CatalogModel( + identifier=model.identifier, + metadata=model.metadata or {}, + api_model_type=model_type, + provider_id=model.provider_id or "", + type=model.object or "model", + provider_resource_id=model.provider_resource_id or "", + model_type=model_type, + ) + + +def parse_anthropic_model(model: AnthropicListModelsResponseData) -> CatalogModel: + """Parse an Anthropic model list entry into a unified catalog model. + + Parameters: + model: Anthropic model object from ``AnthropicListModelsResponse.data``. + + Returns: + CatalogModel: Normalized catalog entry. Treated as an LLM. + """ + metadata: dict[str, Any] = { + "display_name": model.display_name, + "created_at": model.created_at, + } + if model.max_input_tokens is not None: + metadata["max_input_tokens"] = model.max_input_tokens + if model.max_tokens is not None: + metadata["max_tokens"] = model.max_tokens + + return CatalogModel( + identifier=model.id, + metadata=metadata, + api_model_type="llm", + provider_id="anthropic", + type=model.type or "model", + provider_resource_id=model.id, + model_type="llm", + ) + + +def parse_google_model(model: GoogleListModelsResponseModel) -> CatalogModel: + """Parse a Google model list entry into a unified catalog model. + + Parameters: + model: Google model object from ``GoogleListModelsResponse.models``. + + Returns: + CatalogModel: Normalized catalog entry. Treated as an LLM. + """ + metadata: dict[str, Any] = { + "display_name": model.display_name, + } + if model.description is not None: + metadata["description"] = model.description + + return CatalogModel( + identifier=model.name, + metadata=metadata, + api_model_type="llm", + provider_id="google", + type="model", + provider_resource_id=model.name, + model_type="llm", + ) + + +def parse_model_list_response(response: ModelListResponse) -> list[CatalogModel]: + """Normalize an OGX ``models.list()`` union response into catalog models. + + OGX returns one of ``ListModelsResponse``, ``AnthropicListModelsResponse``, + or ``GoogleListModelsResponse``. This helper matches on the concrete type + and parses every entry into :class:`CatalogModel`. + + Parameters: + response: The union response returned by ``client.models.list()``. + + Returns: + list[CatalogModel]: Parsed models in the unified catalog shape. + """ + match response: + case ListModelsResponse(data=data): + return [parse_openai_style_model(model) for model in data] + case AnthropicListModelsResponse(data=data): + return [parse_anthropic_model(model) for model in data] + case GoogleListModelsResponse(models=models): + return [parse_google_model(model) for model in models] + case _: + return [] diff --git a/src/utils/models_dumper.py b/src/utils/models_dumper.py index c0a98b32f..a620a6a96 100644 --- a/src/utils/models_dumper.py +++ b/src/utils/models_dumper.py @@ -243,4 +243,118 @@ def dump_models_group(model_group: str, filename: Optional[str] = None) -> None: filename = f"{model_group}.json" # dump all selected models into one OpenAPI-compatible JSON file + # add all requests data models + for model in [ + r.ConversationUpdateRequest, + r.FeedbackRequest, + r.FeedbackStatusUpdateRequest, + r.MCPServerRegistrationRequest, + r.ModelFilter, + r.PromptCreateRequest, + r.PromptUpdateRequest, + r.QueryRequest, + r.ResponsesRequest, + r.RlsapiV1Attachment, + r.RlsapiV1CLA, + r.RlsapiV1Context, + r.RlsapiV1InferRequest, + r.RlsapiV1SystemInfo, + r.RlsapiV1Terminal, + r.StreamingInterruptRequest, + r.VectorStoreCreateRequest, + r.VectorStoreFileCreateRequest, + r.VectorStoreUpdateRequest, + s.AuthorizedResponse, + s.ConfigurationResponse, + s.ConversationDeleteResponse, + s.ConversationResponse, + s.ConversationUpdateResponse, + s.ConversationsListResponse, + s.ConversationsListResponseV2, + s.FeedbackResponse, + s.FeedbackStatusUpdateResponse, + s.FileResponse, + s.InfoResponse, + s.LivenessResponse, + s.MCPClientAuthOptionsResponse, + s.MCPServerDeleteResponse, + s.MCPServerListResponse, + s.MCPServerRegistrationResponse, + s.ModelsResponse, + s.PromptDeleteResponse, + s.PromptResourceResponse, + s.PromptsListResponse, + s.ProviderResponse, + s.ProvidersListResponse, + s.QueryResponse, + s.RAGInfoResponse, + s.RAGListResponse, + s.ReadinessResponse, + s.ResponsesResponse, + s.RlsapiV1InferData, + s.RlsapiV1InferResponse, + s.ShieldsResponse, + s.StatusResponse, + s.StreamingInterruptResponse, + s.StreamingQueryResponse, + s.ToolsResponse, + s.VectorStoreDeleteResponse, + s.VectorStoreFileDeleteResponse, + s.VectorStoreFileResponse, + s.VectorStoreFilesListResponse, + s.VectorStoreResponse, + s.VectorStoresListResponse, + e.AbstractErrorResponse, + e.BadRequestResponse, + e.ConflictResponse, + e.DetailModel, + e.FileTooLargeResponse, + e.ForbiddenResponse, + e.InternalServerErrorResponse, + e.NotFoundResponse, + e.PromptTooLongResponse, + e.QuotaExceededResponse, + e.ServiceUnavailableResponse, + e.UnauthorizedResponse, + e.UnprocessableEntityResponse, + c.Attachment, + c.ConversationData, + c.ConversationDetails, + c.ConversationTurn, + c.MCPListToolsSummary, + c.MCPServerAuthInfo, + c.MCPServerInfo, + c.Message, + c.ProviderHealthStatus, + c.RAGChunk, + c.RAGContext, + c.ReferencedDocument, + c.CatalogShield, + c.ShieldModerationBlocked, + c.ShieldModerationPassed, + c.SolrVectorSearchRequest, + c.ToolCallSummary, + c.ToolInfoSummary, + c.ToolResultSummary, + c.Transcript, + c.TranscriptMetadata, + c.TurnSummary, + a.EndEventData, + a.EndStreamPayload, + a.ErrorEventData, + a.ErrorStreamPayload, + a.InterruptedEventData, + a.InterruptedStreamPayload, + a.StartEventData, + a.StartStreamPayload, + a.StreamPayloadBase, + a.TokenChunkData, + a.TokenStreamPayload, + a.ToolCallStreamPayload, + a.ToolResultStreamPayload, + a.TurnCompleteStreamPayload, + cr.InputToolMCP, + cr.ResponsesApiParams, + ]: + models.append(model) dump_openapi_schema(models, filename) diff --git a/src/utils/pydantic_ai_helpers.py b/src/utils/pydantic_ai_helpers.py index 38bf71e92..49e2b9736 100644 --- a/src/utils/pydantic_ai_helpers.py +++ b/src/utils/pydantic_ai_helpers.py @@ -5,17 +5,27 @@ import re from typing import Any, Final, Optional -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack_client import AsyncLlamaStackClient +from ogx.core.library_client import AsyncOGXAsLibraryClient +from ogx_client import AsyncOgxClient from pydantic_ai.agent import Agent from pydantic_ai.capabilities import AbstractCapability, AgentCapability from pydantic_ai_skills import SkillsCapability +from configuration import AppConfig from models.common.responses.responses_api_params import ResponsesApiParams -from models.config import SkillsConfiguration +from models.common.tools import CatalogTool, CatalogToolParameter +from models.config import ( + QuestionValidityConfig, + RedactionConfig, + ShieldConfiguration, + SkillsConfiguration, +) +from pydantic_ai_lightspeed.capabilities import QuestionValidity +from pydantic_ai_lightspeed.capabilities.redaction import PiiRedactionCapability from pydantic_ai_lightspeed.llamastack import ( - LlamaStackResponsesModel, + OgxResponsesModel, ) +from utils.shields import get_shields_for_request _AGENT_SKILLS_PROVIDER_ID: Final[str] = "agent-skills" _AGENT_SKILLS_TOOLGROUP_ID: Final[str] = "builtin::agent-skills" @@ -44,13 +54,13 @@ def _skills_capability( def _json_schema_to_parameters( schema: Optional[dict[str, Any]], -) -> list[dict[str, Any]]: +) -> list[CatalogToolParameter]: """Convert a JSON Schema object to the flat parameter list used by ``/tools``.""" if not schema or "properties" not in schema: return [] required_params = set(schema.get("required", [])) - parameters: list[dict[str, Any]] = [] + parameters: list[CatalogToolParameter] = [] for name, prop in schema["properties"].items(): parameter_type = prop.get("type") if parameter_type is None and "anyOf" in prop: @@ -62,13 +72,13 @@ def _json_schema_to_parameters( parameter_type = option["type"] break parameters.append( - { - "name": name, - "description": prop.get("description", ""), - "parameter_type": parameter_type or "string", - "required": name in required_params, - "default": prop.get("default"), - } + CatalogToolParameter( + name=name, + description=prop.get("description", ""), + parameter_type=parameter_type or "string", + required=name in required_params, + default=prop.get("default"), + ) ) return parameters @@ -80,44 +90,42 @@ def _capability_tool_description(description: str) -> str: return description.strip() -def _capability_tools_from_toolset(toolset: Any) -> list[dict[str, Any]]: +def _capability_tools_from_toolset(toolset: Any) -> list[CatalogTool]: """Serialize tools registered on a pydantic-ai capability toolset.""" raw_tools = getattr(toolset, "tools", None) if not raw_tools: return [] - tool_dicts: list[dict[str, Any]] = [] + tools: list[CatalogTool] = [] for tool in raw_tools.values(): - tool_dicts.append( - { - "identifier": tool.name, - "description": _capability_tool_description(tool.description or ""), - "parameters": _json_schema_to_parameters( - tool.function_schema.json_schema - ), - "provider_id": _AGENT_SKILLS_PROVIDER_ID, - "toolgroup_id": _AGENT_SKILLS_TOOLGROUP_ID, - "server_source": _BUILTIN_CAPABILITY_SERVER_SOURCE, - "type": _CAPABILITY_TOOL_TYPE, - } + tools.append( + CatalogTool( + identifier=tool.name, + description=_capability_tool_description(tool.description or ""), + parameters=_json_schema_to_parameters(tool.function_schema.json_schema), + provider_id=_AGENT_SKILLS_PROVIDER_ID, + toolgroup_id=_AGENT_SKILLS_TOOLGROUP_ID, + server_source=_BUILTIN_CAPABILITY_SERVER_SOURCE, + type=_CAPABILITY_TOOL_TYPE, + ) ) - return tool_dicts + return tools def get_agent_capability_tools( skills: Optional[SkillsConfiguration], -) -> list[dict[str, Any]]: +) -> list[CatalogTool]: """Return tool metadata for pydantic-ai capabilities configured for LCS agents. Parameters: skills: Agent skills configuration from LCS, or None when skills are disabled. Returns: - Tool dictionaries compatible with the ``/tools`` endpoint response format. + Catalog tools for the ``/tools`` endpoint response format. """ capabilities = _agent_capabilities(skills) or [] - tools: list[dict[str, Any]] = [] + tools: list[CatalogTool] = [] for capability in capabilities: if not isinstance(capability, AbstractCapability): continue @@ -128,20 +136,51 @@ def get_agent_capability_tools( return tools +def _shield_capability(shield: ShieldConfiguration) -> AgentCapability[object]: + """Build the pydantic-ai capability instance for a single configured shield. + + Parameters: + shield: A single guardrail shield configuration entry. + + Returns: + A ``QuestionValidity`` capability when ``shield.provider_id`` is + ``"question_validity"``, or a ``PiiRedactionCapability`` when it is + ``"redaction"``. + + Raises: + ValueError: If ``shield.config`` doesn't match a known shield config type. + """ + match shield.config: + case QuestionValidityConfig(): + return QuestionValidity(config=shield.config) + case RedactionConfig(): + return PiiRedactionCapability(config=shield.config) + case _: + raise ValueError( + f"Unsupported shield config type for shield '{shield.name}': " + f"{type(shield.config).__name__}" + ) + + def _agent_capabilities( skills: Optional[SkillsConfiguration], + shields: Optional[list[ShieldConfiguration]] = None, no_tools: bool = False, ) -> Optional[list[AgentCapability[object]]]: """Assemble pydantic-ai capabilities for an LCS agent. Args: skills: Agent skills configuration from LCS, or None when skills are disabled. + shields: Configured guardrail shields (question validity, redaction), or + None/empty when no shields are enabled. no_tools: When True, omit capabilities that expose a toolset via ``get_toolset()``. Returns: Configured capabilities, or None when no capabilities are enabled. """ capabilities: list[AgentCapability[object]] = [] + for shield in shields or []: + capabilities.append(_shield_capability(shield)) if skills_capability := _skills_capability(skills): capabilities.append(skills_capability) if no_tools: @@ -157,31 +196,38 @@ def _agent_capabilities( def build_agent( - client: AsyncLlamaStackClient | AsyncLlamaStackAsLibraryClient, + client: AsyncOgxClient | AsyncOGXAsLibraryClient, responses_params: ResponsesApiParams, - skills: Optional[SkillsConfiguration], + config: AppConfig, + shields: Optional[list[str]] = None, no_tools: bool = False, ) -> Agent[None, str]: """Build a Pydantic AI agent that mirrors ``responses_params`` on the Llama Stack backend. - Uses ``LlamaStackProvider`` with the same ``AsyncLlamaStackClient`` (or library client) + Uses ``OgxProvider`` with the same ``AsyncOgxClient`` (or library client) as the query endpoint, and ``OpenAIResponsesModel`` so requests follow the Responses API. Llama-Stack-specific fields (conversation, tools, MCP headers, etc.) are passed via ``model_settings['extra_body']`` so they merge into the OpenAI client request body. Parameters: - client: Initialized Llama Stack client from ``AsyncLlamaStackClientHolder().get_client()``. + client: Initialized Llama Stack client from ``AsyncOgxClientHolder().get_client()``. responses_params: Parameters produced by ``prepare_responses_params`` for this turn. - skills: Agent skills configuration from LCS, or None when skills are disabled. + config: Application configuration. Agent skills (``config.skills``) and the + configured guardrail shields (``config.shields``) are extracted from it. + shields: Optional list of shield names to run for this turn, matching each + shield's configured ``name``. Mirrors ``QueryRequest.shield_ids``: if + ``None``, all shields configured in ``config.shields`` run; an empty + list disables all shields. no_tools: When True, omit capabilities that expose a toolset via ``get_toolset()``. Returns: ``Agent`` configured for ``await agent.run(...)`` (or streaming) against the same stack configuration as ``client.responses.create(**responses_params.model_dump())``. """ - capabilities = _agent_capabilities(skills, no_tools=no_tools) + shield_configs = get_shields_for_request(config.shields, shields) + capabilities = _agent_capabilities(config.skills, shield_configs, no_tools=no_tools) - model = LlamaStackResponsesModel.from_llama_stack_client( + model = OgxResponsesModel.from_ogx_client( responses_params.model, client, responses_params=responses_params ) diff --git a/src/utils/query.py b/src/utils/query.py index fe84c184c..c7cbab6ad 100644 --- a/src/utils/query.py +++ b/src/utils/query.py @@ -6,10 +6,9 @@ import psycopg2 from fastapi import HTTPException -from llama_stack_client import ( +from ogx_client import ( APIStatusError as LLSApiStatusError, ) -from llama_stack_client.types import Shield from openai._exceptions import APIStatusError as OpenAIAPIStatusError from pydantic_ai.messages import ImageUrl, UserContent from sqlalchemy import func @@ -121,52 +120,6 @@ def validate_model_provider_override( raise HTTPException(**response.model_dump()) -def _is_inout_shield(shield: Shield) -> bool: - """ - Determine if the shield identifier indicates an input/output shield. - - Parameters: - ---------- - shield (Shield): The shield to check. - - Returns: - ------- - bool: True if the shield identifier starts with "inout_", otherwise False. - """ - return shield.identifier.startswith("inout_") - - -def is_output_shield(shield: Shield) -> bool: - """ - Determine if the shield is for monitoring output. - - Return True if the given shield is classified as an output or - inout shield. - - A shield is considered an output shield if its identifier - starts with "output_" or "inout_". - """ - return _is_inout_shield(shield) or shield.identifier.startswith("output_") - - -def is_input_shield(shield: Shield) -> bool: - """ - Determine if the shield is for monitoring input. - - Return True if the shield is classified as an input or inout - shield. - - Parameters: - ---------- - shield (Shield): The shield identifier to classify. - - Returns: - ------- - bool: True if the shield is for input or both input/output monitoring; False otherwise. - """ - return _is_inout_shield(shield) or not is_output_shield(shield) - - def prepare_input( query_request: QueryRequest, inline_rag_context: Optional[str] = None ) -> str: diff --git a/src/utils/responses.py b/src/utils/responses.py index 2beff8990..9ea71f247 100644 --- a/src/utils/responses.py +++ b/src/utils/responses.py @@ -7,78 +7,78 @@ from typing import Any, Optional, cast from fastapi import HTTPException -from llama_stack_api import OpenAIResponseObject -from llama_stack_api.openai_responses import ApprovalFilter -from llama_stack_api.openai_responses import ( +from ogx_api import OpenAIResponseObject +from ogx_api.openai_responses import ApprovalFilter +from ogx_api.openai_responses import ( OpenAIResponseContentPartRefusal as ContentPartRefusal, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputMessageContent as InputMessageContent, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputMessageContentFile as InputFilePart, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputMessageContentText as InputTextPart, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoice as ToolChoice, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoiceAllowedTools as AllowedTools, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoiceMode as ToolChoiceMode, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolFileSearch as InputToolFileSearch, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMCPApprovalRequest as MCPApprovalRequest, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMCPApprovalResponse as MCPApprovalResponse, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMessage as ResponseMessage, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseObject as ResponseObject, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutput as ResponseOutput, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageContent as OutputMessageContent, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageContentOutputText as OutputTextPart, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFileSearchToolCall as FileSearchCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFunctionToolCall as FunctionCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPCall as MCPCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPListTools as MCPListTools, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageWebSearchToolCall as WebSearchCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseUsage as ResponseUsage, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseUsageInputTokensDetails as UsageInputTokensDetails, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseUsageOutputTokensDetails as UsageOutputTokensDetails, ) -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient import constants from configuration import configuration @@ -113,6 +113,7 @@ build_mcp_headers, find_unresolved_auth_headers, ) +from utils.model_list import parse_model_list_response from utils.prompts import get_system_prompt, get_topic_summary_system_prompt from utils.query import ( extract_provider_and_model_from_model_id, @@ -126,14 +127,52 @@ logger = get_logger(__name__) +async def get_vector_store_ids( + client: AsyncOgxClient, + vector_store_ids: Optional[list[str]] = None, +) -> list[str]: + """Get vector store IDs for querying. + + If vector_store_ids are provided, returns them. Otherwise fetches all + available vector stores from Llama Stack. + + Args: + client: The AsyncOgxClient to use for fetching stores + vector_store_ids: Optional list of vector store IDs. If provided, + returns this list. If None, fetches all available vector stores. + + Returns: + List of vector store IDs to query + + Raises: + HTTPException: With ServiceUnavailableResponse if connection fails, + or InternalServerErrorResponse if API returns an error status + """ + if vector_store_ids is not None: + return vector_store_ids + + try: + vector_stores = await client.vector_stores.list() + return [vector_store.id for vector_store in vector_stores.data] + 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: + error_response = InternalServerErrorResponse.generic() + raise HTTPException(**error_response.model_dump()) from e + + async def get_topic_summary( # pylint: disable=too-many-nested-blocks - question: str, client: AsyncLlamaStackClient, model_id: str + question: str, client: AsyncOgxClient, model_id: str ) -> str: """Get a topic summary for a question using Responses API. Args: question: The question to generate a topic summary for - client: The AsyncLlamaStackClient to use for the request + client: The AsyncOgxClient to use for the request model_id: The llama stack model ID (full format: provider/model) Returns: @@ -155,7 +194,7 @@ async def get_topic_summary( # pylint: disable=too-many-nested-blocks ) except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e @@ -169,7 +208,7 @@ async def get_topic_summary( # pylint: disable=too-many-nested-blocks async def maybe_get_topic_summary( generate_topic_summary: bool, input_text: str, - client: AsyncLlamaStackClient, + client: AsyncOgxClient, model_id: str, ) -> Optional[str]: """Generate a topic summary when requested for the current response. @@ -282,7 +321,7 @@ def _build_provider_data_headers( async def prepare_responses_params( # pylint: disable=too-many-arguments,too-many-locals,too-many-positional-arguments - client: AsyncLlamaStackClient, + client: AsyncOgxClient, query_request: QueryRequest, user_conversation: Optional[UserConversation], token: str, @@ -295,7 +334,7 @@ async def prepare_responses_params( # pylint: disable=too-many-arguments,too-ma """Prepare API request parameters for Responses API. Args: - client: The AsyncLlamaStackClient instance (must be initialized by caller) + client: The AsyncOgxClient instance (must be initialized by caller) query_request: The query request containing the user's question user_conversation: The user conversation if conversation_id was provided, None otherwise token: The authentication token for authorization @@ -352,7 +391,7 @@ async def prepare_responses_params( # pylint: disable=too-many-arguments,too-ma conversation = await client.conversations.create(metadata={}) except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e @@ -1274,13 +1313,13 @@ def parse_arguments_string(arguments_str: str) -> dict[str, Any]: async def check_model_configured( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, model_id: str, ) -> bool: """Validate that a model is configured and available. Args: - client: The AsyncLlamaStackClient instance + client: The AsyncOgxClient instance model_id: The model identifier in "provider/model" format Returns: @@ -1290,15 +1329,15 @@ async def check_model_configured( HTTPException: If there's a connection error or other API error """ try: - models = await client.models.list() + models = parse_model_list_response(await client.models.list()) for model in models: - if model.id == model_id: + if model.identifier == model_id: return True # Workaround to llama-stack watsonx bug - if model_id.startswith("watsonx/") and model.id == model_id.removeprefix( + if model_id.startswith( "watsonx/" - ): + ) and model.identifier == model_id.removeprefix("watsonx/"): return True return False except APIStatusError as e: @@ -1306,7 +1345,7 @@ async def check_model_configured( raise HTTPException(**response.model_dump()) from e except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e @@ -1314,7 +1353,7 @@ async def check_model_configured( async def select_model_for_responses( request_model: Optional[str], - client: AsyncLlamaStackClient, + client: AsyncOgxClient, user_conversation: Optional[UserConversation], ) -> str: """Select model for Responses API if not explicitly specified in the request. @@ -1327,7 +1366,7 @@ async def select_model_for_responses( Args: request_model: The model explicitly specified in the request, or None if not specified - client: The AsyncLlamaStackClient instance + client: The AsyncOgxClient instance user_conversation: The user conversation if conversation_id was provided, None otherwise Returns: @@ -1356,10 +1395,10 @@ async def select_model_for_responses( # 3. Fetch models list and select the first LLM model (model_type="llm") try: - models = await client.models.list() + models = parse_model_list_response(await client.models.list()) except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e @@ -1367,27 +1406,20 @@ async def select_model_for_responses( error_response = InternalServerErrorResponse.generic() raise HTTPException(**error_response.model_dump()) from e - llm_models = [ - m - for m in models - if m.custom_metadata and m.custom_metadata.get("model_type") == "llm" - ] + llm_models = [m for m in models if m.model_type == "llm"] if not llm_models: logger.error("No LLM model found in available models") response = NotFoundResponse(resource="model", resource_id=None) raise HTTPException(**response.model_dump()) model = llm_models[0] - logger.info("Selected first LLM model: %s", model.id) + logger.info("Selected first LLM model: %s", model.identifier) # Workaround to llama-stack bug for watsonx # model needs to be "watsonx/" in the response request - metadata = model.custom_metadata or {} - if metadata.get("provider_id") == "watsonx": - provider_resource_id = metadata.get("provider_resource_id") - if isinstance(provider_resource_id, str): - return provider_resource_id - return model.id + if model.provider_id == "watsonx" and model.provider_resource_id: + return model.provider_resource_id + return model.identifier def is_server_deployed_output(output_item: ResponseOutput) -> bool: @@ -1468,9 +1500,9 @@ def build_turn_summary( # pylint: disable=too-many-arguments,too-many-positiona continue tool_call, tool_result = build_tool_call_summary(item) if tool_call: - summary.tool_calls.append(tool_call) + summary.tool_calls.append(tool_call) # pylint: disable=no-member if tool_result: - summary.tool_results.append(tool_result) + summary.tool_results.append(tool_result) # pylint: disable=no-member summary.rag_chunks = parse_rag_chunks(response, vector_store_ids, rag_id_mapping) summary.token_usage = extract_token_usage(response.usage, model, endpoint_path) @@ -1567,7 +1599,7 @@ def deduplicate_referenced_documents( async def create_new_conversation( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, ) -> str: """Create a new conversation via the Llama Stack Conversations API. @@ -1582,7 +1614,7 @@ async def create_new_conversation( return conversation.id except APIConnectionError as e: error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause=str(e), ) raise HTTPException(**error_response.model_dump()) from e diff --git a/src/utils/shields.py b/src/utils/shields.py index 5dca71ad3..dd72da72c 100644 --- a/src/utils/shields.py +++ b/src/utils/shields.py @@ -1,88 +1,35 @@ -"""Utility functions for working with Llama Stack shields.""" +"""Utility helpers for shield override validation and moderation.""" -from typing import Any, Optional +from typing import Optional from fastapi import HTTPException -from llama_stack_api import OpenAIResponseMessage -from llama_stack_client import ( - APIConnectionError, - AsyncLlamaStackClient, -) -from llama_stack_client import ( - APIStatusError as LLSApiStatusError, -) -from llama_stack_client.types import ShieldListResponse -from openai._exceptions import APIStatusError as OpenAIAPIStatusError +from ogx_client import AsyncOgxClient +from pydantic_ai.exceptions import AgentRunError from configuration import AppConfig -from constants import DEFAULT_VIOLATION_MESSAGE from log import get_logger -from metrics import recording from models.api.requests import QueryRequest from models.api.responses.error import ( - InternalServerErrorResponse, NotFoundResponse, - ServiceUnavailableResponse, UnprocessableEntityResponse, ) from models.common.moderation import ( - ShieldModerationBlocked, ShieldModerationPassed, ShieldModerationResult, ) -from utils.query import handle_known_apistatus_errors +from models.config import QuestionValidityConfig, RedactionConfig, ShieldConfiguration +from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability +from pydantic_ai_lightspeed.capabilities.question_validity._capability import ( + QuestionValidity, +) +from pydantic_ai_lightspeed.capabilities.redaction._capability import ( + PiiRedactionCapability, +) +from utils.agents.error_handler import map_agent_inference_error logger = get_logger(__name__) -async def get_available_shields(client: AsyncLlamaStackClient) -> list[str]: - """ - Discover and return available shield identifiers. - - Parameters: - ---------- - client: The Llama Stack client to query for available shields. - - Returns: - ------- - list[str]: List of available shield identifiers; empty if no shields are available. - """ - available_shields = [shield.identifier for shield in await client.shields.list()] - if not available_shields: - logger.info("No available shields. Disabling safety") - else: - logger.info("Available shields: %s", available_shields) - return available_shields - - -def detect_shield_violations(output_items: list[Any]) -> bool: - """ - Check output items for shield violations and update metrics. - - Iterates through output items looking for message items with refusal - attributes. If a refusal is found, increments the validation error - metric and logs a warning. - - Parameters: - ---------- - output_items: List of output items from the LLM response to check. - - Returns: - ------- - bool: True if a shield violation was detected, False otherwise. - """ - for output_item in output_items: - item_type = getattr(output_item, "type", None) - if item_type == "message": - refusal = getattr(output_item, "refusal", None) - if refusal: - # Metric for LLM validation errors (shield violations) - recording.record_llm_validation_error() - logger.warning("Shield violation detected: %s", refusal) - return True - return False - - def validate_shield_ids_override( query_request: QueryRequest, config: AppConfig ) -> None: @@ -119,178 +66,128 @@ def validate_shield_ids_override( raise HTTPException(**response.model_dump()) -async def run_shield_moderation( - client: AsyncLlamaStackClient, +async def run_shield_moderation_v2( input_text: str, - endpoint_path: str, - shield_ids: Optional[list[str]] = None, + shield_configs: list[ShieldConfiguration], + selected_shield_ids: Optional[list[str]] = None, ) -> ShieldModerationResult: - """ - Run shield moderation on input text. + """Run v2 shield moderation on input text. Iterates through configured shields and runs moderation checks. - Raises HTTPException if shield model is not found. Parameters: - ---------- - client: The Llama Stack client. input_text: The text to moderate. - endpoint_path: The API endpoint path for metric labeling. - shield_ids: Optional list of shield IDs to use. If None, uses all shields. - If empty list, skips all shields. + shield_configs: List of shield configurations to evaluate. + selected_shield_ids: Optional list of shield names to filter by. Returns: - ------- - ShieldModerationResult: Result indicating if content was blocked and the message. - - Raises: - ------ - HTTPException: If shield's provider_resource_id is not configured or model not found. + Result indicating if content was blocked or passed. """ - shields_to_run = await get_shields_for_request(client, shield_ids) - available_models = {model.id for model in await client.models.list()} - for shield in shields_to_run: - # Lightspeed safety providers configure their model internally - # so provider_resource_id is not necessarily a valid model ID. - if shield.provider_id == "llama-guard" and ( - not shield.provider_resource_id - or shield.provider_resource_id not in available_models - ): - logger.error("Shield model not found: %s", shield.provider_resource_id) - response = NotFoundResponse( - resource="Shield model", resource_id=shield.provider_resource_id or "" - ) - raise HTTPException(**response.model_dump()) + selected_shield_configs = get_shields_for_request( + shield_configs, selected_shield_ids + ) + + for shield_config in selected_shield_configs: + shield = build_shield(shield_config) try: - moderation_result = await client.moderations.create( - input=input_text, model=shield.provider_resource_id - ) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except (LLSApiStatusError, OpenAIAPIStatusError) as e: - error_response = handle_known_apistatus_errors( - e, shield.provider_resource_id or "" - ) - raise HTTPException(**error_response.model_dump()) from e - - if moderation_result.results and moderation_result.results[0].flagged: - result = moderation_result.results[0] - recording.record_llm_validation_error(endpoint_path) - logger.warning( - "Shield '%s' flagged content: categories=%s", - shield.identifier, - result.categories, - ) - violation_message = result.user_message or DEFAULT_VIOLATION_MESSAGE - return ShieldModerationBlocked( - message=violation_message, - moderation_id=moderation_result.id, - refusal_response=create_refusal_response(violation_message), - ) + shield_result = await shield.run(input_text) + # APIConnectionError and APIStatusError from ogx should not be raised from model_request, + # because they will be caught inside AsyncOpenAI and transferred into openai's + # APIConnectionError. The openai's exceptions will further transferred into ModelHTTPError + # or ModelAPIError by _map_api_errors in OpenAIResponseModel. + except (AgentRunError, RuntimeError) as exc: + model_id = getattr(shield_config.config, "model_id", "unknown-shield-model") + response = map_agent_inference_error(exc, model_id) + raise HTTPException(**response.model_dump()) from exc + + if shield_result.decision == "blocked": + return shield_result return ShieldModerationPassed() -async def append_turn_to_conversation( - client: AsyncLlamaStackClient, - conversation_id: str, - user_message: str, - assistant_message: str, -) -> None: - """ - Append a user/assistant turn to a conversation after shield violation. - - Used to record the conversation turn when a shield blocks the request, - storing both the user's original message and the violation response. +def build_shield(shield_config: ShieldConfiguration) -> AbstractSafetyCapability: + """Build a safety capability instance from a shield configuration. Parameters: - ---------- - client: The Llama Stack client. - conversation_id: The Llama Stack conversation ID. - user_message: The user's input message. - assistant_message: The shield violation response message. + shield_config: The shield configuration to build from. + + Returns: + The constructed safety capability. """ - try: - await client.conversations.items.create( - conversation_id, - items=[ - {"type": "message", "role": "user", "content": user_message}, - {"type": "message", "role": "assistant", "content": assistant_message}, - ], - ) - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(e), - ) - raise HTTPException(**error_response.model_dump()) from e - except LLSApiStatusError as e: - error_response = InternalServerErrorResponse.generic() - raise HTTPException(**error_response.model_dump()) from e + match shield_config.config: + case QuestionValidityConfig(): + return QuestionValidity(shield_config.config) + case RedactionConfig(): + return PiiRedactionCapability(shield_config.config) -def create_refusal_response(refusal_message: str) -> OpenAIResponseMessage: - """Create a refusal response message object. +async def run_shield_moderation( + _client: AsyncOgxClient, + _input_text: str, + _endpoint_path: str, + _shield_ids: Optional[list[str]] = None, +) -> ShieldModerationResult: + """ + Run shield moderation on input text. - Args: - refusal_message: The refusal message text. + Iterates through configured shields and runs moderation checks. + Raises HTTPException if shield model is not found. + + Parameters: + ---------- + client: The Llama Stack client. + input_text: The text to moderate. + endpoint_path: The API endpoint path for metric labeling. + shield_ids: Optional list of shield IDs to use. If None, uses all shields. + If empty list, skips all shields. Returns: - OpenAIResponseMessage with refusal message. + ------- + ShieldModerationResult: Result indicating if content was blocked and the message. + + Raises: + ------ + HTTPException: If shield's provider_resource_id is not configured or model not found. """ - return OpenAIResponseMessage( - role="assistant", - content=refusal_message, - ) + # Currently stubbed to always pass until LCS-owned input shields are wired. + return ShieldModerationPassed() -async def get_shields_for_request( - client: AsyncLlamaStackClient, +def get_shields_for_request( + shields: list[ShieldConfiguration], shield_ids: Optional[list[str]] = None, -) -> ShieldListResponse: - """Resolve shields for the request: filtered by shield_ids or all configured. +) -> list[ShieldConfiguration]: + """Return configured shields, optionally filtered by request shield_ids. Args: - client: Llama Stack client. - shield_ids: Optional list of shield IDs. If provided, only shields - with these identifiers are returned; if None, all configured - shields are returned. + shields: Configured LCS shields. + shield_ids: Optional list of shield names. If None, all shields are + returned. An empty list skips all shields. Otherwise only shields + whose name is in this list are returned. Returns: - ShieldListResponse: List of Shield objects to run for this request. + list[ShieldConfiguration]: Shield configurations to run for this request. Raises: - HTTPException: 404 if shield_ids is provided and any requested - shield is not configured in Llama Stack. + HTTPException: 404 if shield_ids is provided and any requested shield + name is not present in shields. """ + if shield_ids is None: + return list(shields) + if shield_ids == []: return [] - try: - configured_shields: ShieldListResponse = await client.shields.list() - if shield_ids is None: - return configured_shields - requested = set(shield_ids) - configured_ids = {s.identifier for s in configured_shields} - missing = requested - configured_ids - if missing: - response = NotFoundResponse( - resource=f"Shield{'s' if len(missing) > 1 else ''}", - resource_id=", ".join(missing), - ) - raise HTTPException(**response.model_dump()) - - return [s for s in configured_shields if s.identifier in requested] - except APIConnectionError as e: - error_response = ServiceUnavailableResponse( - backend_name="Llama Stack", - cause=str(e), + + requested = set(shield_ids) + configured_names = {shield.name for shield in shields} + missing = requested - configured_names + if missing: + response = NotFoundResponse( + resource=f"Shield{'s' if len(missing) > 1 else ''}", + resource_id=", ".join(sorted(missing)), ) - raise HTTPException(**error_response.model_dump()) from e - except LLSApiStatusError as e: - error_response = InternalServerErrorResponse.generic() - raise HTTPException(**error_response.model_dump()) from e + raise HTTPException(**response.model_dump()) + + return [shield for shield in shields if shield.name in requested] diff --git a/src/utils/stream_interrupts.py b/src/utils/stream_interrupts.py index 5afaf92f8..acdd771b3 100644 --- a/src/utils/stream_interrupts.py +++ b/src/utils/stream_interrupts.py @@ -8,7 +8,7 @@ from threading import Lock from typing import Any, Optional, cast -from llama_stack_api import OpenAIResponseMessage +from ogx_api import OpenAIResponseMessage from constants import ( INTERRUPTED_RESPONSE_MESSAGE, @@ -19,11 +19,13 @@ from models.common.responses.responses_api_params import ResponsesApiParams from models.common.responses.types import ResponseInput from models.common.turn_summary import TurnSummary -from utils.conversations import append_turn_items_to_conversation +from utils.conversations import ( + append_turn_items_to_conversation, + append_turn_to_conversation, +) from utils.markdown_repair import close_open_markdown from utils.query import store_query_results, update_conversation_topic_summary from utils.responses import get_topic_summary -from utils.shields import append_turn_to_conversation from utils.types import Singleton logger = get_logger(__name__) diff --git a/src/utils/tool_formatter.py b/src/utils/tool_formatter.py index 4b55141ea..ba67a9917 100644 --- a/src/utils/tool_formatter.py +++ b/src/utils/tool_formatter.py @@ -1,54 +1,45 @@ """Utility functions for formatting and parsing MCP tool descriptions.""" +from __future__ import annotations + from collections.abc import Mapping -from typing import Any +from typing import Any, Optional from log import get_logger +from models.common.tools import ( + CatalogTool, + CatalogToolParameter, + ListedMcpTool, +) logger = get_logger(__name__) -def format_tool_response(tool_dict: dict[str, Any]) -> dict[str, Any]: - """ - Format a tool dictionary to include only required fields. - - If the input description contains structured metadata (e.g., - lines starting with `TOOL_NAME=` or `DISPLAY_NAME=`), the - description will be replaced with a cleaned, human-readable - version extracted by `extract_clean_description`. +def input_schema_to_parameters( + schema: Optional[dict[str, Any]], +) -> list[CatalogToolParameter]: + """Convert a JSON Schema object to the flat parameter list used by ``/tools``. Parameters: - ---------- - tool_dict: Raw tool dictionary from Llama Stack + schema: JSON Schema dict with ``properties`` and ``required`` keys. Returns: - ------- - dict[str, Any]: Formatted tool dictionary containing the following keys: - - identifier: tool identifier string (defaults to ""). - - description: cleaned or original description string. - - parameters: list of parameter definitions (defaults to empty list). - - provider_id: provider identifier string (defaults to ""). - - toolgroup_id: tool group identifier string (defaults to ""). - - server_source: server source string (defaults to ""). - - type: tool type string (defaults to ""). + Flat parameter models for the tools endpoint response. """ - # Clean up description if it contains structured metadata - description = tool_dict.get("description", "") - if description and ("TOOL_NAME=" in description or "DISPLAY_NAME=" in description): - # Extract clean description from structured metadata - clean_description = extract_clean_description(description) - description = clean_description - - # Extract only the required fields - return { - "identifier": tool_dict.get("identifier", ""), - "description": description, - "parameters": tool_dict.get("parameters", []), - "provider_id": tool_dict.get("provider_id", ""), - "toolgroup_id": tool_dict.get("toolgroup_id", ""), - "server_source": tool_dict.get("server_source", ""), - "type": tool_dict.get("type", ""), - } + if not schema or "properties" not in schema: + return [] + + required_params = set(schema.get("required", [])) + return [ + CatalogToolParameter( + name=name, + description=prop.get("description", ""), + parameter_type=prop.get("type", "string"), + required=name in required_params, + default=prop.get("default"), + ) + for name, prop in schema["properties"].items() + ] def extract_clean_description(description: str) -> str: @@ -123,20 +114,37 @@ def extract_clean_description(description: str) -> str: ) -def format_tools_list(tools: list[dict[str, Any]]) -> list[dict[str, Any]]: - """ - Format a list of tools with structured description parsing. +def build_catalog_tool( + tool: ListedMcpTool, + provider_id: str, + toolgroup_id: str, + server_source: str, +) -> CatalogTool: + """Build a ``/tools`` catalog entry from discovered tool metadata. Parameters: - ---------- - tools: (list[dict[str, Any]]): List of raw tool dictionaries + tool: MCP tool definition with name, description, and schema. + provider_id: Provider ID serving the tool. + toolgroup_id: Tool group identifier. + server_source: Human-readable source label for the tool. Returns: - ------- - list[dict[str, Any]]: Formatted tool dictionaries with normalized - fields and cleaned descriptions. + Typed catalog tool entry. """ - return [format_tool_response(tool) for tool in tools] + cleaned_description = tool.description or "" + if cleaned_description and ( + "TOOL_NAME=" in cleaned_description or "DISPLAY_NAME=" in cleaned_description + ): + cleaned_description = extract_clean_description(cleaned_description) + + return CatalogTool( + identifier=tool.name, + description=cleaned_description, + parameters=input_schema_to_parameters(tool.input_schema), + provider_id=provider_id, + toolgroup_id=toolgroup_id, + server_source=server_source, + ) def translate_vector_store_ids_to_user_facing( diff --git a/src/utils/types.py b/src/utils/types.py index 93ca6e055..0af0296bc 100644 --- a/src/utils/types.py +++ b/src/utils/types.py @@ -3,7 +3,7 @@ from re import Pattern from typing import Any -from llama_stack_api import ImageContentItem, TextContentItem +from ogx_api import ImageContentItem, TextContentItem type SingletonInstances = dict[type, Any] diff --git a/src/utils/vector_search.py b/src/utils/vector_search.py index 2e0fd3dea..9cf4ffff4 100644 --- a/src/utils/vector_search.py +++ b/src/utils/vector_search.py @@ -9,10 +9,10 @@ from typing import Any, Optional, cast from urllib.parse import urljoin -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMessage as ResponseMessage, ) -from llama_stack_client import AsyncLlamaStackClient +from ogx_client import AsyncOgxClient from pydantic import AnyUrl import constants @@ -246,7 +246,7 @@ def _format_rag_context(rag_chunks: list[RAGChunk], query: str) -> str: async def _query_store_for_byok_rag( - client: AsyncLlamaStackClient, + client: AsyncOgxClient, vector_store_id: str, query: str, weight: float, @@ -255,7 +255,7 @@ async def _query_store_for_byok_rag( """Query a single vector store for BYOK RAG. Args: - client: AsyncLlamaStackClient for vector_io queries + client: AsyncOgxClient for vector_io queries vector_store_id: ID of the vector store to query query: Search query string weight: Score multiplier to apply @@ -440,7 +440,7 @@ def _process_solr_chunks_for_documents( async def _fetch_byok_rag( # pylint: disable=too-many-locals - client: AsyncLlamaStackClient, + client: AsyncOgxClient, query: str, vector_store_ids: Optional[list[str]] = None, max_chunks: Optional[int] = None, @@ -448,7 +448,7 @@ async def _fetch_byok_rag( # pylint: disable=too-many-locals """Fetch chunks and documents from BYOK RAG sources. Args: - client: The AsyncLlamaStackClient to use for the request + client: The AsyncOgxClient to use for the request query: The search query vector_store_ids: Optional list of vector store IDs to query. If provided, only these stores will be queried. If None, all stores @@ -551,14 +551,14 @@ async def _fetch_byok_rag( # pylint: disable=too-many-locals async def _fetch_solr_rag( # pylint: disable=too-many-locals - client: AsyncLlamaStackClient, + client: AsyncOgxClient, query: str, solr: Optional[SolrVectorSearchRequest] = None, ) -> tuple[list[RAGChunk], list[ReferencedDocument]]: """Fetch chunks and documents from Solr RAG source. Args: - client: The AsyncLlamaStackClient to use for the request + client: The AsyncOgxClient to use for the request query: The user's query solr: Structured Solr inline RAG request from the API (optional). max_chunks: Maximum number of chunks to return. If None, uses @@ -630,7 +630,7 @@ async def _fetch_solr_rag( # pylint: disable=too-many-locals async def build_rag_context( # pylint: disable=too-many-locals,too-many-branches - client: AsyncLlamaStackClient, + client: AsyncOgxClient, moderation_decision: str, # pylint: disable=unused-argument query: str, vector_store_ids: Optional[list[str]], @@ -644,7 +644,7 @@ async def build_rag_context( # pylint: disable=too-many-locals,too-many-branche and/or Solr OKP. Args: - client: The AsyncLlamaStackClient to use for the request + client: The AsyncOgxClient to use for the request query: The user's query vector_store_ids: The vector store IDs to query solr: Structured Solr inline RAG request from the API (optional). diff --git a/tests/benchmarks/data/python_1000_lines.py b/tests/benchmarks/data/python_1000_lines.py index c81f27d59..d5399573e 100644 --- a/tests/benchmarks/data/python_1000_lines.py +++ b/tests/benchmarks/data/python_1000_lines.py @@ -439,8 +439,8 @@ from langchain.prompts import PromptTemplate from langchain_core.output_parsers import StrOutputParser from langchain_openai import ChatOpenAI -from llama_stack.distribution.library_client import LlamaStackAsLibraryClient -from llama_stack_client import LlamaStackClient +from ogx.core.library_client import OGXAsLibraryClient +from ogx_client import OgxClient from openai import OpenAI client = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) @@ -678,7 +678,7 @@ # Získání seznamu všech dostupných modelů -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") @@ -695,7 +695,7 @@ # Získání seznamu všech dostupných modelů -client = LlamaStackAsLibraryClient("run.yaml") +client = OGXAsLibraryClient("run.yaml") client.initialize() print(f"Using Llama Stack version {client._version}") @@ -709,7 +709,7 @@ # # ### Komunikace s LLM -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") @@ -739,7 +739,7 @@ # ### Využití novějšího API -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") @@ -774,7 +774,7 @@ # * vytvoření nové vektorové databáze # * inicializace vektorové databáze -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") vector_store_name = f"vec_{str(uuid.uuid4())[0:8]}" diff --git a/tests/benchmarks/data/python_100_lines.py b/tests/benchmarks/data/python_100_lines.py index 70fd5c5e6..1ccb24397 100644 --- a/tests/benchmarks/data/python_100_lines.py +++ b/tests/benchmarks/data/python_100_lines.py @@ -10,10 +10,10 @@ # Získání seznamu všech dostupných modelů -from llama_stack.distribution.library_client import LlamaStackAsLibraryClient -from llama_stack_client import LlamaStackClient +from ogx.core.library_client import OGXAsLibraryClient +from ogx_client import OgxClient -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") @@ -29,7 +29,7 @@ # Získání seznamu všech dostupných modelů -client = LlamaStackAsLibraryClient("run.yaml") +client = OGXAsLibraryClient("run.yaml") client.initialize() print(f"Using Llama Stack version {client._version}") @@ -43,7 +43,7 @@ # # ### Komunikace s LLM -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") @@ -73,7 +73,7 @@ # ### Využití novějšího API -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") print(f"Using Llama Stack version {client._version}") diff --git a/tests/benchmarks/data/python_10_lines.py b/tests/benchmarks/data/python_10_lines.py index 2f9dc742e..da8853cb0 100644 --- a/tests/benchmarks/data/python_10_lines.py +++ b/tests/benchmarks/data/python_10_lines.py @@ -1,8 +1,8 @@ """Source to be tokenized.""" -from llama_stack_client import LlamaStackClient +from ogx_client import OgxClient -client = LlamaStackClient(base_url="http://localhost:8321") +client = OgxClient(base_url="http://localhost:8321") models = client.models.list() diff --git a/tests/configuration/minimal-stack.yaml b/tests/configuration/minimal-stack.yaml index 9f4ea1491..c0e5f5305 100644 --- a/tests/configuration/minimal-stack.yaml +++ b/tests/configuration/minimal-stack.yaml @@ -1,5 +1,5 @@ version: '2' -image_name: llamastack-minimal-stack +distro_name: llamastack-minimal-stack container_image: null external_providers_dir: /tmp @@ -26,5 +26,8 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default diff --git a/tests/configuration/run.yaml b/tests/configuration/run.yaml index a0668bb90..374dab491 100644 --- a/tests/configuration/run.yaml +++ b/tests/configuration/run.yaml @@ -1,53 +1,27 @@ version: '2' -image_name: minimal-viable-llama-stack-configuration +distro_name: minimal-viable-llama-stack-configuration apis: - - agents + - responses - datasetio - eval - inference - post_training - - safety - scoring - tool_runtime - vector_io -benchmarks: [] container_image: null -datasets: [] external_providers_dir: null logging: null providers: - agents: - - provider_id: meta-reference - provider_type: inline::meta-reference + responses: + - provider_id: builtin + provider_type: inline::builtin config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - datasetio: - - provider_id: huggingface - provider_type: remote::huggingface - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - - provider_id: localfs - provider_type: inline::localfs - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - eval: - - provider_id: meta-reference - provider_type: inline::meta-reference - config: - kvstore: - namespace: eval_store - backend: kv_default inference: - provider_id: openai provider_type: remote::openai @@ -64,32 +38,16 @@ providers: distributed_backend: null provider_id: huggingface provider_type: inline::huggingface - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - config: {} - provider_id: basic - provider_type: inline::basic - - config: {} - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - - config: - openai_api_key: '********' - provider_id: braintrust - provider_type: inline::braintrust telemetry: - config: service_name: '' sinks: sqlite sqlite_db_path: .llama/distributions/ollama/trace_store.db - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin tool_runtime: - - provider_id: rag-runtime - provider_type: inline::rag-runtime + - provider_id: file-search + provider_type: inline::file-search config: {} - provider_id: model-context-protocol provider_type: remote::model-context-protocol @@ -101,7 +59,6 @@ providers: backend: kv_rag provider_id: faiss provider_type: inline::faiss -scoring_fns: [] server: auth: null host: null @@ -115,6 +72,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -129,18 +89,14 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: [] - shields: [] vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag - provider_id: rag-runtime telemetry: enabled: true vector_stores: diff --git a/tests/e2e-prow/rhoai/configs/run.yaml b/tests/e2e-prow/rhoai/configs/run.yaml index 9348af6b9..ec50feabb 100644 --- a/tests/e2e-prow/rhoai/configs/run.yaml +++ b/tests/e2e-prow/rhoai/configs/run.yaml @@ -1,21 +1,15 @@ version: 2 -image_name: starter +distro_name: starter apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] - providers: inference: - provider_id: vllm @@ -36,26 +30,10 @@ providers: storage_dir: /opt/app-root/src/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol @@ -66,45 +44,21 @@ providers: backend: kv_rag provider_id: faiss provider_type: inline::faiss - agents: + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -131,8 +85,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: meta-llama/Llama-3.1-8B-Instruct @@ -140,20 +97,8 @@ registered_resources: model_type: llm provider_model_id: null vector_stores: [] - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: vllm/meta-llama/Llama-3.1-8B-Instruct - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - provider_id: rag-runtime - toolgroup_id: builtin::rag vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-openai.yaml b/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-openai.yaml index 80651a3cb..11e76dcf8 100644 --- a/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-openai.yaml +++ b/tests/e2e-prow/rhoai/manifests/lightspeed/llama-stack-openai.yaml @@ -205,12 +205,12 @@ spec: if [[ -f "$ENRICHED_CONFIG" ]] && [[ "$ENRICHMENT_FAILED" -eq 0 ]]; then echo "Using enriched config: $ENRICHED_CONFIG" restore_rag_seed - exec llama stack run "$ENRICHED_CONFIG" + exec ogx stack run "$ENRICHED_CONFIG" fi fi echo "Using original config: $INPUT_CONFIG" restore_rag_seed - exec llama stack run "$INPUT_CONFIG" + exec ogx stack run "$INPUT_CONFIG" ports: - containerPort: 8321 readinessProbe: 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 f33f9519b..271304cbd 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 llama stack run "$ENRICHED_CONFIG" + exec ogx stack run "$ENRICHED_CONFIG" fi fi echo "Using original config: $INPUT_CONFIG" restore_rag_seed - exec llama stack run "$INPUT_CONFIG" + exec ogx stack run "$INPUT_CONFIG" ports: - containerPort: 8321 readinessProbe: diff --git a/tests/e2e-prow/rhoai/pipeline-konflux.sh b/tests/e2e-prow/rhoai/pipeline-konflux.sh index 2e9c8f865..c711826ed 100755 --- a/tests/e2e-prow/rhoai/pipeline-konflux.sh +++ b/tests/e2e-prow/rhoai/pipeline-konflux.sh @@ -189,13 +189,14 @@ RAG_DB_PATH="$REPO_ROOT/tests/e2e/rag/kv_store.db" if [ -f "$RAG_DB_PATH" ]; then # Extract vector store ID from kv_store.db using Python (sqlite3 CLI may not be available) log "Extracting vector store ID from kv_store.db..." - # Key format is: vector_stores:v3::vs_xxx or openai_vector_stores:v3::vs_xxx + # OGX 1.0 FAISS keys use persistence.namespace prefix, e.g.: + # vector_io::faiss:vector_stores:v3::vs_xxx export FAISS_VECTOR_STORE_ID=$(python3 -c " import sqlite3 import re conn = sqlite3.connect('$RAG_DB_PATH') cursor = conn.cursor() -cursor.execute(\"SELECT key FROM kvstore WHERE key LIKE 'vector_stores:v%::%' LIMIT 1\") +cursor.execute(\"SELECT key FROM kvstore WHERE key LIKE 'vector_io::faiss:vector_stores:v%::%' LIMIT 1\") row = cursor.fetchone() if row: # Extract the vs_xxx ID from the key @@ -326,7 +327,7 @@ PF_JWKS_PID=$! # Behave runs in this shell; pipeline-services-konflux.sh cannot export here. MCP hooks call # Llama Stack directly — mirror LCS and forward llama-stack-service-svc to localhost:8321. -log "Starting port-forward for llama-stack (MCP / llama_stack_client hooks)..." +log "Starting port-forward for llama-stack (MCP / ogx_client hooks)..." oc port-forward svc/llama-stack-service-svc 8321:8321 -n $NAMESPACE & PF_LLAMA_PID=$! echo "$PF_LLAMA_PID" >"$E2E_LLAMA_PORT_FORWARD_PID_FILE" diff --git a/tests/e2e-prow/rhoai/pipeline.sh b/tests/e2e-prow/rhoai/pipeline.sh index 374904f46..69ea17ad2 100755 --- a/tests/e2e-prow/rhoai/pipeline.sh +++ b/tests/e2e-prow/rhoai/pipeline.sh @@ -264,13 +264,14 @@ RAG_DB_PATH="$REPO_ROOT/tests/e2e/rag/kv_store.db" if [ -f "$RAG_DB_PATH" ]; then # Extract vector store ID from kv_store.db using Python (sqlite3 CLI may not be available) echo "Extracting vector store ID from kv_store.db..." - # Key format is: vector_stores:v3::vs_xxx or openai_vector_stores:v3::vs_xxx + # OGX 1.0 FAISS keys use persistence.namespace prefix, e.g.: + # vector_io::faiss:vector_stores:v3::vs_xxx export FAISS_VECTOR_STORE_ID=$(python3 -c " import sqlite3 import re conn = sqlite3.connect('$RAG_DB_PATH') cursor = conn.cursor() -cursor.execute(\"SELECT key FROM kvstore WHERE key LIKE 'vector_stores:v%::%' LIMIT 1\") +cursor.execute(\"SELECT key FROM kvstore WHERE key LIKE 'vector_io::faiss:vector_stores:v%::%' LIMIT 1\") row = cursor.fetchone() if row: # Extract the vs_xxx ID from the key diff --git a/tests/e2e/configs/run-azure.yaml b/tests/e2e/configs/run-azure.yaml index 6f627ba8b..7880b4f41 100644 --- a/tests/e2e/configs/run-azure.yaml +++ b/tests/e2e/configs/run-azure.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: azure-configuration +distro_name: azure-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -42,69 +37,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -112,6 +73,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -128,8 +92,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: gpt-4o-mini @@ -142,21 +109,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configs/run-bedrock.yaml b/tests/e2e/configs/run-bedrock.yaml index 02f902b2d..2de83e64d 100644 --- a/tests/e2e/configs/run-bedrock.yaml +++ b/tests/e2e/configs/run-bedrock.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: starter +distro_name: starter apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -40,69 +35,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -110,6 +71,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -126,8 +90,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: custom-bedrock-model @@ -140,21 +107,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configs/run-ci.yaml b/tests/e2e/configs/run-ci.yaml index 5c9dd8cb1..eeae1ec81 100644 --- a/tests/e2e/configs/run-ci.yaml +++ b/tests/e2e/configs/run-ci.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: starter +distro_name: starter apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -35,69 +30,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -105,6 +66,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -121,8 +85,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: all-mpnet-base-v2 @@ -131,21 +98,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configs/run-rhaiis.yaml b/tests/e2e/configs/run-rhaiis.yaml index 74c0b654b..542e7cc4c 100644 --- a/tests/e2e/configs/run-rhaiis.yaml +++ b/tests/e2e/configs/run-rhaiis.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: rhaiis-configuration +distro_name: rhaiis-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -42,69 +37,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -112,6 +73,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -128,8 +92,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: ${env.RHAIIS_MODEL} @@ -142,21 +109,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configs/run-rhelai.yaml b/tests/e2e/configs/run-rhelai.yaml index 4fa263387..29e377e82 100644 --- a/tests/e2e/configs/run-rhelai.yaml +++ b/tests/e2e/configs/run-rhelai.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: rhelai-configuration +distro_name: rhelai-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -42,69 +37,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -112,6 +73,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -128,8 +92,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: ${env.VLLM_MODEL} @@ -142,21 +109,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configs/run-vertexai.yaml b/tests/e2e/configs/run-vertexai.yaml index aa5f4925e..341413097 100644 --- a/tests/e2e/configs/run-vertexai.yaml +++ b/tests/e2e/configs/run-vertexai.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: vertexai-configuration +distro_name: vertexai-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -41,69 +36,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -111,6 +72,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -127,8 +91,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: all-mpnet-base-v2 @@ -137,21 +104,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configs/run-watsonx.yaml b/tests/e2e/configs/run-watsonx.yaml index a472471b6..f410eb44a 100644 --- a/tests/e2e/configs/run-watsonx.yaml +++ b/tests/e2e/configs/run-watsonx.yaml @@ -1,20 +1,15 @@ version: 2 -image_name: watsonx-configuration +distro_name: watsonx-configuration apis: -- agents +- responses - batches -- datasetio -- eval - files - inference -- safety -- scoring - tool_runtime +- conversations - vector_io -benchmarks: [] -datasets: [] # external_providers_dir: /opt/app-root/src/.llama/providers.d providers: @@ -42,69 +37,35 @@ providers: storage_dir: ~/.llama/storage/files provider_id: meta-reference-files provider_type: inline::localfs - safety: - - config: - excluded_categories: [] - provider_id: llama-guard - provider_type: inline::llama-guard - scoring: - - provider_id: basic - provider_type: inline::basic - config: {} - - provider_id: llm-as-judge - provider_type: inline::llm-as-judge - config: {} - - provider_id: braintrust - provider_type: inline::braintrust - config: - openai_api_key: '********' tool_runtime: - config: {} # Enable the RAG tool - provider_id: rag-runtime - provider_type: inline::rag-runtime + provider_id: file-search + provider_type: inline::file-search - config: {} # Enable MCP (Model Context Protocol) support provider_id: model-context-protocol provider_type: remote::model-context-protocol - vector_io: [] - agents: + vector_io: + - config: # Define the storage backend for RAG + persistence: + namespace: vector_io::faiss + backend: kv_rag + provider_id: faiss + provider_type: inline::faiss + responses: - config: persistence: - agent_state: - namespace: agents_state - backend: kv_default responses: table_name: agents_responses backend: sql_default - provider_id: meta-reference - provider_type: inline::meta-reference + provider_id: builtin + provider_type: inline::builtin batches: - config: - kvstore: - namespace: batches_store - backend: kv_default + sqlstore: + table_name: batches + backend: sql_default provider_id: reference provider_type: inline::reference - datasetio: - - config: - kvstore: - namespace: huggingface_datasetio - backend: kv_default - provider_id: huggingface - provider_type: remote::huggingface - - config: - kvstore: - namespace: localfs_datasetio - backend: kv_default - provider_id: localfs - provider_type: inline::localfs - eval: - - config: - kvstore: - namespace: eval_store - backend: kv_default - provider_id: meta-reference - provider_type: inline::meta-reference -scoring_fns: [] server: port: 8321 storage: @@ -112,6 +73,9 @@ storage: kv_default: type: kv_sqlite db_path: ${env.KV_STORE_PATH:=~/.llama/storage/kv_store.db} + kv_rag: # Define the storage backend type for RAG + type: kv_sqlite + db_path: ${env.KV_RAG_PATH:=~/.llama/storage/rag/kv_store.db} sql_default: type: sql_sqlite db_path: ${env.SQL_STORE_PATH:=~/.llama/storage/sql_store.db} @@ -128,8 +92,11 @@ storage: table_name: openai_conversations backend: sql_default prompts: - namespace: prompts - backend: kv_default + table_name: prompts + backend: sql_default + connectors: + table_name: connectors + backend: sql_default registered_resources: models: - model_id: all-mpnet-base-v2 @@ -138,21 +105,9 @@ registered_resources: provider_model_id: all-mpnet-base-v2 metadata: embedding_dimension: 768 - shields: - - shield_id: llama-guard - provider_id: llama-guard - provider_shield_id: openai/gpt-4o-mini vector_stores: [] - datasets: [] - scoring_fns: [] - benchmarks: [] - tool_groups: - - toolgroup_id: builtin::rag # Register the RAG tool - provider_id: rag-runtime vector_stores: default_provider_id: faiss default_embedding_model: # Define the default embedding model for RAG provider_id: sentence-transformers model_id: all-mpnet-base-v2 -safety: - default_shield_id: llama-guard diff --git a/tests/e2e/configuration/library-mode/lightspeed-stack.yaml b/tests/e2e/configuration/library-mode/lightspeed-stack.yaml index 077b2e9ab..825c188fb 100644 --- a/tests/e2e/configuration/library-mode/lightspeed-stack.yaml +++ b/tests/e2e/configuration/library-mode/lightspeed-stack.yaml @@ -35,3 +35,12 @@ byok_rag: rag: tool: - e2e-test-docs + +shields: + - name: pii-redaction + provider_id: redaction + config: + rules: + - pattern: '\d+' + replacement: '[NUM]' + diff --git a/tests/e2e/configuration/server-mode/lightspeed-stack.yaml b/tests/e2e/configuration/server-mode/lightspeed-stack.yaml index de945bb2e..f708052a8 100644 --- a/tests/e2e/configuration/server-mode/lightspeed-stack.yaml +++ b/tests/e2e/configuration/server-mode/lightspeed-stack.yaml @@ -32,4 +32,12 @@ byok_rag: rag: tool: - - e2e-test-docs \ No newline at end of file + - e2e-test-docs + +shields: + - name: pii-redaction + provider_id: redaction + config: + rules: + - pattern: '\d+' + replacement: '[NUM]' diff --git a/tests/e2e/features/info.feature b/tests/e2e/features/info.feature index 8d85aabc9..eac3fbcc3 100644 --- a/tests/e2e/features/info.feature +++ b/tests/e2e/features/info.feature @@ -19,7 +19,7 @@ Feature: Info tests When I access REST API endpoint "info" using HTTP GET method Then The status code of the response is 200 And The body of the response has proper name Lightspeed Core Service (LCS) and version 0.6.0rc2 - And The body of the response has llama-stack version 0.6.0 + And The body of the response has llama-stack version 1.0.2 Scenario: Check if shields endpoint is working When I access REST API endpoint "shields" using HTTP GET method @@ -27,12 +27,10 @@ Feature: Info tests And The body of the response has proper shield structure - #https://issues.redhat.com/browse/LCORE-1211 - @skip Scenario: Check if tools endpoint is working When I access REST API endpoint "tools" using HTTP GET method Then The status code of the response is 200 - And The response contains 2 tools listed for provider rag-runtime + And The response contains 2 tools listed for provider file-search And The body of the response has the following schema """ { @@ -61,13 +59,13 @@ Feature: Info tests } } """ - And The body of the response has proper structure for provider rag-runtime + And The body of the response has proper structure for provider file-search """ { "identifier": "insert_into_memory", "description": "Insert documents into memory", - "provider_id": "rag-runtime", - "toolgroup_id": "builtin::rag", + "provider_id": "file-search", + "toolgroup_id": "builtin::file_search", "server_source": "builtin", "type": "tool" } diff --git a/tests/e2e/features/llama_stack_disrupted.feature b/tests/e2e/features/llama_stack_disrupted.feature index e82b6dc19..5d63e82c6 100644 --- a/tests/e2e/features/llama_stack_disrupted.feature +++ b/tests/e2e/features/llama_stack_disrupted.feature @@ -3,8 +3,8 @@ Feature: Llama Stack connection disrupted End-to-end scenarios that stop the Llama Stack container (or simulate disconnect) and assert degraded responses (503, readiness, etc.). Config order matches test_list.txt: - default stack, then noop-token (query/conversations/…), then rbac (rlsapi errors), then - mcp (immediately before mcp.feature). Skipped in library mode. + default stack, then noop-token (query/conversations/…), then rbac (rlsapi errors). + Skipped in library mode. Background: Given The service is started locally @@ -23,7 +23,7 @@ Feature: Llama Stack connection disrupted Then The status code of the response is 503 And The body of the response is the following """ - {"detail": {"response": "Unable to connect to Llama Stack", "cause": "Connection error."}} + {"detail": {"response": "Unable to connect to OGX", "cause": "Connection error."}} """ Scenario: Check if service report proper readiness state when llama stack is not available @@ -58,18 +58,7 @@ Feature: Llama Stack connection disrupted Then The status code of the response is 503 And The body of the response is the following """ - {"detail": {"response": "Unable to connect to Llama Stack", "cause": "Connection error."}} - """ - - Scenario: Check if shields endpoint reports error when llama-stack is unreachable - Given The service uses the lightspeed-stack.yaml configuration - And The service is restarted - And The llama-stack connection is disrupted - When I access REST API endpoint "shields" using HTTP GET method - Then The status code of the response is 503 - And The body of the response is the following - """ - {"detail": {"response": "Unable to connect to Llama Stack", "cause": "Connection error."}} + {"detail": {"response": "Unable to connect to OGX", "cause": "Connection error."}} """ Scenario: Check if tools endpoint reports error when llama-stack is unreachable @@ -80,7 +69,7 @@ Feature: Llama Stack connection disrupted Then The status code of the response is 503 And The body of the response is the following """ - {"detail": {"response": "Unable to connect to Llama Stack", "cause": "Connection error."}} + {"detail": {"response": "Unable to connect to OGX", "cause": "Connection error."}} """ @@ -96,7 +85,7 @@ Feature: Llama Stack connection disrupted {"query": "Say hello"} """ Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Responses returns error when unable to connect to llama-stack Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -109,7 +98,7 @@ Feature: Llama Stack connection disrupted {"input": "Say hello", "model": "{PROVIDER}/{MODEL}", "stream": false} """ Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Streaming responses returns error when unable to connect to llama-stack Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -121,7 +110,7 @@ Feature: Llama Stack connection disrupted {"input": "Say hello", "model": "{PROVIDER}/{MODEL}", "stream": true} """ Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if rags endpoint fails when llama-stack is unavailable Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -130,7 +119,7 @@ Feature: Llama Stack connection disrupted And The llama-stack connection is disrupted When I access REST API endpoint rags using HTTP GET method Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if prompts list endpoint fails when llama-stack is unavailable Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -139,7 +128,7 @@ Feature: Llama Stack connection disrupted And The llama-stack connection is disrupted When I access REST API endpoint "prompts" using HTTP GET method Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if prompts create endpoint fails when llama-stack is unavailable Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -151,7 +140,7 @@ Feature: Llama Stack connection disrupted {"prompt": "Summarize: {{text}}", "variables": ["text"]} """ Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if prompts get by id endpoint fails when llama-stack is unavailable Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -160,7 +149,7 @@ Feature: Llama Stack connection disrupted And The llama-stack connection is disrupted When I access REST API endpoint "prompts/pmpt_5c76d7f7c633ef97477adeb2f642150d8d08e8a6526e9909" using HTTP GET method Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if prompts update endpoint fails when llama-stack is unavailable Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -172,7 +161,7 @@ Feature: Llama Stack connection disrupted {"prompt": "Summarize in bullets: {{text}}", "version": 1, "set_as_default": true, "variables": ["text"]} """ Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if prompts delete endpoint fails when llama-stack is unavailable Given The service uses the lightspeed-stack-auth-noop-token.yaml configuration @@ -181,7 +170,7 @@ Feature: Llama Stack connection disrupted And The llama-stack connection is disrupted When I access REST API endpoint "prompts/pmpt_5c76d7f7c633ef97477adeb2f642150d8d08e8a6526e9909" using HTTP DELETE method Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if conversations/{conversation_id} GET endpoint fails when llama-stack is unavailable Given Llama Stack is restarted @@ -197,7 +186,7 @@ Feature: Llama Stack connection disrupted And The llama-stack connection is disrupted When I use REST API conversation endpoint with conversation_id from above using HTTP GET method Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check if conversations/{conversation_id} DELETE endpoint fails when llama-stack is unavailable Given Llama Stack is restarted @@ -213,7 +202,7 @@ Feature: Llama Stack connection disrupted And The llama-stack connection is disrupted When I use REST API conversation endpoint with conversation_id from above using HTTP DELETE method Then The status code of the response is 503 - And The body of the response contains Unable to connect to Llama Stack + And The body of the response contains Unable to connect to OGX Scenario: Check conversations/{conversation_id} works when llama-stack is down Given Llama Stack is restarted @@ -276,19 +265,4 @@ Feature: Llama Stack connection disrupted {"question": "How do I list files?"} """ Then The status code of the response is 503 - And The body of the response contains Llama Stack - - - # --- lightspeed-stack-mcp.yaml (aligned with mcp.feature / mcp_servers_api.feature next in test_list) --- - @MCP - Scenario: Register MCP server returns 503 when Llama Stack is unreachable - Given Llama Stack is restarted - And The service uses the lightspeed-stack-mcp.yaml configuration - And The service is restarted - And The llama-stack connection is disrupted - When I access REST API endpoint "mcp-servers" using HTTP POST method - """ - {"name": "unreachable-server", "url": "http://mock-mcp:3000", "provider_id": "model-context-protocol"} - """ - Then The status code of the response is 503 - And The body of the response contains Llama Stack + And The body of the response contains OGX diff --git a/tests/e2e/features/mcp.feature b/tests/e2e/features/mcp.feature index 3d0e42aed..d1c722b81 100644 --- a/tests/e2e/features/mcp.feature +++ b/tests/e2e/features/mcp.feature @@ -11,7 +11,7 @@ Feature: MCP tests # File-based (valid token) — lightspeed-stack-mcp-file-auth.yaml @MCPFileAuthConfig Scenario: Check if tools endpoint succeeds when MCP file-based auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-file-auth.yaml configuration And The service is restarted And The mcp-file mcp server Authorization header is set to "/tmp/mcp-token" @@ -21,7 +21,7 @@ Feature: MCP tests @MCPFileAuthConfig @flaky Scenario: Check if query endpoint succeeds when MCP file-based auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-file-auth.yaml configuration And The service is restarted And The mcp-file mcp server Authorization header is set to "/tmp/mcp-token" @@ -38,7 +38,7 @@ Feature: MCP tests @MCPFileAuthConfig @flaky Scenario: Check if streaming_query endpoint succeeds when MCP file-based auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-file-auth.yaml configuration And The service is restarted And The mcp-file mcp server Authorization header is set to "/tmp/mcp-token" @@ -57,7 +57,7 @@ Feature: MCP tests # File-based (invalid token) — lightspeed-stack-invalid-mcp-file-auth.yaml @InvalidMCPFileAuthConfig Scenario: Check if tools endpoint reports error when MCP file-based invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-invalid-mcp-file-auth.yaml configuration And The service is restarted And The mcp-file mcp server Authorization header is set to "/tmp/invalid-mcp-token" @@ -75,7 +75,7 @@ Feature: MCP tests @InvalidMCPFileAuthConfig Scenario: Check if query endpoint reports error when MCP file-based invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-invalid-mcp-file-auth.yaml configuration And The service is restarted And The mcp-file mcp server Authorization header is set to "/tmp/invalid-mcp-token" @@ -96,7 +96,7 @@ Feature: MCP tests @InvalidMCPFileAuthConfig Scenario: Check if streaming_query endpoint reports error when MCP file-based invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-invalid-mcp-file-auth.yaml configuration And The service is restarted And The mcp-file mcp server Authorization header is set to "/tmp/invalid-mcp-token" @@ -118,7 +118,7 @@ Feature: MCP tests # Kubernetes — lightspeed-stack-mcp-kubernetes-auth.yaml (success paths then invalid token) @MCPKubernetesAuthConfig Scenario: Check if tools endpoint succeeds when MCP kubernetes auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-kubernetes-auth.yaml configuration And The service is restarted And I set the Authorization header to Bearer kubernetes-test-token @@ -128,7 +128,7 @@ Feature: MCP tests @MCPKubernetesAuthConfig @flaky Scenario: Check if query endpoint succeeds when MCP kubernetes auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-kubernetes-auth.yaml configuration And The service is restarted And I set the Authorization header to Bearer kubernetes-test-token @@ -145,7 +145,7 @@ Feature: MCP tests @MCPKubernetesAuthConfig @flaky Scenario: Check if streaming_query endpoint succeeds when MCP kubernetes auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-kubernetes-auth.yaml configuration And The service is restarted And I set the Authorization header to Bearer kubernetes-test-token @@ -163,7 +163,7 @@ Feature: MCP tests @MCPKubernetesAuthConfig Scenario: Check if tools endpoint reports error when MCP kubernetes invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-kubernetes-auth.yaml configuration And The service is restarted And I set the Authorization header to Bearer kubernetes-invalid-token @@ -181,7 +181,7 @@ Feature: MCP tests @MCPKubernetesAuthConfig Scenario: Check if query endpoint reports error when MCP kubernetes invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-kubernetes-auth.yaml configuration And The service is restarted And I set the Authorization header to Bearer kubernetes-invalid-token @@ -202,7 +202,7 @@ Feature: MCP tests @MCPKubernetesAuthConfig Scenario: Check if streaming_query endpoint reports error when MCP kubernetes invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-kubernetes-auth.yaml configuration And The service is restarted And I set the Authorization header to Bearer kubernetes-invalid-token @@ -224,7 +224,7 @@ Feature: MCP tests # Client-provided — lightspeed-stack-mcp-clientauth.yaml @MCPClientAuthConfig Scenario: Check if tools endpoint succeeds when MCP client-provided auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -237,7 +237,7 @@ Feature: MCP tests @MCPClientAuthConfig @flaky Scenario: Check if query endpoint succeeds when MCP client-provided auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -257,7 +257,7 @@ Feature: MCP tests @MCPClientAuthConfig @flaky Scenario: Check if streaming_query endpoint succeeds when MCP client-provided auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -278,7 +278,7 @@ Feature: MCP tests @MCPClientAuthConfig Scenario: Check if tools endpoint succeeds by skipping when MCP client-provided auth token is omitted - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted When I access REST API endpoint "tools" using HTTP GET method @@ -287,7 +287,7 @@ Feature: MCP tests @MCPClientAuthConfig @flaky Scenario: Check if query endpoint succeeds by skipping when MCP client-provided auth token is omitted - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I capture the current token metrics @@ -304,7 +304,7 @@ Feature: MCP tests @MCPClientAuthConfig @flaky Scenario: Check if streaming_query endpoint succeeds by skipping when MCP client-provided auth token is omitted - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I capture the current token metrics @@ -322,7 +322,7 @@ Feature: MCP tests @MCPClientAuthConfig Scenario: Check if tools endpoint reports error when MCP client-provided invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -343,7 +343,7 @@ Feature: MCP tests @MCPClientAuthConfig Scenario: Check if query endpoint reports error when MCP client-provided invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -367,7 +367,7 @@ Feature: MCP tests @MCPClientAuthConfig Scenario: Check if streaming_query endpoint reports error when MCP client-provided invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-client-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -392,7 +392,7 @@ Feature: MCP tests # OAuth — lightspeed-stack-mcp-oauth-auth.yaml (valid token, then unauthenticated, then invalid token) @MCPOAuthAuthConfig Scenario: Check if tools endpoint succeeds when MCP OAuth auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -405,7 +405,7 @@ Feature: MCP tests @MCPOAuthAuthConfig @flaky Scenario: Check if query endpoint succeeds when MCP OAuth auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -425,7 +425,7 @@ Feature: MCP tests @MCPOAuthAuthConfig @flaky Scenario: Check if streaming_query endpoint succeeds when MCP OAuth auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -446,7 +446,7 @@ Feature: MCP tests @MCPOAuthAuthConfig Scenario: Check if tools endpoint reports error when MCP OAuth requires authentication - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted When I access REST API endpoint "tools" using HTTP GET method @@ -464,7 +464,7 @@ Feature: MCP tests @MCPOAuthAuthConfig Scenario: Check if query endpoint reports error when MCP OAuth requires authentication - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted When I use "query" to ask question @@ -485,7 +485,7 @@ Feature: MCP tests @MCPOAuthAuthConfig Scenario: Check if streaming_query endpoint reports error when MCP OAuth requires authentication - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted When I use "streaming_query" to ask question @@ -506,7 +506,7 @@ Feature: MCP tests @MCPOAuthAuthConfig Scenario: Check if tools endpoint reports error when MCP OAuth invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -528,7 +528,7 @@ Feature: MCP tests @MCPOAuthAuthConfig Scenario: Check if query endpoint reports error when MCP OAuth invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -553,7 +553,7 @@ Feature: MCP tests @MCPOAuthAuthConfig Scenario: Check if streaming_query endpoint reports error when MCP OAuth invalid auth token is passed - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp-oauth-auth.yaml configuration And The service is restarted And I set the "MCP-HEADERS" header to @@ -577,7 +577,7 @@ Feature: MCP tests And The headers of the response contains the following header "www-authenticate" Scenario: Check if MCP client auth options endpoint is working - Given MCP toolgroups are reset for a new MCP configuration + Given MCP configuration is reset for a new scenario And The service uses the lightspeed-stack-mcp.yaml configuration And The service is restarted When I access REST API endpoint "mcp-auth/client-options" using HTTP GET method diff --git a/tests/e2e/features/query.feature b/tests/e2e/features/query.feature index 0ad3d9cde..f6a20ef04 100644 --- a/tests/e2e/features/query.feature +++ b/tests/e2e/features/query.feature @@ -266,6 +266,7 @@ Scenario: Check if LLM responds for query request with error for missing query | error | | image | + @skip Scenario: Check if query with shields returns 413 when question is too long for model context When I use "query" to ask question with too-long query and authorization header Then The status code of the response is 413 diff --git a/tests/e2e/features/skills.feature b/tests/e2e/features/skills.feature index a8d86e800..3914a8872 100644 --- a/tests/e2e/features/skills.feature +++ b/tests/e2e/features/skills.feature @@ -12,7 +12,7 @@ Feature: Agent skills tests @SkillsConfig Scenario: Skill tools are registered when skills are configured Given The service uses the lightspeed-stack-skills.yaml configuration - And MCP toolgroups are reset for a new MCP configuration + And MCP configuration is reset for a new scenario And The service is restarted When I access REST API endpoint "tools" using HTTP GET method Then The status code of the response is 200 @@ -24,14 +24,14 @@ Feature: Agent skills tests "identifier": "insert_into_memory", "description": "Insert documents into memory", "parameters": [], - "provider_id": "rag-runtime", - "toolgroup_id": "builtin::rag", + "provider_id": "file-search", + "toolgroup_id": "builtin::file_search", "server_source": "builtin", - "type": "tool_group" + "type": "tool" }, { - "identifier": "knowledge_search", - "description": "Search for information in a database.", + "identifier": "file_search", + "description": "Search files for relevant information", "parameters": [ { "name": "query", @@ -41,10 +41,10 @@ Feature: Agent skills tests "default": null } ], - "provider_id": "rag-runtime", - "toolgroup_id": "builtin::rag", + "provider_id": "file-search", + "toolgroup_id": "builtin::file_search", "server_source": "builtin", - "type": "tool_group" + "type": "tool" }, { "identifier": "list_skills", @@ -140,7 +140,7 @@ Feature: Agent skills tests Scenario: Skill tools are not registered when no skills are configured Given The service uses the lightspeed-stack.yaml configuration - And MCP toolgroups are reset for a new MCP configuration + And MCP configuration is reset for a new scenario And The service is restarted When I access REST API endpoint "tools" using HTTP GET method Then The status code of the response is 200 @@ -152,14 +152,14 @@ Feature: Agent skills tests "identifier": "insert_into_memory", "description": "Insert documents into memory", "parameters": [], - "provider_id": "rag-runtime", - "toolgroup_id": "builtin::rag", + "provider_id": "file-search", + "toolgroup_id": "builtin::file_search", "server_source": "builtin", - "type": "tool_group" + "type": "tool" }, { - "identifier": "knowledge_search", - "description": "Search for information in a database.", + "identifier": "file_search", + "description": "Search files for relevant information", "parameters": [ { "name": "query", @@ -169,10 +169,10 @@ Feature: Agent skills tests "default": null } ], - "provider_id": "rag-runtime", - "toolgroup_id": "builtin::rag", + "provider_id": "file-search", + "toolgroup_id": "builtin::file_search", "server_source": "builtin", - "type": "tool_group" + "type": "tool" } ] } diff --git a/tests/e2e/features/steps/common.py b/tests/e2e/features/steps/common.py index 1c38f43b4..f0b2ccc2f 100644 --- a/tests/e2e/features/steps/common.py +++ b/tests/e2e/features/steps/common.py @@ -6,7 +6,6 @@ from behave import given # pyright: ignore[reportAttributeAccessIssue] from behave.runner import Context -from tests.e2e.utils.llama_stack_utils import unregister_mcp_toolgroups from tests.e2e.utils.utils import ( absolute_repo_path, clear_llama_stack_storage, @@ -87,10 +86,10 @@ def configure_service(context: Context, config_name: str) -> None: state, not only ``context``, so it survives per-scenario context resets), returns immediately: no backup, no copy, and sets ``context.lightspeed_stack_skip_restart`` so the next ``The service is - restarted`` step can no-op—except after ``MCP toolgroups are reset for a new - MCP configuration`` (library ``~/.llama`` clear or server-mode unregister), - in which case the restart is not skipped so Lightspeed reloads config and - Llama MCP state stays consistent. When the basename differs from the last apply, creates the + restarted`` step can no-op—except after ``MCP configuration is reset for a new + scenario`` (library mode clears embedded Llama Stack storage), in which case + the restart is not skipped so Lightspeed reloads config and MCP state stays + consistent. When the basename differs from the last apply, creates the backup on first use, copies the YAML, updates ``context.feature_config`` / override flags, and stores the basename for the next check. Cleared in ``before_feature`` so a @@ -109,15 +108,12 @@ def configure_service(context: Context, config_name: str) -> None: """ config_name = config_name.strip() if _active_lightspeed_stack_config_basename["basename"] == config_name: - # ``MCP toolgroups are reset for a new MCP configuration`` may have run - # (library: clear ``~/.llama``; server: unregister toolgroups). The next - # restart must not be skipped or SQLite handles / MCP registration state - # diverges from the running process. - if getattr( - context, "force_lightspeed_restart_after_mcp_toolgroup_reset", False - ): + # ``MCP configuration is reset for a new scenario`` may have run (library: + # clear ``~/.llama``). The next restart must not be skipped or SQLite + # handles / MCP state diverges from the running process. + if getattr(context, "force_lightspeed_restart_after_mcp_config_reset", False): context.lightspeed_stack_skip_restart = False - context.force_lightspeed_restart_after_mcp_toolgroup_reset = False + context.force_lightspeed_restart_after_mcp_config_reset = False else: context.lightspeed_stack_skip_restart = True return @@ -159,27 +155,26 @@ def configure_service(context: Context, config_name: str) -> None: _active_lightspeed_stack_config_basename["basename"] = config_name context.active_lightspeed_stack_config_basename = config_name context.lightspeed_stack_skip_restart = False - context.force_lightspeed_restart_after_mcp_toolgroup_reset = False + context.force_lightspeed_restart_after_mcp_config_reset = False +@given("MCP configuration is reset for a new scenario") @given("MCP toolgroups are reset for a new MCP configuration") -def reset_mcp_toolgroups_for_new_configuration(context: Context) -> None: - """Clear MCP toolgroups on Llama Stack (server) or ~/.llama storage (library). +def reset_mcp_configuration_for_new_scenario(context: Context) -> None: + """Reset MCP-related state before applying a different MCP config. - Run before applying a different MCP-related ``lightspeed-stack-*.yaml`` in a - scenario so tool registration matches the new config. Sets - ``force_lightspeed_restart_after_mcp_toolgroup_reset`` so the next + Llama Stack 0.7 no longer registers MCP servers as toolgroups. In library + mode, clear embedded Llama Stack storage so the next config applies cleanly. + In server mode, only force a Lightspeed restart on the next config apply. + + Sets ``force_lightspeed_restart_after_mcp_config_reset`` so the next ``The service uses ...`` step cannot skip ``The service is restarted`` when - the YAML basename is unchanged—library mode needs a real restart after - clearing ``~/.llama`` (SQLite); server mode needs it after unregister so - Lightspeed re-registers MCP toolgroups on startup. + the YAML basename is unchanged. """ - context.force_lightspeed_restart_after_mcp_toolgroup_reset = True + context.force_lightspeed_restart_after_mcp_config_reset = True context.lightspeed_stack_skip_restart = False if context.is_library_mode: clear_llama_stack_storage() - else: - unregister_mcp_toolgroups() @given("The service is restarted") diff --git a/tests/e2e/features/steps/info.py b/tests/e2e/features/steps/info.py index 52fed7eb8..2f07c5c43 100644 --- a/tests/e2e/features/steps/info.py +++ b/tests/e2e/features/steps/info.py @@ -58,15 +58,11 @@ def check_shield_structure(context: Context) -> None: # Validate structure and values assert found_shield["type"] == "shield", "type should be 'shield'" assert ( - found_shield["provider_id"] == "llama-guard" - ), "provider_id should be 'llama-guard'" - assert found_shield["provider_resource_id"] == "openai/gpt-4o-mini", ( - f"provider_resource_id should be 'openai/gpt-4o-mini', " - f"but is '{found_shield['provider_resource_id']}'" + found_shield["provider_id"] == "redaction" + ), "provider_id should be 'redaction'" + assert found_shield["name"] == "pii-redaction", ( + f"name should be 'pii-redaction', " f"but is '{found_shield['name']}'" ) - assert ( - found_shield["identifier"] == "llama-guard" - ), f"identifier should be 'llama-guard', but is '{found_shield["identifier"]}'" @then("The response contains {count:d} tools listed for provider {provider_name}") diff --git a/tests/e2e/features/streaming_query.feature b/tests/e2e/features/streaming_query.feature index 33eabea7b..332ed0dde 100644 --- a/tests/e2e/features/streaming_query.feature +++ b/tests/e2e/features/streaming_query.feature @@ -257,6 +257,7 @@ Feature: streaming_query endpoint API tests | error | | image | + @skip Scenario: Check if streaming_query with shields returns 413 when question is too long for model context When I use "streaming_query" to ask question with too-long query and authorization header Then The status code of the response is 413 diff --git a/tests/e2e/rag/README.md b/tests/e2e/rag/README.md index 60db1b797..9b5bbd804 100644 --- a/tests/e2e/rag/README.md +++ b/tests/e2e/rag/README.md @@ -2,6 +2,22 @@ This directory holds committed BYOK vector stores used by the e2e suite. +## OGX 1.0 KV key namespace + +OGX 1.0 FAISS persistence uses +`persistence.namespace: vector_io::faiss` (see `run-ci.yaml` and LCS BYOK +enrichment). SQLite KV keys are therefore stored as: + +```text +vector_io::faiss:vector_stores:v3:: +vector_io::faiss:faiss_index:v3:: +… +``` + +Fixtures committed here must use that prefix. Pre-1.0 un-namespaced keys +(`vector_stores:v3::…`) are invisible to OGX 1.0 and yield empty search +results even when file_search / `vector_io.query` run successfully. + ## `kv_store.db` Faiss BYOK store used by `faiss.feature` and `inline_rag.feature` (the diff --git a/tests/e2e/rag/kv_store.db b/tests/e2e/rag/kv_store.db index 187f497e0..cd892c0ba 100644 Binary files a/tests/e2e/rag/kv_store.db and b/tests/e2e/rag/kv_store.db differ diff --git a/tests/e2e/rag/pdf_kv_store.db b/tests/e2e/rag/pdf_kv_store.db index f56861dc0..f6b98388d 100644 Binary files a/tests/e2e/rag/pdf_kv_store.db and b/tests/e2e/rag/pdf_kv_store.db differ diff --git a/tests/e2e/utils/README.md b/tests/e2e/utils/README.md index 5e51501fd..1ae5059af 100644 --- a/tests/e2e/utils/README.md +++ b/tests/e2e/utils/README.md @@ -7,7 +7,7 @@ Helpers for reading and updating Llama Stack run.yaml across environments. Thin Prow/OpenShift wrappers for Llama Stack run.yaml ConfigMap operations. ## [llama_stack_utils.py](llama_stack_utils.py) -E2E test utilities for Llama Stack (toolgroups and shields). +E2E test utilities for Llama Stack shields. ## [prow_utils.py](prow_utils.py) Prow/OpenShift-specific utility functions for E2E tests. diff --git a/tests/e2e/utils/llama_stack_utils.py b/tests/e2e/utils/llama_stack_utils.py index 70b444170..bfb7d4fe6 100644 --- a/tests/e2e/utils/llama_stack_utils.py +++ b/tests/e2e/utils/llama_stack_utils.py @@ -1,9 +1,8 @@ -"""E2E test utilities for Llama Stack (toolgroups and shields). +"""E2E test utilities for Llama Stack shields. -This module provides functions to manage MCP toolgroups and shields on a running -Llama Stack instance during end-to-end tests: unregister MCP toolgroups when -switching configurations or testing MCP auth, and unregister/re-register shields -(e.g. from the ``Given shields are disabled for this scenario`` step). +This module provides functions to manage shields on a running Llama Stack +instance during end-to-end tests: unregister/re-register shields (e.g. from the +``Given shields are disabled for this scenario`` step). Only applies when running Llama Stack as a separate service (server mode). Requires E2E_LLAMA_STACK_URL or E2E_LLAMA_HOSTNAME and E2E_LLAMA_PORT. @@ -13,17 +12,17 @@ import os from typing import Optional -from llama_stack_client import ( +from ogx_client import ( APIConnectionError, APIStatusError, - AsyncLlamaStackClient, + AsyncOgxClient, ) from tests.e2e.utils.utils import is_prow_environment -def _get_llama_stack_client() -> AsyncLlamaStackClient: - """Build an AsyncLlamaStackClient from env (for e2e test use).""" +def _get_ogx_client() -> AsyncOgxClient: + """Build an AsyncOgxClient from env (for e2e test use).""" base_url = os.getenv("E2E_LLAMA_STACK_URL") if not base_url: if is_prow_environment(): @@ -34,50 +33,7 @@ def _get_llama_stack_client() -> AsyncLlamaStackClient: base_url = f"http://{host}:{port}" api_key = os.getenv("E2E_LLAMA_STACK_API_KEY", "xyzzy") timeout = int(os.getenv("E2E_LLAMA_STACK_TIMEOUT", "60")) - return AsyncLlamaStackClient(base_url=base_url, api_key=api_key, timeout=timeout) - - -# ----------------------------------------------------------------------------- -# Toolgroups -# ----------------------------------------------------------------------------- - - -async def _unregister_toolgroup_async(identifier: str) -> None: - """Unregister a toolgroup by identifier.""" - client = _get_llama_stack_client() - try: - await client.toolgroups.unregister(identifier) - except APIConnectionError: - raise - except APIStatusError as e: - # 400 "not found": toolgroup already absent, scenario can proceed - if e.status_code == 400 and "not found" in str(e).lower(): - return None - raise - finally: - await client.close() - - -async def _unregister_mcp_toolgroups_async() -> None: - """Unregister all MCP toolgroups.""" - client = _get_llama_stack_client() - try: - toolgroups = await client.toolgroups.list() - for toolgroup in toolgroups: - if ( - toolgroup.identifier - and toolgroup.provider_id == "model-context-protocol" - ): - await _unregister_toolgroup_async(toolgroup.identifier) - except APIConnectionError: - raise - finally: - await client.close() - - -def unregister_mcp_toolgroups() -> None: - """Unregister all MCP toolgroups.""" - asyncio.run(_unregister_mcp_toolgroups_async()) + return AsyncOgxClient(base_url=base_url, api_key=api_key, timeout=timeout) # ----------------------------------------------------------------------------- @@ -87,7 +43,7 @@ def unregister_mcp_toolgroups() -> None: async def _unregister_shield_async(identifier: str) -> Optional[tuple[str, str]]: """Unregister a shield by identifier; return (provider_id, provider_shield_id) for restore.""" - client = _get_llama_stack_client() + client = _get_ogx_client() try: shields = await client.shields.list() provider_id = None @@ -126,7 +82,7 @@ async def _register_shield_async( provider_shield_id: str, ) -> None: """Register a shield (restore after unregister).""" - client = _get_llama_stack_client() + client = _get_ogx_client() try: await client.shields.register( shield_id=shield_id, diff --git a/tests/e2e/utils/utils.py b/tests/e2e/utils/utils.py index b23163340..0597c4846 100644 --- a/tests/e2e/utils/utils.py +++ b/tests/e2e/utils/utils.py @@ -409,9 +409,8 @@ def remove_config_backup(backup_path: str) -> None: def clear_llama_stack_storage(container_name: str = "lightspeed-stack") -> None: """Clear Llama Stack storage in library mode (embedded Llama Stack). - Removes the ~/.llama directory so that toolgroups and other persisted - state are reset. Used before MCP config scenarios when not running in - server mode (no separate Llama Stack to unregister toolgroups from). + Removes the ~/.llama directory so embedded Llama Stack persisted state is + reset. Used before MCP config scenarios in library mode. Only runs when using Docker (skipped in Prow). Parameters: diff --git a/tests/integration/README.md b/tests/integration/README.md index ca29f810c..2388fcdea 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -66,7 +66,7 @@ def test_example(mock_request_with_auth: Request) -> None: ### Mocking Fixtures -#### `mock_llama_stack_client` (function-scoped) +#### `mock_ogx_client` (function-scoped) Mocks the external Llama Stack client with sensible defaults: - Returns a mock response with "This is a test response about Ansible." - Mocks `models.list`, `shields.list`, `vector_stores.list` @@ -74,9 +74,9 @@ Mocks the external Llama Stack client with sensible defaults: - Can be customized in individual tests ```python -def test_example(mock_llama_stack_client: Any) -> None: +def test_example(mock_ogx_client: Any) -> None: # Customize the mock for this specific test - mock_llama_stack_client.responses.create.return_value = custom_response + mock_ogx_client.responses.create.return_value = custom_response ``` ## Helper Functions @@ -185,7 +185,7 @@ from configuration import AppConfig @pytest.mark.asyncio async def test_example_endpoint_success( test_config: AppConfig, - mock_llama_stack_client: Any, + mock_ogx_client: Any, test_request: Request, test_auth: AuthTuple, ) -> None: @@ -198,7 +198,7 @@ async def test_example_endpoint_success( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client test_request: FastAPI request test_auth: noop authentication tuple """ @@ -367,9 +367,9 @@ def my_custom_client(mocker): # ... duplicate code # ✅ GOOD - Using common fixture -def test_example(mock_llama_stack_client: Any): +def test_example(mock_ogx_client: Any): # Customize if needed - mock_llama_stack_client.responses.create.return_value = custom_response + mock_ogx_client.responses.create.return_value = custom_response ``` ### 2. Use Test Constants @@ -422,7 +422,7 @@ Include what the test verifies and parameters: @pytest.mark.asyncio async def test_example( test_config: AppConfig, - mock_llama_stack_client: Any, + mock_ogx_client: Any, ) -> None: """Test that example endpoint handles errors correctly. @@ -433,7 +433,7 @@ async def test_example( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client """ ``` diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 7aa24a099..fa5c96444 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -8,8 +8,9 @@ import pytest from fastapi import Request, Response from fastapi.testclient import TestClient -from llama_stack_api.openai_responses import OpenAIResponseObject -from llama_stack_client.types import VersionInfo +from ogx_api.openai_responses import OpenAIResponseObject +from ogx_client.types import ListModelsResponse, VersionInfo +from ogx_client.types.model import Model from pydantic_ai import AgentRunResultEvent from pydantic_ai.messages import ( ModelMessage, @@ -714,8 +715,8 @@ def mock_request_with_auth_fixture() -> Request: return request -@pytest.fixture(name="mock_llama_stack_client") -def mock_llama_stack_client_fixture( +@pytest.fixture(name="mock_ogx_client") +def mock_ogx_client_fixture( mocker: MockerFixture, ) -> Generator[Any, None, None]: """Mock only the external Llama Stack client for integration tests. @@ -724,7 +725,7 @@ def mock_llama_stack_client_fixture( defaults for integration tests. Individual tests can override specific behaviors as needed. - Patches AsyncLlamaStackClientHolder in both app.endpoints.query and app.main + Patches AsyncOgxClientHolder in both app.endpoints.query and app.main to ensure the mock is active during TestClient startup (when app.main imports and initializes the client) and during endpoint execution. @@ -734,13 +735,13 @@ def mock_llama_stack_client_fixture( Yields: mock_client: The mocked Llama Stack client instance. """ - # Patch AsyncLlamaStackClientHolder at multiple import locations + # Patch AsyncOgxClientHolder at multiple import locations # This ensures the mock is active both during app startup (app.main) # and during endpoint execution (query, conversations_v1, responses, etc.) - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") - mocker.patch("app.main.AsyncLlamaStackClientHolder", mock_holder_class) + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") + mocker.patch("app.main.AsyncOgxClientHolder", mock_holder_class) mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder", mock_holder_class + "app.endpoints.conversations_v1.AsyncOgxClientHolder", mock_holder_class ) mock_client = mocker.AsyncMock() @@ -767,13 +768,20 @@ def mock_llama_stack_client_fixture( mock_client.responses.create.return_value = mock_response # Mock models.list - mock_model = mocker.MagicMock() - mock_model.id = "test-provider/test-model" - mock_model.custom_metadata = { - "provider_id": "test-provider", - "model_type": "llm", - } - mock_client.models.list.return_value = [mock_model] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="test-provider/test-model", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "provider_id": "test-provider", + "model_type": "llm", + }, + ) + ] + ) # Mock shields.list (empty by default) mock_client.shields.list.return_value = [] diff --git a/tests/integration/endpoints/test_conversations_v1_integration.py b/tests/integration/endpoints/test_conversations_v1_integration.py index 3e82f3d55..3a1515e33 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 llama_stack_client import APIConnectionError, APIStatusError +from ogx_client import APIConnectionError, APIStatusError from pytest_mock import AsyncMockType, MockerFixture from sqlalchemy.orm import Session @@ -316,7 +316,7 @@ async def test_conversation_validation_errors( async def test_conversation_error_handling( # pylint: disable=too-many-locals test_case: dict, test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -332,7 +332,7 @@ async def test_conversation_error_handling( # pylint: disable=too-many-locals Parameters: test_case: Dictionary containing test parameters test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -360,7 +360,7 @@ async def test_conversation_error_handling( # pylint: disable=too-many-locals patch_db_session.commit() # Configure mock to raise appropriate error - mock_method = mock_llama_stack_client + mock_method = mock_ogx_client for attr in mock_path.split("."): mock_method = getattr(mock_method, attr) @@ -403,7 +403,7 @@ async def test_conversation_error_handling( # pylint: disable=too-many-locals @pytest.mark.asyncio async def test_get_conversation_returns_chat_history( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -419,7 +419,7 @@ async def test_get_conversation_returns_chat_history( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -454,9 +454,7 @@ async def test_get_conversation_returns_chat_history( mock_items = mocker.Mock() mock_items.data = [mock_user_message, mock_assistant_message] mock_items.has_next_page.return_value = False - mock_llama_stack_client.conversations.items.list = mocker.AsyncMock( - return_value=mock_items - ) + mock_ogx_client.conversations.items.list = mocker.AsyncMock(return_value=mock_items) response = await get_conversation_endpoint_handler( request=non_admin_test_request, @@ -485,7 +483,7 @@ async def test_get_conversation_returns_chat_history( @pytest.mark.asyncio async def test_get_conversation_with_turns_metadata( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -500,7 +498,7 @@ async def test_get_conversation_with_turns_metadata( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -546,9 +544,7 @@ async def test_get_conversation_with_turns_metadata( mock_items = mocker.Mock() mock_items.data = [mock_user_message, mock_assistant_message] mock_items.has_next_page.return_value = False - mock_llama_stack_client.conversations.items.list = mocker.AsyncMock( - return_value=mock_items - ) + mock_ogx_client.conversations.items.list = mocker.AsyncMock(return_value=mock_items) response = await get_conversation_endpoint_handler( request=non_admin_test_request, @@ -588,7 +584,7 @@ async def test_get_conversation_with_turns_metadata( @pytest.mark.asyncio async def test_delete_conversation_deletes_from_database_and_llama_stack( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -604,7 +600,7 @@ async def test_delete_conversation_deletes_from_database_and_llama_stack( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -629,7 +625,7 @@ async def test_delete_conversation_deletes_from_database_and_llama_stack( # Mock Llama Stack delete response mock_delete_response = mocker.MagicMock() mock_delete_response.deleted = True - mock_llama_stack_client.conversations.delete.return_value = mock_delete_response + mock_ogx_client.conversations.delete.return_value = mock_delete_response response = await delete_conversation_endpoint_handler( request=non_admin_test_request, @@ -654,7 +650,7 @@ async def test_delete_conversation_deletes_from_database_and_llama_stack( @pytest.mark.asyncio async def test_delete_conversation_handles_not_found_in_llama_stack( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -669,7 +665,7 @@ async def test_delete_conversation_handles_not_found_in_llama_stack( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -692,7 +688,7 @@ async def test_delete_conversation_handles_not_found_in_llama_stack( patch_db_session.commit() # Configure mock to raise not found error - mock_llama_stack_client.conversations.delete.side_effect = APIStatusError( + mock_ogx_client.conversations.delete.side_effect = APIStatusError( message="Not found", response=mocker.Mock(status_code=404), body=None, @@ -721,7 +717,7 @@ async def test_delete_conversation_handles_not_found_in_llama_stack( @pytest.mark.asyncio async def test_delete_conversation_non_existent_returns_success( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -736,7 +732,7 @@ async def test_delete_conversation_non_existent_returns_success( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -748,7 +744,7 @@ async def test_delete_conversation_non_existent_returns_success( # Mock Llama Stack delete response mock_delete_response = mocker.MagicMock() mock_delete_response.deleted = False - mock_llama_stack_client.conversations.delete.return_value = mock_delete_response + mock_ogx_client.conversations.delete.return_value = mock_delete_response response = await delete_conversation_endpoint_handler( request=non_admin_test_request, @@ -769,7 +765,7 @@ async def test_delete_conversation_non_existent_returns_success( @pytest.mark.asyncio async def test_update_conversation_updates_topic_summary( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, non_admin_test_request: Request, test_auth: AuthTuple, patch_db_session: Session, @@ -784,7 +780,7 @@ async def test_update_conversation_updates_topic_summary( Parameters: test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client non_admin_test_request: FastAPI request with standard user permissions test_auth: noop authentication tuple patch_db_session: Test database session @@ -806,7 +802,7 @@ async def test_update_conversation_updates_topic_summary( patch_db_session.commit() # Mock Llama Stack update response - mock_llama_stack_client.conversations.update.return_value = None + mock_ogx_client.conversations.update.return_value = None update_request = ConversationUpdateRequest(topic_summary="New topic summary") diff --git a/tests/integration/endpoints/test_health_integration.py b/tests/integration/endpoints/test_health_integration.py index 9dd2846d0..4a0dafea3 100644 --- a/tests/integration/endpoints/test_health_integration.py +++ b/tests/integration/endpoints/test_health_integration.py @@ -17,8 +17,8 @@ from models.common import HealthStatus -@pytest.fixture(name="mock_llama_stack_client_health") -def mock_llama_stack_client_fixture( +@pytest.fixture(name="mock_ogx_client_health") +def mock_ogx_client_fixture( mocker: MockerFixture, ) -> Generator[Any, None, None]: """Mock only the external Llama Stack client. @@ -30,7 +30,7 @@ def mock_llama_stack_client_fixture( mock_client: An AsyncMock representing the Llama Stack client whose `inspect.version` returns an empty list. """ - mock_holder_class = mocker.patch("app.endpoints.health.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.health.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() # Mock the version endpoint to return a known version @@ -74,7 +74,7 @@ async def test_health_liveness( @pytest.mark.asyncio async def test_health_readiness_provider_statuses( - mock_llama_stack_client_health: AsyncMockType, + mock_ogx_client_health: AsyncMockType, mocker: MockerFixture, ) -> None: """Test that get_providers_health_statuses correctly retrieves and returns @@ -88,11 +88,11 @@ async def test_health_readiness_provider_statuses( Parameters: ---------- - mock_llama_stack_client_health: Mocked Llama Stack client + mock_ogx_client_health: Mocked Llama Stack client mocker: pytest-mock fixture for creating mock objects """ # Arrange: Set up mock provider list with mixed health statuses - mock_llama_stack_client_health.providers.list.return_value = [ + mock_ogx_client_health.providers.list.return_value = [ mocker.Mock( provider_id="unhealthy-provider-1", health={ @@ -150,13 +150,13 @@ async def test_health_readiness_client_error( with pytest.raises(RuntimeError) as exc_info: await readiness_probe_get_method(auth=test_auth, response=test_response) - assert "AsyncLlamaStackClient has not been initialised" in str(exc_info.value) + assert "AsyncOgxClient has not been initialised" in str(exc_info.value) assert "Ensure 'load(..)' has been called" in str(exc_info.value) @pytest.mark.asyncio async def test_health_readiness( - mock_llama_stack_client_health: AsyncMockType, + mock_ogx_client_health: AsyncMockType, test_response: Response, test_auth: AuthTuple, mocker: MockerFixture, @@ -171,7 +171,7 @@ async def test_health_readiness( Parameters: ---------- - mock_llama_stack_client_health: Mocked Llama Stack client + mock_ogx_client_health: Mocked Llama Stack client test_response: FastAPI response object test_auth: noop authentication tuple @@ -179,7 +179,7 @@ async def test_health_readiness( ------- None """ - _ = mock_llama_stack_client_health + _ = mock_ogx_client_health # Mock check_default_model_available since configuration is not loaded mock_check_model = mocker.patch( diff --git a/tests/integration/endpoints/test_info_integration.py b/tests/integration/endpoints/test_info_integration.py index 2b3d46838..bdd546a07 100644 --- a/tests/integration/endpoints/test_info_integration.py +++ b/tests/integration/endpoints/test_info_integration.py @@ -5,8 +5,8 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError -from llama_stack_client.types import VersionInfo +from ogx_client import APIConnectionError +from ogx_client.types import VersionInfo from pytest_mock import AsyncMockType, MockerFixture from app.endpoints.info import info_endpoint_handler @@ -15,8 +15,8 @@ from version import __version__ -@pytest.fixture(name="mock_llama_stack_client") -def mock_llama_stack_client_fixture( +@pytest.fixture(name="mock_ogx_client") +def mock_ogx_client_fixture( mocker: MockerFixture, ) -> Generator[Any, None, None]: """Mock only the external Llama Stack client. @@ -32,7 +32,7 @@ def mock_llama_stack_client_fixture( ------ AsyncMock: A mocked Llama Stack client configured for tests. """ - mock_holder_class = mocker.patch("app.endpoints.info.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.info.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() # Mock the version endpoint to return a known version @@ -48,7 +48,7 @@ def mock_llama_stack_client_fixture( @pytest.mark.asyncio async def test_info_endpoint_returns_service_information( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, test_request: Request, test_auth: AuthTuple, ) -> None: @@ -64,7 +64,7 @@ async def test_info_endpoint_returns_service_information( Parameters: ---------- test_config: Loads real configuration (required for endpoint to access config) - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client test_request: FastAPI request test_auth: noop authentication tuple @@ -83,13 +83,13 @@ async def test_info_endpoint_returns_service_information( assert response.llama_stack_version == "0.2.22" # Verify the Llama Stack client was called - mock_llama_stack_client.inspect.version.assert_called_once() + mock_ogx_client.inspect.version.assert_called_once() @pytest.mark.asyncio async def test_info_endpoint_handles_connection_error( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, test_request: Request, test_auth: AuthTuple, mocker: MockerFixture, @@ -104,7 +104,7 @@ async def test_info_endpoint_handles_connection_error( Parameters: ---------- test_config: Loads real configuration (required for endpoint to access config) - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client test_request: FastAPI request test_auth: noop authentication tuple mocker: pytest-mock fixture for creating mocks @@ -112,7 +112,7 @@ async def test_info_endpoint_handles_connection_error( # test_config fixture loads configuration, which is required for the endpoint _ = test_config # Configure mock to raise connection error - mock_llama_stack_client.inspect.version.side_effect = APIConnectionError( + mock_ogx_client.inspect.version.side_effect = APIConnectionError( request=mocker.Mock() ) @@ -123,7 +123,7 @@ async def test_info_endpoint_handles_connection_error( # Verify error details assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE assert isinstance(exc_info.value.detail, dict) - expected = "Unable to connect to Llama Stack" + expected = "Unable to connect to OGX" assert exc_info.value.detail["response"] == expected # type: ignore[reportArgumentType] assert "cause" in exc_info.value.detail @@ -131,7 +131,7 @@ async def test_info_endpoint_handles_connection_error( @pytest.mark.asyncio async def test_info_endpoint_uses_configuration_values( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, test_request: Request, test_auth: AuthTuple, ) -> None: @@ -145,12 +145,12 @@ async def test_info_endpoint_uses_configuration_values( Parameters: ---------- test_config: Loads real configuration (required for endpoint to access config) - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client test_request: Real FastAPI request test_auth: Real noop authentication tuple """ # Fixtures with side effects (needed but not directly used) - _ = mock_llama_stack_client + _ = mock_ogx_client response = await info_endpoint_handler(auth=test_auth, request=test_request) diff --git a/tests/integration/endpoints/test_model_list.py b/tests/integration/endpoints/test_model_list.py index cf5f0a7d1..47a62b4d5 100644 --- a/tests/integration/endpoints/test_model_list.py +++ b/tests/integration/endpoints/test_model_list.py @@ -6,7 +6,9 @@ import pytest from fastapi import Request from fastapi.exceptions import HTTPException -from llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pytest_mock import AsyncMockType, MockerFixture from app.endpoints.models import models_endpoint_handler @@ -15,8 +17,8 @@ from models.api.requests import ModelFilter -@pytest.fixture(name="mock_llama_stack_client") -def mock_llama_stack_client_fixture( +@pytest.fixture(name="mock_ogx_client") +def mock_ogx_client_fixture( mocker: MockerFixture, ) -> Generator[Any, None, None]: """Mock only the external Llama Stack client. @@ -33,24 +35,35 @@ def mock_llama_stack_client_fixture( mock_client: The mocked Llama Stack client instance configured as described above. """ # Patch in app.endpoints.models where it's actually used by models_endpoint_handler_base - mock_holder_class = mocker.patch("app.endpoints.models.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.models.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() # Mock models list (required for model selection) - mock_model1 = mocker.MagicMock() - mock_model1.id = "test-provider/test-model-1" - mock_model1.custom_metadata = { - "provider_id": "test-provider", - "model_type": "llm", - } - mock_model2 = mocker.MagicMock() - mock_model2.id = "test-provider/test-model-2" - mock_model2.custom_metadata = { - "provider_id": "test-provider", - "model_type": "embedding", - } - mock_client.models.list.return_value = [mock_model1, mock_model2] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="test-provider/test-model-1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "provider_id": "test-provider", + "model_type": "llm", + }, + ), + Model.model_construct( + id="test-provider/test-model-2", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "provider_id": "test-provider", + "model_type": "embedding", + }, + ), + ] + ) # Create a mock holder instance mock_holder_instance = mock_holder_class.return_value @@ -59,8 +72,8 @@ def mock_llama_stack_client_fixture( yield mock_client -@pytest.fixture(name="mock_llama_stack_client_failing") -def mock_llama_stack_client_failing_fixture( +@pytest.fixture(name="mock_ogx_client_failing") +def mock_ogx_client_failing_fixture( mocker: MockerFixture, ) -> Generator[Any, None, None]: """Mock only the external Llama Stack client. @@ -77,7 +90,7 @@ def mock_llama_stack_client_failing_fixture( mock_client: The mocked Llama Stack client instance configured as described above. """ # Patch in app.endpoints.models where it's actually used by models_endpoint_handler_base - mock_holder_class = mocker.patch("app.endpoints.models.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.models.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() @@ -131,7 +144,7 @@ def mock_llama_stack_client_failing_fixture( async def test_models_list_with_filter( test_case: dict, test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, test_request: Request, test_auth: AuthTuple, ) -> None: @@ -147,12 +160,12 @@ async def test_models_list_with_filter( test_case: Dictionary containing test parameters (filter_type, expected_count, expected_models) test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client test_request: FastAPI request test_auth: noop authentication tuple """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client filter_type = test_case["filter_type"] expected_count = test_case["expected_count"] @@ -170,14 +183,14 @@ async def test_models_list_with_filter( # Verify each expected model for i, expected_model in enumerate(expected_models): - assert response.models[i]["identifier"] == expected_model["identifier"] - assert response.models[i]["api_model_type"] == expected_model["api_model_type"] + assert response.models[i].identifier == expected_model["identifier"] + assert response.models[i].api_model_type == expected_model["api_model_type"] @pytest.mark.asyncio async def test_models_list_on_api_connection_error( test_config: AppConfig, - mock_llama_stack_client_failing: AsyncMockType, + mock_ogx_client_failing: AsyncMockType, test_request: Request, test_auth: AuthTuple, ) -> None: @@ -190,12 +203,12 @@ async def test_models_list_on_api_connection_error( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client_failing: Mocked Llama Stack client that raises APIConnectionError + mock_ogx_client_failing: Mocked Llama Stack client that raises APIConnectionError test_request: FastAPI request test_auth: noop authentication tuple """ _ = test_config - _ = mock_llama_stack_client_failing + _ = mock_ogx_client_failing # we should catch HTTPException, not APIConnectionError! with pytest.raises(HTTPException) as exc_info: @@ -207,6 +220,6 @@ async def test_models_list_on_api_connection_error( assert exc_info.value.status_code == 503 assert isinstance(exc_info.value.detail, dict) - expected = "Unable to connect to Llama Stack" + expected = "Unable to connect to OGX" assert exc_info.value.detail["response"] == expected # type: ignore[reportArgumentType] assert "cause" in exc_info.value.detail diff --git a/tests/integration/endpoints/test_query_byok_integration.py b/tests/integration/endpoints/test_query_byok_integration.py index a4e0205aa..faa5c35e4 100644 --- a/tests/integration/endpoints/test_query_byok_integration.py +++ b/tests/integration/endpoints/test_query_byok_integration.py @@ -7,7 +7,8 @@ import pytest from fastapi import Request -from llama_stack_client.types import VersionInfo +from ogx_client.types import ListModelsResponse, VersionInfo +from ogx_client.types.model import Model from pytest_mock import AsyncMockType, MockerFixture import constants @@ -98,13 +99,20 @@ def _build_base_mock_client(mocker: MockerFixture) -> Any: mock_client = mocker.AsyncMock() # Model list - mock_model = mocker.MagicMock() - mock_model.id = "test-provider/test-model" - mock_model.custom_metadata = { - "provider_id": "test-provider", - "model_type": "llm", - } - mock_client.models.list.return_value = [mock_model] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="test-provider/test-model", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "provider_id": "test-provider", + "model_type": "llm", + }, + ) + ] + ) # Shields (empty) mock_client.shields.list.return_value = [] @@ -150,7 +158,7 @@ def mock_byok_client_fixture( output_tokens=20, ) - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) # BYOK vector_io returns results @@ -177,7 +185,7 @@ def mock_byok_tool_rag_client_fixture( Configures vector_stores.list with a BYOK store and agent.run to return a file_search tool result alongside the assistant message. """ - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) # vector_io returns empty (no inline RAG) @@ -429,7 +437,7 @@ async def test_query_byok_inline_rag_with_request_vector_store_ids( test_config.configuration.byok_rag = [entry_a, entry_b] test_config.configuration.rag.inline = ["source-a"] - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) mock_client.vector_io.query = mocker.AsyncMock( @@ -502,7 +510,7 @@ async def test_query_byok_request_vector_store_ids_filters_configured_stores( test_config.configuration.byok_rag = [entry_a, entry_b] test_config.configuration.rag.inline = ["source-a", "source-b"] - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) mock_client.vector_io.query = mocker.AsyncMock( @@ -765,7 +773,7 @@ async def test_query_byok_combined_inline_and_tool_rag( # pylint: disable=too-m test_config.configuration.rag.tool = ["test-knowledge"] # Mock Llama Stack client - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) # Inline RAG returns chunks via vector_io @@ -874,7 +882,7 @@ async def test_query_byok_inline_rag_only_configured_rag_id_is_queried( test_config.configuration.byok_rag = [entry_a, entry_b] test_config.configuration.rag.inline = ["source-a"] - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) mock_client.vector_io.query = mocker.AsyncMock( @@ -960,7 +968,7 @@ async def test_query_byok_score_multiplier_shifts_chunk_priority( # pylint: dis test_config.configuration.byok_rag = [entry_a, entry_b] test_config.configuration.rag.inline = ["source-a", "source-b"] - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) # Source A: high base similarity @@ -1062,7 +1070,7 @@ async def test_query_rag_content_limit_caps_retrieved_results( # pylint: disabl # Disable reranker for this test since it's testing chunk capping, not reranking test_config.configuration.reranker.enabled = False - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) # Generate more chunks than INLINE_RAG_MAX_CHUNKS @@ -1156,7 +1164,7 @@ async def test_query_rag_content_limit_caps_across_multiple_sources( # pylint: test_config.configuration.byok_rag = [entry_a, entry_b] test_config.configuration.rag.inline = ["source-a", "source-b"] - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) # Overlapping score bands so top-k must pick from both sources @@ -1263,7 +1271,7 @@ async def test_query_rag_content_limit_caps_inline_rag( # pylint: disable=too-m test_config.configuration.rag.inline = ["big-source"] test_config.configuration.reranker.enabled = False - mock_holder_class = mocker.patch("app.endpoints.query.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.query.AsyncOgxClientHolder") mock_client = _build_base_mock_client(mocker) num_chunks = constants.BYOK_RAG_MAX_CHUNKS diff --git a/tests/integration/endpoints/test_query_integration.py b/tests/integration/endpoints/test_query_integration.py index a1e29cadf..5cc90b762 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 llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError from pytest_mock import AsyncMockType, MockerFixture from sqlalchemy.orm import Session @@ -33,7 +33,7 @@ SPECIFIC_CONV_ID = "c9d40813-d64d-41eb-8060-3b2446929a02" EXISTING_CONV_ID = "22222222-2222-2222-2222-222222222222" -# Note: mock_llama_stack_client and patch_db_session are now provided by +# Note: mock_ogx_client and patch_db_session are now provided by # tests/integration/conftest.py (patch_db_session is autouse for all tests) # ========================================== @@ -44,7 +44,7 @@ @pytest.mark.asyncio async def test_query_v2_endpoint_successful_response( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -60,13 +60,13 @@ async def test_query_v2_endpoint_successful_response( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent query_request = QueryRequest( @@ -93,7 +93,7 @@ async def test_query_v2_endpoint_successful_response( @pytest.mark.asyncio async def test_query_v2_endpoint_handles_connection_error( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -109,7 +109,7 @@ async def test_query_v2_endpoint_handles_connection_error( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -120,7 +120,7 @@ async def test_query_v2_endpoint_handles_connection_error( None """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent mock_query_agent.run.side_effect = APIConnectionError(request=mocker.Mock()) @@ -138,7 +138,7 @@ async def test_query_v2_endpoint_handles_connection_error( # Verify error details assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE assert isinstance(exc_info.value.detail, dict) - expected = "Unable to connect to Llama Stack" + expected = "Unable to connect to OGX" assert exc_info.value.detail["response"] == expected # type: ignore[reportArgumentType] assert "cause" in exc_info.value.detail @@ -166,7 +166,7 @@ async def test_query_v2_endpoint_handles_connection_error( async def test_query_v2_endpoint_returns_401_for_mcp_oauth( test_case: dict, test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -186,14 +186,14 @@ async def test_query_v2_endpoint_returns_401_for_mcp_oauth( test_case: Dictionary containing test parameters (www_authenticate, expect_www_authenticate) test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client 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 """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent www_authenticate = test_case["www_authenticate"] @@ -237,7 +237,7 @@ async def test_query_v2_endpoint_returns_401_for_mcp_oauth( @pytest.mark.asyncio async def test_query_v2_endpoint_empty_query( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -252,7 +252,7 @@ async def test_query_v2_endpoint_empty_query( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -262,7 +262,7 @@ async def test_query_v2_endpoint_empty_query( None """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent query_request = QueryRequest(query="") @@ -370,7 +370,7 @@ async def test_query_v2_endpoint_empty_query( async def test_query_v2_endpoint_attachment_handling( test_case: dict, test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -390,13 +390,13 @@ async def test_query_v2_endpoint_attachment_handling( test_case: Dictionary containing test parameters (attachments, expected_status, expected_error) test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent attachments = test_case["attachments"] @@ -446,7 +446,7 @@ async def test_query_v2_endpoint_attachment_handling( @pytest.mark.asyncio async def test_query_v2_endpoint_with_tool_calls( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -462,14 +462,14 @@ async def test_query_v2_endpoint_with_tool_calls( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client 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 """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent tool_run_result = create_file_search_agent_run_result( @@ -508,7 +508,7 @@ async def test_query_v2_endpoint_with_tool_calls( @pytest.mark.asyncio async def test_query_v2_endpoint_with_mcp_list_tools( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -524,14 +524,14 @@ async def test_query_v2_endpoint_with_mcp_list_tools( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client 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 """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent mcp_run_result = create_mcp_list_tools_agent_run_result( @@ -569,7 +569,7 @@ async def test_query_v2_endpoint_with_mcp_list_tools( @pytest.mark.asyncio async def test_query_v2_endpoint_with_multiple_tool_types( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -585,14 +585,14 @@ async def test_query_v2_endpoint_with_multiple_tool_types( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client 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 """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent mock_query_agent.run.return_value = create_multi_tool_agent_run_result(mocker) @@ -617,7 +617,7 @@ async def test_query_v2_endpoint_with_multiple_tool_types( @pytest.mark.asyncio async def test_query_v2_endpoint_bypasses_tools_when_no_tools_true( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -635,7 +635,7 @@ async def test_query_v2_endpoint_bypasses_tools_when_no_tools_true( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -655,7 +655,7 @@ async def test_query_v2_endpoint_bypasses_tools_when_no_tools_true( mock_list_result = mocker.MagicMock() mock_list_result.data = [mock_vector_store] - mock_llama_stack_client.vector_stores.list.return_value = mock_list_result + mock_ogx_client.vector_stores.list.return_value = mock_list_result query_request = QueryRequest(query="What is Ansible?", no_tools=True) @@ -677,7 +677,7 @@ async def test_query_v2_endpoint_bypasses_tools_when_no_tools_true( @pytest.mark.asyncio async def test_query_v2_endpoint_uses_tools_when_available( # pylint: disable=unused-argument test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -694,7 +694,7 @@ async def test_query_v2_endpoint_uses_tools_when_available( # pylint: disable=u Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -737,7 +737,7 @@ async def test_query_v2_endpoint_uses_tools_when_available( # pylint: disable=u @pytest.mark.asyncio async def test_query_v2_endpoint_persists_conversation_to_database( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -753,14 +753,14 @@ async def test_query_v2_endpoint_persists_conversation_to_database( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple patch_db_session: Test database session """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent query_request = QueryRequest(query="What is Ansible?") @@ -791,7 +791,7 @@ async def test_query_v2_endpoint_persists_conversation_to_database( @pytest.mark.asyncio async def test_query_v2_endpoint_updates_existing_conversation( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -809,14 +809,14 @@ async def test_query_v2_endpoint_updates_existing_conversation( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple patch_db_session: Test database session """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent # Create an existing conversation in the database @@ -866,7 +866,7 @@ async def test_query_v2_endpoint_updates_existing_conversation( @pytest.mark.asyncio async def test_query_v2_endpoint_conversation_ownership_validation( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -882,14 +882,14 @@ async def test_query_v2_endpoint_conversation_ownership_validation( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple patch_db_session: Test database session """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent # Create conversation owned by the authenticated user in database @@ -922,7 +922,7 @@ async def test_query_v2_endpoint_conversation_ownership_validation( @pytest.mark.asyncio async def test_query_v2_endpoint_creates_valid_cache_entry( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -942,7 +942,7 @@ async def test_query_v2_endpoint_creates_valid_cache_entry( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -954,7 +954,7 @@ async def test_query_v2_endpoint_creates_valid_cache_entry( None """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent _ = patch_db_session @@ -992,7 +992,7 @@ async def test_query_v2_endpoint_creates_valid_cache_entry( @pytest.mark.asyncio async def test_query_v2_endpoint_conversation_not_found_returns_404( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1009,14 +1009,14 @@ async def test_query_v2_endpoint_conversation_not_found_returns_404( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple patch_db_session: Test database session """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent query_request = QueryRequest( @@ -1047,7 +1047,7 @@ async def test_query_v2_endpoint_conversation_not_found_returns_404( @pytest.mark.asyncio async def test_query_v2_endpoint_with_shield_violation( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1069,7 +1069,7 @@ async def test_query_v2_endpoint_with_shield_violation( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -1077,7 +1077,7 @@ async def test_query_v2_endpoint_with_shield_violation( mocker: pytest-mock fixture (only for Llama Stack response) """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent set_query_agent_run( @@ -1110,7 +1110,7 @@ async def test_query_v2_endpoint_with_shield_violation( @pytest.mark.asyncio async def test_query_v2_endpoint_without_shields( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1127,7 +1127,7 @@ async def test_query_v2_endpoint_without_shields( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -1137,7 +1137,7 @@ async def test_query_v2_endpoint_without_shields( _ = patch_db_session # Configure Llama Stack client mock to return no shields (default behavior) - mock_llama_stack_client.shields.list.return_value = [] + mock_ogx_client.shields.list.return_value = [] query_request = QueryRequest(query="What is Ansible?") @@ -1160,7 +1160,7 @@ async def test_query_v2_endpoint_without_shields( @pytest.mark.asyncio async def test_query_v2_endpoint_handles_empty_llm_response( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1177,14 +1177,14 @@ async def test_query_v2_endpoint_handles_empty_llm_response( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client 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 """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent set_query_agent_run( @@ -1217,7 +1217,7 @@ async def test_query_v2_endpoint_handles_empty_llm_response( @pytest.mark.asyncio async def test_query_v2_endpoint_quota_integration( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1235,7 +1235,7 @@ async def test_query_v2_endpoint_quota_integration( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -1243,7 +1243,7 @@ async def test_query_v2_endpoint_quota_integration( mocker: pytest-mock fixture (only for spying on quota functions) """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent _ = patch_db_session @@ -1286,7 +1286,7 @@ async def test_query_v2_endpoint_quota_integration( @pytest.mark.asyncio async def test_query_v2_endpoint_rejects_query_when_quota_exceeded( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1304,7 +1304,7 @@ async def test_query_v2_endpoint_rejects_query_when_quota_exceeded( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple @@ -1316,7 +1316,7 @@ async def test_query_v2_endpoint_rejects_query_when_quota_exceeded( None """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent # Mock check_tokens_available to simulate quota exceeded @@ -1355,7 +1355,7 @@ async def test_query_v2_endpoint_rejects_query_when_quota_exceeded( @pytest.mark.asyncio async def test_query_v2_endpoint_transcript_behavior( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1373,14 +1373,14 @@ async def test_query_v2_endpoint_transcript_behavior( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple patch_db_session: Test database session mocker: pytest-mock fixture """ - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent # Mock store_transcript to prevent file creation mocker.patch("utils.query.store_transcript") @@ -1449,7 +1449,7 @@ async def test_query_v2_endpoint_transcript_behavior( @pytest.mark.asyncio async def test_query_v2_endpoint_uses_conversation_history_model( test_config: AppConfig, - mock_llama_stack_client: AsyncMockType, + mock_ogx_client: AsyncMockType, mock_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -1467,14 +1467,14 @@ async def test_query_v2_endpoint_uses_conversation_history_model( Parameters: ---------- test_config: Test configuration - mock_llama_stack_client: Mocked Llama Stack client + mock_ogx_client: Mocked Llama Stack client mock_query_agent: Mocked Pydantic AI agent for build_agent/agent.run test_request: FastAPI request test_auth: noop authentication tuple patch_db_session: Test database session """ _ = test_config - _ = mock_llama_stack_client + _ = mock_ogx_client _ = mock_query_agent user_id, _, _, _ = test_auth diff --git a/tests/integration/endpoints/test_responses_byok_integration.py b/tests/integration/endpoints/test_responses_byok_integration.py index 8a702dd9e..f0e8ce65c 100644 --- a/tests/integration/endpoints/test_responses_byok_integration.py +++ b/tests/integration/endpoints/test_responses_byok_integration.py @@ -75,12 +75,12 @@ def _build_responses_mock_client(mocker: MockerFixture) -> Any: def _patch_all_client_holders(mocker: MockerFixture, mock_client: Any) -> None: - """Patch AsyncLlamaStackClientHolder in all modules used by the responses endpoint.""" + """Patch AsyncOgxClientHolder in all modules used by the responses endpoint.""" for module in ( "app.endpoints.responses", "utils.endpoints", ): - holder = mocker.patch(f"{module}.AsyncLlamaStackClientHolder") + holder = mocker.patch(f"{module}.AsyncOgxClientHolder") holder.return_value.get_client.return_value = mock_client original_cls = ResponsesContext diff --git a/tests/integration/endpoints/test_responses_integration.py b/tests/integration/endpoints/test_responses_integration.py index 96016e118..808cec3f4 100644 --- a/tests/integration/endpoints/test_responses_integration.py +++ b/tests/integration/endpoints/test_responses_integration.py @@ -11,6 +11,8 @@ import pytest from fastapi import Request from fastapi.responses import StreamingResponse +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pytest_mock import MockerFixture from sqlalchemy.orm import Session @@ -19,6 +21,7 @@ from configuration import AppConfig from models.api.requests import ResponsesRequest from models.api.responses.successful import ResponsesResponse +from models.common.moderation import ShieldModerationBlocked from models.common.responses.contexts import ResponsesContext from models.database.conversations import UserConversation, UserTurn @@ -88,13 +91,20 @@ def _build_mock_client(mocker: MockerFixture) -> Any: mock_response.model_dump.return_value = _RESPONSE_DUMP.copy() mock_client.responses.create = mocker.AsyncMock(return_value=mock_response) - mock_model = mocker.MagicMock() - mock_model.id = "test-provider/test-model" - mock_model.custom_metadata = { - "provider_id": "test-provider", - "model_type": "llm", - } - mock_client.models.list.return_value = [mock_model] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="test-provider/test-model", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "provider_id": "test-provider", + "model_type": "llm", + }, + ) + ] + ) mock_client.shields.list.return_value = [] @@ -110,7 +120,7 @@ def _build_mock_client(mocker: MockerFixture) -> Any: def _patch_client_holders(mocker: MockerFixture, mock_client: Any) -> None: - """Patch AsyncLlamaStackClientHolder in all modules used by the responses endpoint. + """Patch AsyncOgxClientHolder in all modules used by the responses endpoint. Patches three import locations (responses endpoint, utils.endpoints, utils.responses) and bypasses ResponsesContext Pydantic validation. @@ -119,7 +129,7 @@ def _patch_client_holders(mocker: MockerFixture, mock_client: Any) -> None: "app.endpoints.responses", "utils.endpoints", ): - holder = mocker.patch(f"{module}.AsyncLlamaStackClientHolder") + holder = mocker.patch(f"{module}.AsyncOgxClientHolder") holder.return_value.get_client.return_value = mock_client original_cls = ResponsesContext @@ -149,29 +159,22 @@ def _setup_test(mocker: MockerFixture) -> Any: def _configure_shield_blocked( mocker: MockerFixture, - mock_client: Any, moderation_id: str, ) -> None: - """Configure mock client to simulate shield-blocked moderation. + """Configure stub moderation to return a blocked result. Args: mocker: pytest-mock fixture. - mock_client: The mock Llama Stack client to configure. moderation_id: The moderation ID for the blocked response. """ - mock_shield = mocker.MagicMock() - mock_shield.identifier = "test-shield" - mock_shield.provider_resource_id = "test-shield-model" - mock_shield.provider_id = "test-shield-provider" - mock_client.shields.list.return_value = [mock_shield] - - mock_moderation = mocker.MagicMock() - mock_moderation.id = moderation_id - mock_result = mocker.MagicMock() - mock_result.flagged = True - mock_result.user_message = "Content blocked by safety shield" - mock_moderation.results = [mock_result] - mock_client.moderations.create = mocker.AsyncMock(return_value=mock_moderation) + blocked = ShieldModerationBlocked( + message="Content blocked by safety shield", + moderation_id=moderation_id, + ) + mocker.patch( + "app.endpoints.responses.run_shield_moderation_v2", + return_value=blocked, + ) @pytest.mark.asyncio @@ -236,7 +239,7 @@ async def test_shield_blocked_persists_moderation_turn( """Test shield-blocked response persists moderation ID and skips last_response_id.""" _ = test_config mock_client = _setup_test(mocker) - _configure_shield_blocked(mocker, mock_client, "modr_blocked_integ_123") + _configure_shield_blocked(mocker, "modr_blocked_integ_123") request = ResponsesRequest( input="Some blocked content", @@ -319,7 +322,7 @@ async def test_streaming_blocked_returns_sse_and_persists_turn( """Test that shield-blocked streaming returns valid SSE events and persists to DB.""" _ = test_config mock_client = _setup_test(mocker) - _configure_shield_blocked(mocker, mock_client, "modr_stream_blocked_123") + _configure_shield_blocked(mocker, "modr_stream_blocked_123") request = ResponsesRequest( input="Some blocked content", diff --git a/tests/integration/endpoints/test_rlsapi_v1_integration.py b/tests/integration/endpoints/test_rlsapi_v1_integration.py index 9830cb40f..3fdf37e86 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 llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError from pytest_mock import MockerFixture import constants @@ -85,7 +85,7 @@ def mock_authorization_fixture(mocker: MockerFixture) -> None: def mock_shield_passed_fixture(mocker: MockerFixture) -> None: """Mock shield moderation to pass for all integration tests.""" mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=ShieldModerationPassed()), ) @@ -126,9 +126,7 @@ def _setup_responses_mock( mock_client = mocker.Mock() mock_client.responses = mock_responses - mock_holder_class = mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder" - ) + mock_holder_class = mocker.patch("app.endpoints.rlsapi_v1.AsyncOgxClientHolder") mock_holder_class.return_value.get_client.return_value = mock_client return mock_client @@ -264,7 +262,7 @@ async def test_rlsapi_v1_infer_connection_error_returns_503( test_auth: AuthTuple, mocker: MockerFixture, ) -> None: - """Test /v1/infer returns 503 when Llama Stack is unavailable.""" + """Test /v1/infer returns 503 when OGX is unavailable.""" _ = rlsapi_config mock_responses = mocker.Mock() @@ -275,9 +273,7 @@ async def test_rlsapi_v1_infer_connection_error_returns_503( mock_client = mocker.Mock() mock_client.responses = mock_responses - mock_holder_class = mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder" - ) + mock_holder_class = mocker.patch("app.endpoints.rlsapi_v1.AsyncOgxClientHolder") mock_holder_class.return_value.get_client.return_value = mock_client with pytest.raises(HTTPException) as exc_info: @@ -292,7 +288,7 @@ async def test_rlsapi_v1_infer_connection_error_returns_503( assert isinstance(exc_info.value.detail, dict) assert "response" in exc_info.value.detail detail = cast(dict[str, str], exc_info.value.detail) - assert "Llama Stack" in detail["response"] + assert "OGX" in detail["response"] @pytest.mark.asyncio @@ -318,9 +314,7 @@ async def test_rlsapi_v1_infer_fallback_response_empty_output( mock_client = mocker.Mock() mock_client.responses = mock_responses - mock_holder_class = mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder" - ) + mock_holder_class = mocker.patch("app.endpoints.rlsapi_v1.AsyncOgxClientHolder") mock_holder_class.return_value.get_client.return_value = mock_client response = await infer_endpoint( @@ -361,9 +355,7 @@ async def test_rlsapi_v1_infer_input_source_combination( mock_client = mocker.Mock() mock_client.responses = mock_responses - mock_holder_class = mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder" - ) + mock_holder_class = mocker.patch("app.endpoints.rlsapi_v1.AsyncOgxClientHolder") mock_holder_class.return_value.get_client.return_value = mock_client await infer_endpoint( @@ -424,9 +416,7 @@ async def test_rlsapi_v1_infer_no_mcp_servers_passes_empty_tools( mock_client = mocker.Mock() mock_client.responses = mock_responses - mock_holder_class = mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder" - ) + mock_holder_class = mocker.patch("app.endpoints.rlsapi_v1.AsyncOgxClientHolder") mock_holder_class.return_value.get_client.return_value = mock_client mocker.patch( @@ -469,9 +459,7 @@ async def test_rlsapi_v1_infer_mcp_tools_passed_to_llm( mock_client = mocker.Mock() mock_client.responses = mock_responses - mock_holder_class = mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder" - ) + mock_holder_class = mocker.patch("app.endpoints.rlsapi_v1.AsyncOgxClientHolder") mock_holder_class.return_value.get_client.return_value = mock_client mcp_tools = [ diff --git a/tests/integration/endpoints/test_root_endpoint.py b/tests/integration/endpoints/test_root_endpoint.py index 192678773..e894e289b 100644 --- a/tests/integration/endpoints/test_root_endpoint.py +++ b/tests/integration/endpoints/test_root_endpoint.py @@ -5,7 +5,7 @@ import pytest from fastapi import Request, status -from llama_stack_client.types import VersionInfo +from ogx_client.types import VersionInfo from pytest_mock import MockerFixture from app.endpoints.root import root_endpoint_handler @@ -13,8 +13,8 @@ from configuration import AppConfig -@pytest.fixture(name="mock_llama_stack_client") -def mock_llama_stack_client_fixture( +@pytest.fixture(name="mock_ogx_client") +def mock_ogx_client_fixture( mocker: MockerFixture, ) -> Generator[Any, None, None]: """Mock only the external Llama Stack client. @@ -30,7 +30,7 @@ def mock_llama_stack_client_fixture( ------ AsyncMock: A mocked Llama Stack client configured for tests. """ - mock_holder_class = mocker.patch("app.endpoints.info.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.info.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() # Mock the version endpoint to return a known version diff --git a/tests/integration/endpoints/test_streaming_query_byok_integration.py b/tests/integration/endpoints/test_streaming_query_byok_integration.py index 1a371fdb8..5db1ffff4 100644 --- a/tests/integration/endpoints/test_streaming_query_byok_integration.py +++ b/tests/integration/endpoints/test_streaming_query_byok_integration.py @@ -99,7 +99,7 @@ def mock_streaming_byok_client_fixture( ) mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -152,7 +152,7 @@ def mock_streaming_byok_tool_client_fixture( # pylint: disable=too-many-stateme ) mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -298,7 +298,7 @@ async def test_streaming_query_byok_inline_rag_with_request_vector_store_ids( test_config.configuration.rag.inline = ["source-a"] mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -363,7 +363,7 @@ async def test_streaming_query_byok_request_vector_store_ids_filters_configured_ test_config.configuration.rag.inline = ["source-a", "source-b"] mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -641,7 +641,7 @@ async def test_streaming_query_byok_combined_inline_and_tool_rag( # Mock Llama Stack client mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -719,7 +719,7 @@ async def test_streaming_query_byok_only_configured_rag_id_is_queried( test_config.configuration.rag.inline = ["source-a"] mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -797,7 +797,7 @@ async def test_streaming_query_byok_score_multiplier_shifts_priority( # pylint: test_config.configuration.rag.inline = ["source-a", "source-b"] mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -878,7 +878,7 @@ async def test_streaming_query_rag_content_limit_caps_context( # pylint: disabl test_config.configuration.rag.inline = ["big-source"] mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -953,7 +953,7 @@ async def test_streaming_query_rag_content_limit_caps_across_multiple_sources( test_config.configuration.rag.inline = ["source-a", "source-b"] mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) @@ -1041,7 +1041,7 @@ async def test_streaming_query_rag_content_limit_caps_inline_rag( # pylint: dis test_config.configuration.reranker.enabled = False mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = _build_base_streaming_mock_client(mocker) diff --git a/tests/integration/endpoints/test_streaming_query_integration.py b/tests/integration/endpoints/test_streaming_query_integration.py index efe09b72d..21f8df39e 100644 --- a/tests/integration/endpoints/test_streaming_query_integration.py +++ b/tests/integration/endpoints/test_streaming_query_integration.py @@ -7,6 +7,8 @@ from fastapi import HTTPException, Request, status from fastapi.responses import StreamingResponse from fastapi.testclient import TestClient +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pytest_mock import AsyncMockType, MockerFixture from app.endpoints.streaming_query import streaming_query_endpoint_handler @@ -16,7 +18,7 @@ from models.common.query import Attachment -@pytest.fixture(name="mock_streaming_llama_stack_client") +@pytest.fixture(name="mock_streaming_ogx_client") def mock_llama_stack_streaming_fixture( mocker: MockerFixture, mock_streaming_query_agent: AsyncMockType, @@ -29,17 +31,24 @@ def mock_llama_stack_streaming_fixture( """ _ = mock_streaming_query_agent mock_holder_class = mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder" + "app.endpoints.streaming_query.AsyncOgxClientHolder" ) mock_client = mocker.AsyncMock() - mock_model = mocker.MagicMock() - mock_model.id = "test-provider/test-model" - mock_model.custom_metadata = { - "provider_id": "test-provider", - "model_type": "llm", - } - mock_client.models.list.return_value = [mock_model] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="test-provider/test-model", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "provider_id": "test-provider", + "model_type": "llm", + }, + ) + ] + ) mock_vector_stores_response = mocker.MagicMock() mock_vector_stores_response.data = [] @@ -146,7 +155,7 @@ async def _responses_create(**_kwargs: Any) -> Any: async def test_streaming_query_v2_endpoint_attachment_handling( # pylint: disable=too-many-arguments,too-many-positional-arguments test_case: dict, test_config: AppConfig, - mock_streaming_llama_stack_client: AsyncMockType, + mock_streaming_ogx_client: AsyncMockType, mock_streaming_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -164,13 +173,13 @@ async def test_streaming_query_v2_endpoint_attachment_handling( # pylint: disab test_case: Dictionary containing test parameters (attachments, expected_status, expected_error) test_config: Test configuration - mock_streaming_llama_stack_client: Mocked Llama Stack client + mock_streaming_ogx_client: Mocked Llama Stack client mock_streaming_query_agent: Mocked Pydantic AI agent for build_agent test_request: FastAPI request test_auth: noop authentication tuple """ _ = test_config - _ = mock_streaming_llama_stack_client + _ = mock_streaming_ogx_client _ = mock_streaming_query_agent attachments = test_case["attachments"] @@ -250,7 +259,7 @@ def test_streaming_query_v2_endpoint_empty_body_returns_422( async def test_streaming_query_endpoint_returns_401_for_mcp_oauth( # pylint: disable=too-many-arguments,too-many-positional-arguments test_case: dict, test_config: AppConfig, - mock_streaming_llama_stack_client: Any, + mock_streaming_ogx_client: Any, mock_streaming_query_agent: AsyncMockType, test_request: Request, test_auth: AuthTuple, @@ -269,14 +278,14 @@ async def test_streaming_query_endpoint_returns_401_for_mcp_oauth( # pylint: di test_case: Dictionary containing test parameters (www_authenticate, expect_www_authenticate) test_config: Test configuration - mock_streaming_llama_stack_client: Mocked Llama Stack client + mock_streaming_ogx_client: Mocked Llama Stack client mock_streaming_query_agent: Mocked Pydantic AI agent for build_agent test_request: FastAPI request test_auth: noop authentication tuple mocker: pytest-mock fixture """ _ = test_config - _ = mock_streaming_llama_stack_client + _ = mock_streaming_ogx_client _ = mock_streaming_query_agent www_authenticate = test_case["www_authenticate"] diff --git a/tests/integration/endpoints/test_tools_integration.py b/tests/integration/endpoints/test_tools_integration.py index f013bb20c..5b4cd3853 100644 --- a/tests/integration/endpoints/test_tools_integration.py +++ b/tests/integration/endpoints/test_tools_integration.py @@ -21,7 +21,7 @@ def mock_llama_stack_tools_fixture( Returns: Mock client with toolgroups.list and tools.list configured. """ - mock_holder_class = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") + mock_holder_class = mocker.patch("app.endpoints.tools.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() mock_holder_class.return_value.get_client.return_value = mock_client yield mock_client diff --git a/tests/integration/test_openapi_json.py b/tests/integration/test_openapi_json.py index aa93eaf78..a6d31e1f5 100644 --- a/tests/integration/test_openapi_json.py +++ b/tests/integration/test_openapi_json.py @@ -215,7 +215,7 @@ def test_servers_section_present_from_url(spec_from_url: dict[str, Any]) -> None ("/v1/info", "get", {"200", "401", "403", "503"}), ("/v1/models", "get", {"200", "401", "403", "500", "503"}), ("/v1/tools", "get", {"200", "401", "403", "500", "503"}), - ("/v1/shields", "get", {"200", "401", "403", "500", "503"}), + ("/v1/shields", "get", {"200", "401", "403", "500"}), ("/v1/providers", "get", {"200", "401", "403", "500", "503"}), ( "/v1/providers/{provider_id}", @@ -231,13 +231,13 @@ def test_servers_section_present_from_url(spec_from_url: dict[str, Any]) -> None ( "/v1/mcp-servers", "post", - {"201", "401", "403", "409", "500", "503"}, + {"201", "401", "403", "409", "500"}, ), ("/v1/mcp-servers", "get", {"200", "401", "403", "500"}), ( "/v1/mcp-servers/{name}", "delete", - {"200", "401", "403", "500", "503"}, + {"200", "401", "403", "500"}, ), ("/v1/query", "post", {"200", "401", "403", "404", "422", "429", "500", "503"}), ( @@ -331,7 +331,7 @@ def test_paths_and_responses_exist_from_file( ("/v1/info", "get", {"200", "401", "403", "503"}), ("/v1/models", "get", {"200", "401", "403", "500", "503"}), ("/v1/tools", "get", {"200", "401", "403", "500", "503"}), - ("/v1/shields", "get", {"200", "401", "403", "500", "503"}), + ("/v1/shields", "get", {"200", "401", "403", "500"}), ("/v1/providers", "get", {"200", "401", "403", "500", "503"}), ( "/v1/providers/{provider_id}", @@ -347,13 +347,13 @@ def test_paths_and_responses_exist_from_file( ( "/v1/mcp-servers", "post", - {"201", "401", "403", "409", "500", "503"}, + {"201", "401", "403", "409", "500"}, ), ("/v1/mcp-servers", "get", {"200", "401", "403", "500"}), ( "/v1/mcp-servers/{name}", "delete", - {"200", "401", "403", "500", "503"}, + {"200", "401", "403", "500"}, ), ("/v1/query", "post", {"200", "401", "403", "404", "422", "429", "500", "503"}), ( diff --git a/tests/unit/app/endpoints/test_a2a.py b/tests/unit/app/endpoints/test_a2a.py index c317e0bba..308e185f8 100644 --- a/tests/unit/app/endpoints/test_a2a.py +++ b/tests/unit/app/endpoints/test_a2a.py @@ -22,7 +22,8 @@ ) from a2a.utils import new_agent_text_message from fastapi import HTTPException, Request -from llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError +from ogx_client.types import ListModelsResponse from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError from pydantic_ai.messages import ( @@ -724,7 +725,7 @@ async def test_process_task_streaming_handles_api_connection_error_on_models_lis request=mock_request, ) mocker.patch( - "app.endpoints.a2a.AsyncLlamaStackClientHolder" + "app.endpoints.a2a.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client # prepare_responses_params raises HTTPException when APIConnectionError occurs @@ -735,7 +736,7 @@ async def test_process_task_streaming_handles_api_connection_error_on_models_lis assert exc_info.value.status_code == 503 # Verify error detail contains helpful info - assert "Unable to connect to Llama Stack" in str(exc_info.value.detail) + assert "Unable to connect to OGX" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_process_task_streaming_handles_api_connection_error( # pylint: disable=too-many-locals @@ -776,9 +777,11 @@ async def test_process_task_streaming_handles_api_connection_error( # pylint: d # Mock the client mock_client = mocker.AsyncMock() mock_models = [mocker.MagicMock()] - mock_client.models.list = mocker.AsyncMock(return_value=mock_models) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct(data=mock_models) + ) mocker.patch( - "app.endpoints.a2a.AsyncLlamaStackClientHolder" + "app.endpoints.a2a.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client # Mock prepare_responses_params @@ -860,9 +863,11 @@ async def test_process_task_streaming_handles_agent_run_error( # pylint: disabl ) mock_client = mocker.AsyncMock() - mock_client.models.list = mocker.AsyncMock(return_value=[mocker.MagicMock()]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct(data=[mocker.MagicMock()]) + ) mocker.patch( - "app.endpoints.a2a.AsyncLlamaStackClientHolder" + "app.endpoints.a2a.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client mock_responses_params = mocker.Mock() @@ -935,9 +940,11 @@ async def test_process_task_streaming_applies_compaction( # pylint: disable=too ) mock_client = mocker.AsyncMock() - mock_client.models.list = mocker.AsyncMock(return_value=[mocker.MagicMock()]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct(data=[mocker.MagicMock()]) + ) mocker.patch( - "app.endpoints.a2a.AsyncLlamaStackClientHolder" + "app.endpoints.a2a.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client mock_params = mocker.Mock() diff --git a/tests/unit/app/endpoints/test_conversations.py b/tests/unit/app/endpoints/test_conversations.py index c5a185e17..3a634bc8c 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 llama_stack_client import APIConnectionError, APIStatusError, NotFoundError +from ogx_client import APIConnectionError, APIStatusError, NotFoundError from pytest_mock import MockerFixture, MockType from sqlalchemy.exc import SQLAlchemyError @@ -533,7 +533,7 @@ async def test_llama_stack_connection_error( request=None # type: ignore[arg-type] ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -550,7 +550,7 @@ async def test_llama_stack_connection_error( detail = exc_info.value.detail assert isinstance(detail, dict) response = detail["response"] # pyright: ignore[reportArgumentType] - assert response == "Unable to connect to Llama Stack" + assert response == "Unable to connect to OGX" @pytest.mark.asyncio async def test_llama_stack_not_found_error( @@ -584,7 +584,7 @@ async def test_llama_stack_not_found_error( body=None, ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -691,7 +691,7 @@ async def test_get_others_conversations_allowed_for_authorized_user( ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client response = await get_conversation_endpoint_handler( @@ -745,7 +745,7 @@ async def test_successful_conversation_retrieval( mock_client.conversations.items.list = mocker.AsyncMock(return_value=mock_items) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -776,7 +776,7 @@ async def test_retrieve_conversation_returns_none( mock_database_session(mocker, query_result=[]) mock_client = mocker.AsyncMock() mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -822,7 +822,7 @@ async def test_no_items_found_in_get_conversation( return_value=mock_items_response ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -870,7 +870,7 @@ async def test_api_status_error_in_get_conversation( body=None, ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -970,7 +970,7 @@ def query_side_effect(model_class: type[Any]) -> Any: ] mock_client.conversations.items.list.return_value = mock_items_response mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1112,7 +1112,7 @@ async def test_llama_stack_connection_error( request=None # type: ignore ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1127,7 +1127,7 @@ async def test_llama_stack_connection_error( detail = exc_info.value.detail assert isinstance(detail, dict) response = detail["response"] # pyright: ignore[reportArgumentType] - assert response == "Unable to connect to Llama Stack" + assert response == "Unable to connect to OGX" @pytest.mark.asyncio async def test_llama_stack_not_found_error( @@ -1156,7 +1156,7 @@ async def test_llama_stack_not_found_error( body=None, ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1258,7 +1258,7 @@ async def test_delete_others_conversations_allowed_for_authorized_user( mock_delete_response.deleted = True mock_client.conversations.delete.return_value = mock_delete_response mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1300,7 +1300,7 @@ async def test_successful_conversation_deletion( mock_delete_response.deleted = True mock_client.conversations.delete.return_value = mock_delete_response mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1342,7 +1342,7 @@ async def test_retrieve_conversation_returns_none_in_delete( mock_delete_response.deleted = True mock_client.conversations.delete.return_value = mock_delete_response mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1384,7 +1384,7 @@ async def test_sqlalchemy_error_in_delete( mock_client.agents.session.list.return_value = mock_session_list_response mock_client.agents.session.delete.return_value = None mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -1986,11 +1986,11 @@ async def test_successful_conversation_update( return_value=mock_session_context, ) - # Mock AsyncLlamaStackClientHolder + # Mock AsyncOgxClientHolder mock_client = mocker.AsyncMock() mock_client.conversations.update.return_value = None mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -2030,13 +2030,13 @@ async def test_llama_stack_connection_error_in_update( return_value=mock_conversation, ) - # Mock AsyncLlamaStackClientHolder to raise APIConnectionError + # Mock AsyncOgxClientHolder to raise APIConnectionError mock_client = mocker.AsyncMock() mock_client.conversations.update.side_effect = APIConnectionError( request=None # type: ignore ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -2054,7 +2054,7 @@ async def test_llama_stack_connection_error_in_update( detail = exc_info.value.detail assert isinstance(detail, dict) response = detail["response"] # pyright: ignore[reportArgumentType] - assert response == "Unable to connect to Llama Stack" + assert response == "Unable to connect to OGX" @pytest.mark.asyncio async def test_llama_stack_not_found_error_in_update( @@ -2076,7 +2076,7 @@ async def test_llama_stack_not_found_error_in_update( return_value=mock_conversation, ) - # Mock AsyncLlamaStackClientHolder to raise APIStatusError + # Mock AsyncOgxClientHolder to raise APIStatusError mock_client = mocker.AsyncMock() mock_client.conversations.update.side_effect = APIStatusError( message="Conversation not found", @@ -2084,7 +2084,7 @@ async def test_llama_stack_not_found_error_in_update( body=None, ) mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client @@ -2124,11 +2124,11 @@ async def test_sqlalchemy_error_in_database_update( return_value=mock_conversation, ) - # Mock AsyncLlamaStackClientHolder - update succeeds + # Mock AsyncOgxClientHolder - update succeeds mock_client = mocker.AsyncMock() mock_client.conversations.update.return_value = None mock_client_holder = mocker.patch( - "app.endpoints.conversations_v1.AsyncLlamaStackClientHolder" + "app.endpoints.conversations_v1.AsyncOgxClientHolder" ) mock_client_holder.return_value.get_client.return_value = mock_client diff --git a/tests/unit/app/endpoints/test_health.py b/tests/unit/app/endpoints/test_health.py index e7e9160aa..d32561153 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 llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError from pytest_mock import MockerFixture from app.endpoints.health import ( @@ -210,7 +210,7 @@ async def test_get_providers_health_statuses(self, mocker: MockerFixture) -> Non - unhealthy_provider: status ERROR, message "Connection failed" """ # Mock the imports - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") + mock_lsc = mocker.patch("client.AsyncOgxClientHolder.get_client") # Mock the client and its methods mock_client = mocker.AsyncMock() @@ -264,9 +264,9 @@ async def test_get_providers_health_statuses_connection_error( ) -> None: """Test get_providers_health_statuses when connection fails.""" # Mock the imports - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") + mock_lsc = mocker.patch("client.AsyncOgxClientHolder.get_client") - # Mock get_llama_stack_client to raise an exception + # Mock get_ogx_client to raise an exception mock_lsc.side_effect = APIConnectionError(request=mocker.Mock()) result = await get_providers_health_statuses() @@ -329,7 +329,7 @@ async def test_delegates_to_client_holder( mocker: MockerFixture, ) -> None: """Test delegates to client holder with correct model ID.""" - mock_holder = mocker.patch("app.endpoints.health.AsyncLlamaStackClientHolder") + mock_holder = mocker.patch("app.endpoints.health.AsyncOgxClientHolder") mock_holder.return_value.check_model_available = mocker.AsyncMock( return_value=(True, f"Model {self.EXPECTED_MODEL_ID} is available") ) @@ -349,7 +349,7 @@ async def test_returns_holder_failure( mocker: MockerFixture, ) -> None: """Test passes through failure result from client holder.""" - mock_holder = mocker.patch("app.endpoints.health.AsyncLlamaStackClientHolder") + mock_holder = mocker.patch("app.endpoints.health.AsyncOgxClientHolder") mock_holder.return_value.check_model_available = mocker.AsyncMock( return_value=( False, diff --git a/tests/unit/app/endpoints/test_info.py b/tests/unit/app/endpoints/test_info.py index 71ffafec4..7fe55ee08 100644 --- a/tests/unit/app/endpoints/test_info.py +++ b/tests/unit/app/endpoints/test_info.py @@ -4,8 +4,8 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError -from llama_stack_client.types import VersionInfo +from ogx_client import APIConnectionError +from ogx_client.types import VersionInfo from pytest_mock import MockerFixture from app.endpoints.info import info_endpoint_handler @@ -48,7 +48,7 @@ async def test_info_endpoint(mocker: MockerFixture) -> None: # Mock the LlamaStack client mock_client = mocker.AsyncMock() mock_client.inspect.version.return_value = VersionInfo(version="0.1.2") - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") + mock_lsc = mocker.patch("client.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() mocker.patch("app.endpoints.models.configuration", mock_config) @@ -120,7 +120,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_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") + mock_lsc = mocker.patch("client.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() mocker.patch("app.endpoints.models.configuration", mock_config) @@ -144,4 +144,4 @@ async def test_info_endpoint_connection_error(mocker: MockerFixture) -> None: await info_endpoint_handler(auth=auth, request=request) assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE assert e.value.detail["response"] == "Service unavailable" # type: ignore - assert "Unable to connect to Llama Stack" in e.value.detail["cause"] # type: ignore + assert "Unable to connect to OGX" in e.value.detail["cause"] # type: ignore diff --git a/tests/unit/app/endpoints/test_mcp_servers.py b/tests/unit/app/endpoints/test_mcp_servers.py index ce868c869..595935811 100644 --- a/tests/unit/app/endpoints/test_mcp_servers.py +++ b/tests/unit/app/endpoints/test_mcp_servers.py @@ -7,7 +7,6 @@ import pytest from fastapi import HTTPException, status -from llama_stack_client import APIConnectionError, NotFoundError from pydantic import AnyHttpUrl, SecretStr from pytest_mock import MockerFixture @@ -35,7 +34,7 @@ @pytest.fixture def mock_configuration() -> Configuration: - """Create a mock configuration with MCP servers.""" + """Create a mock configuration with one static MCP server.""" return Configuration( name="test", service=ServiceConfiguration( @@ -95,126 +94,59 @@ def _make_app_config(mocker: MockerFixture, config: Configuration) -> AppConfig: return app_config -def _mock_client(mocker: MockerFixture) -> Any: - """Create and patch a mock Llama Stack client.""" - mock_holder = mocker.patch("app.endpoints.mcp_servers.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_holder.return_value.get_client.return_value = mock_client - return mock_client - - @pytest.mark.asyncio async def test_register_mcp_server_success( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Test successful MCP server registration.""" + """Register a dynamic MCP server in local configuration only.""" app_config = _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None body = MCPServerRegistrationRequest( name="new-mcp-server", - url="http://localhost:8888/mcp", - provider_id="MCP provider ID", - ) - - result = await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH + url="http://localhost:4000", + provider_id="model-context-protocol", ) + request = mocker.Mock() - assert isinstance(result, MCPServerRegistrationResponse) - assert result.name == "new-mcp-server" - assert result.url == "http://localhost:8888/mcp" - assert result.provider_id == "MCP provider ID" - assert "registered successfully" in result.message - - client.toolgroups.register.assert_called_once_with( - toolgroup_id="new-mcp-server", - provider_id="MCP provider ID", - mcp_endpoint={"uri": "http://localhost:8888/mcp"}, + response = await mcp_servers.register_mcp_server_handler( + request=request, + body=body, + auth=MOCK_AUTH, ) + assert isinstance(response, MCPServerRegistrationResponse) + assert response.name == "new-mcp-server" + assert response.url == "http://localhost:4000" + assert response.provider_id == "model-context-protocol" + assert "registered successfully" in response.message assert app_config.is_dynamic_mcp_server("new-mcp-server") - assert any(s.name == "new-mcp-server" for s in app_config.mcp_servers) + assert any(server.name == "new-mcp-server" for server in app_config.mcp_servers) @pytest.mark.asyncio -async def test_register_mcp_server_duplicate_name( +async def test_register_mcp_server_conflict( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Test registration fails when name already exists.""" + """Return 409 when the MCP server name already exists.""" _make_app_config(mocker, mock_configuration) - _mock_client(mocker) body = MCPServerRegistrationRequest( name="static-mcp", - url="http://localhost:9999/mcp", - provider_id="MCP provider ID", - ) - - with pytest.raises(HTTPException) as exc_info: - await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH - ) - assert exc_info.value.status_code == 409 - - -@pytest.mark.asyncio -async def test_register_mcp_server_llama_stack_failure( - mocker: MockerFixture, - mock_configuration: Configuration, -) -> None: - """Test registration rolls back on Llama Stack connection failure.""" - app_config = _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.side_effect = APIConnectionError(request=mocker.Mock()) - - body = MCPServerRegistrationRequest( - name="failing-server", - url="http://localhost:8888/mcp", - provider_id="MCP provider ID", - ) - - with pytest.raises(HTTPException) as exc_info: - await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH - ) - assert exc_info.value.status_code == 503 - - assert not app_config.is_dynamic_mcp_server("failing-server") - assert not any(s.name == "failing-server" for s in app_config.mcp_servers) - - -@pytest.mark.asyncio -async def test_register_mcp_server_toolgroup_failure_returns_500( - mocker: MockerFixture, - mock_configuration: Configuration, -) -> None: - """Registration returns generic 500 without leaking toolgroup exception text.""" - app_config = _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.side_effect = RuntimeError("upstream secret detail") - - body = MCPServerRegistrationRequest( - name="boom-server", - url="http://localhost:8888/mcp", - provider_id="MCP provider ID", + url="http://localhost:4000", + provider_id="model-context-protocol", ) + request = mocker.Mock() with pytest.raises(HTTPException) as exc_info: await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH + request=request, + body=body, + auth=MOCK_AUTH, ) - assert exc_info.value.status_code == 500 - raw_detail = exc_info.value.detail - assert isinstance(raw_detail, dict) - detail: dict[str, Any] = raw_detail - assert detail["response"] == "Failed to register MCP server" - assert not app_config.is_dynamic_mcp_server("boom-server") - assert not any(s.name == "boom-server" for s in app_config.mcp_servers) + assert exc_info.value.status_code == status.HTTP_409_CONFLICT @pytest.mark.asyncio @@ -224,8 +156,6 @@ async def test_register_mcp_server_with_all_fields( ) -> None: """Test registration with all optional fields provided.""" _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None body = MCPServerRegistrationRequest( name="full-mcp-server", @@ -244,12 +174,39 @@ async def test_register_mcp_server_with_all_fields( assert result.provider_id == "custom-provider" +@pytest.mark.asyncio +async def test_list_mcp_servers( + mocker: MockerFixture, + mock_configuration: Configuration, +) -> None: + """List MCP servers from local configuration.""" + app_config = _make_app_config(mocker, mock_configuration) + app_config.add_mcp_server( + ModelContextProtocolServer( + name="dynamic-mcp", + provider_id="model-context-protocol", + url="http://localhost:4001", + ) + ) + + request = mocker.Mock() + response = await mcp_servers.list_mcp_servers_handler( + request=request, + auth=MOCK_AUTH, + ) + + assert isinstance(response, MCPServerListResponse) + sources = {server.name: server.source for server in response.servers} + assert sources["static-mcp"] == "config" + assert sources["dynamic-mcp"] == "api" + + @pytest.mark.asyncio async def test_list_mcp_servers_empty( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Test listing servers returns static servers.""" + """Test listing servers returns static servers only.""" _make_app_config(mocker, mock_configuration) result = await mcp_servers.list_mcp_servers_handler( @@ -263,84 +220,52 @@ async def test_list_mcp_servers_empty( @pytest.mark.asyncio -async def test_list_mcp_servers_with_dynamic( +async def test_delete_mcp_server_static_forbidden( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Test listing shows both static and dynamic servers.""" + """Reject deletion of statically configured MCP servers.""" _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None - - body = MCPServerRegistrationRequest( - name="dynamic-server", - url="http://localhost:9999/mcp", - provider_id="MCP provider ID", - ) - await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH - ) + request = mocker.Mock() - result = await mcp_servers.list_mcp_servers_handler( - request=mocker.Mock(), auth=MOCK_AUTH - ) + with pytest.raises(HTTPException) as exc_info: + await mcp_servers.delete_mcp_server_handler( + request=request, + name="static-mcp", + auth=MOCK_AUTH, + ) - assert len(result.servers) == 2 - sources = {s.name: s.source for s in result.servers} - assert sources["static-mcp"] == "config" - assert sources["dynamic-server"] == "api" + assert exc_info.value.status_code == status.HTTP_403_FORBIDDEN @pytest.mark.asyncio -async def test_delete_dynamic_mcp_server_success( +async def test_delete_mcp_server_dynamic_success( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Test successful deletion of a dynamically registered server.""" + """Delete a dynamic MCP server from local configuration only.""" app_config = _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None - client.toolgroups.unregister.return_value = None - - body = MCPServerRegistrationRequest( - name="to-delete", - url="http://localhost:7777/mcp", - provider_id="MCP provider ID", - ) - await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH + app_config.add_mcp_server( + ModelContextProtocolServer( + name="dynamic-mcp", + provider_id="model-context-protocol", + url="http://localhost:4001", + ) ) - assert app_config.is_dynamic_mcp_server("to-delete") - result = await mcp_servers.delete_mcp_server_handler( - request=mocker.Mock(), name="to-delete", auth=MOCK_AUTH + request = mocker.Mock() + response = await mcp_servers.delete_mcp_server_handler( + request=request, + name="dynamic-mcp", + auth=MOCK_AUTH, ) - assert isinstance(result, MCPServerDeleteResponse) - assert result.name == "to-delete" - assert result.deleted is True - assert result.response == "MCP server deleted successfully" - - assert not app_config.is_dynamic_mcp_server("to-delete") - assert not any(s.name == "to-delete" for s in app_config.mcp_servers) - - client.toolgroups.unregister.assert_called_once_with(toolgroup_id="to-delete") - - -@pytest.mark.asyncio -async def test_delete_static_mcp_server_forbidden( - mocker: MockerFixture, - mock_configuration: Configuration, -) -> None: - """Test that deleting a statically configured server is forbidden.""" - _make_app_config(mocker, mock_configuration) - _mock_client(mocker) - - with pytest.raises(HTTPException) as exc_info: - await mcp_servers.delete_mcp_server_handler( - request=mocker.Mock(), name="static-mcp", auth=MOCK_AUTH - ) - assert exc_info.value.status_code == 403 + assert isinstance(response, MCPServerDeleteResponse) + assert response.deleted is True + assert response.name == "dynamic-mcp" + assert response.response == "MCP server deleted successfully" + assert not app_config.is_dynamic_mcp_server("dynamic-mcp") + assert not any(server.name == "dynamic-mcp" for server in app_config.mcp_servers) @pytest.mark.asyncio @@ -350,7 +275,6 @@ async def test_delete_nonexistent_mcp_server( ) -> None: """Deleting an unknown name returns 200 with deleted=False (idempotent delete).""" _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) result = await mcp_servers.delete_mcp_server_handler( request=mocker.Mock(), name="no-such-server", auth=MOCK_AUTH @@ -360,92 +284,61 @@ async def test_delete_nonexistent_mcp_server( assert result.name == "no-such-server" assert result.deleted is False assert result.response == "MCP server not found" - client.toolgroups.unregister.assert_called_once_with(toolgroup_id="no-such-server") @pytest.mark.asyncio -async def test_delete_mcp_server_llama_stack_failure( +async def test_list_mcp_servers_configuration_not_loaded( mocker: MockerFixture, - mock_configuration: Configuration, ) -> None: - """Test deletion handles Llama Stack connection failure gracefully.""" - _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None - client.toolgroups.unregister.side_effect = APIConnectionError(request=mocker.Mock()) - - body = MCPServerRegistrationRequest( - name="to-delete-fail", - url="http://localhost:7777/mcp", - provider_id="MCP provider ID", - ) - await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH - ) + """Test listing MCP servers returns 500 when configuration is not loaded.""" + mock_config = AppConfig() + mock_config._configuration = None # pylint: disable=protected-access + mocker.patch("app.endpoints.mcp_servers.configuration", mock_config) + mocker.patch("app.endpoints.mcp_servers.authorize", lambda _: lambda func: func) with pytest.raises(HTTPException) as exc_info: - await mcp_servers.delete_mcp_server_handler( - request=mocker.Mock(), name="to-delete-fail", auth=MOCK_AUTH + await mcp_servers.list_mcp_servers_handler( + request=mocker.Mock(), auth=MOCK_AUTH ) - assert exc_info.value.status_code == 503 + + assert exc_info.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR + raw_detail = exc_info.value.detail + assert isinstance(raw_detail, dict) + detail: dict[str, Any] = raw_detail + assert detail["response"] == "Configuration is not loaded" @pytest.mark.asyncio -async def test_delete_mcp_server_toolgroup_not_found_is_idempotent( +async def test_register_and_delete_roundtrip( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Deleting a dynamic server succeeds when Llama Stack toolgroup is already gone.""" - app_config = _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None - client.toolgroups.unregister.side_effect = NotFoundError( - message="Toolgroup not found", - response=mocker.Mock(request=None), - body=None, - ) + """Test full register -> list -> delete -> list cycle.""" + _make_app_config(mocker, mock_configuration) body = MCPServerRegistrationRequest( - name="orphan-server", - url="http://localhost:7777/mcp", + name="roundtrip-server", + url="http://localhost:5555/mcp", provider_id="MCP provider ID", ) await mcp_servers.register_mcp_server_handler( request=mocker.Mock(), body=body, auth=MOCK_AUTH ) - assert app_config.is_dynamic_mcp_server("orphan-server") - result = await mcp_servers.delete_mcp_server_handler( - request=mocker.Mock(), name="orphan-server", auth=MOCK_AUTH + list_result = await mcp_servers.list_mcp_servers_handler( + request=mocker.Mock(), auth=MOCK_AUTH ) + assert len(list_result.servers) == 2 - assert isinstance(result, MCPServerDeleteResponse) - assert result.name == "orphan-server" - assert result.deleted is True - assert not app_config.is_dynamic_mcp_server("orphan-server") - assert not any(s.name == "orphan-server" for s in app_config.mcp_servers) - - -@pytest.mark.asyncio -async def test_list_mcp_servers_configuration_not_loaded( - mocker: MockerFixture, -) -> None: - """Test listing MCP servers returns 500 when configuration is not loaded.""" - mock_config = AppConfig() - mock_config._configuration = None # pylint: disable=protected-access - mocker.patch("app.endpoints.mcp_servers.configuration", mock_config) - mocker.patch("app.endpoints.mcp_servers.authorize", lambda _: lambda func: func) - - with pytest.raises(HTTPException) as exc_info: - await mcp_servers.list_mcp_servers_handler( - request=mocker.Mock(), auth=MOCK_AUTH - ) + await mcp_servers.delete_mcp_server_handler( + request=mocker.Mock(), name="roundtrip-server", auth=MOCK_AUTH + ) - assert exc_info.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR - raw_detail = exc_info.value.detail - assert isinstance(raw_detail, dict) - detail: dict[str, Any] = raw_detail - assert detail["response"] == "Configuration is not loaded" + list_result = await mcp_servers.list_mcp_servers_handler( + request=mocker.Mock(), auth=MOCK_AUTH + ) + assert len(list_result.servers) == 1 + assert list_result.servers[0].name == "static-mcp" def test_mcp_server_registration_request_validation() -> None: @@ -504,39 +397,3 @@ def test_mcp_server_registration_rejects_arbitrary_value() -> None: authorization_headers={"Authorization": "Bearer my-static-token"}, provider_id="MCP provider ID", ) - - -@pytest.mark.asyncio -async def test_register_and_delete_roundtrip( - mocker: MockerFixture, - mock_configuration: Configuration, -) -> None: - """Test full register -> list -> delete -> list cycle.""" - _make_app_config(mocker, mock_configuration) - client = _mock_client(mocker) - client.toolgroups.register.return_value = None - client.toolgroups.unregister.return_value = None - - body = MCPServerRegistrationRequest( - name="roundtrip-server", - url="http://localhost:5555/mcp", - provider_id="MCP provider ID", - ) - await mcp_servers.register_mcp_server_handler( - request=mocker.Mock(), body=body, auth=MOCK_AUTH - ) - - list_result = await mcp_servers.list_mcp_servers_handler( - request=mocker.Mock(), auth=MOCK_AUTH - ) - assert len(list_result.servers) == 2 - - await mcp_servers.delete_mcp_server_handler( - request=mocker.Mock(), name="roundtrip-server", auth=MOCK_AUTH - ) - - list_result = await mcp_servers.list_mcp_servers_handler( - request=mocker.Mock(), auth=MOCK_AUTH - ) - assert len(list_result.servers) == 1 - assert list_result.servers[0].name == "static-mcp" diff --git a/tests/unit/app/endpoints/test_models.py b/tests/unit/app/endpoints/test_models.py index a7a9bc7fd..ed245076a 100644 --- a/tests/unit/app/endpoints/test_models.py +++ b/tests/unit/app/endpoints/test_models.py @@ -4,7 +4,9 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError +from ogx_client import APIConnectionError +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pytest_mock import MockerFixture from pytest_subtests import SubTests @@ -15,17 +17,18 @@ from tests.unit.utils.auth_helpers import mock_authorization_resolvers -# pylint: disable=R0903 -class Model: - """Model information returned in response.""" - - def __init__(self, model_id: str, provider_id: str, model_type: str) -> None: - """Initialize model information.""" - self.id = model_id - self.custom_metadata = { +def _make_model(model_id: str, provider_id: str, model_type: str) -> Model: + """Build an OGX Model for models-endpoint tests.""" + return Model.model_construct( + id=model_id, + created=0, + owned_by="test", + object="model", + custom_metadata={ "model_type": model_type, "provider_id": provider_id, - } + }, + ) @pytest.mark.asyncio @@ -67,10 +70,10 @@ async def test_models_endpoint_handler_configuration_loaded( the Llama Stack client cannot connect. Loads an AppConfig from a test dictionary, patches the endpoint's - configuration and AsyncLlamaStackClientHolder so that get_client raises + configuration and AsyncOgxClientHolder so that get_client raises APIConnectionError, 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 Llama Stack". + and a detail response of "Unable to connect to OGX". """ mock_authorization_resolvers(mocker) @@ -101,9 +104,7 @@ async def test_models_endpoint_handler_configuration_loaded( cfg.init_from_dict(config_dict) mocker.patch("app.endpoints.models.configuration", cfg) - mock_client_holder = mocker.patch( - "app.endpoints.models.AsyncLlamaStackClientHolder" - ) + mock_client_holder = mocker.patch("app.endpoints.models.AsyncOgxClientHolder") mock_client_holder.return_value.get_client.side_effect = APIConnectionError( request=mocker.Mock() ) @@ -123,7 +124,7 @@ async def test_models_endpoint_handler_configuration_loaded( request=request, auth=auth, model_type=ModelFilter(model_type=None) ) assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE - assert e.value.detail["response"] == "Unable to connect to Llama Stack" # type: ignore + assert e.value.detail["response"] == "Unable to connect to OGX" # type: ignore @pytest.mark.asyncio @@ -161,10 +162,8 @@ async def test_models_endpoint_handler_unable_to_retrieve_models_list( # Mock the LlamaStack client mock_client = mocker.AsyncMock() - mock_client.models.list.return_value = [] - mock_lsc = mocker.patch( - "app.endpoints.models.AsyncLlamaStackClientHolder.get_client" - ) + mock_client.models.list.return_value = ListModelsResponse.model_construct(data=[]) + mock_lsc = mocker.patch("app.endpoints.models.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() mocker.patch("app.endpoints.models.configuration", mock_config) @@ -220,10 +219,8 @@ async def test_models_endpoint_handler_model_type_query_parameter( # Mock the LlamaStack client mock_client = mocker.AsyncMock() - mock_client.models.list.return_value = [] - mock_lsc = mocker.patch( - "app.endpoints.models.AsyncLlamaStackClientHolder.get_client" - ) + mock_client.models.list.return_value = ListModelsResponse.model_construct(data=[]) + mock_lsc = mocker.patch("app.endpoints.models.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() mocker.patch("app.endpoints.models.configuration", mock_config) @@ -278,15 +275,15 @@ async def test_models_endpoint_handler_model_list_retrieved( # Mock the LlamaStack client mock_client = mocker.AsyncMock() - mock_client.models.list.return_value = [ - Model("model1", "provider1", "llm"), - Model("model2", "provider2", "embedding"), - Model("model3", "provider3", "llm"), - Model("model4", "provider4", "embedding"), - ] - mock_lsc = mocker.patch( - "app.endpoints.models.AsyncLlamaStackClientHolder.get_client" + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + _make_model("model1", "provider1", "llm"), + _make_model("model2", "provider2", "embedding"), + _make_model("model3", "provider3", "llm"), + _make_model("model4", "provider4", "embedding"), + ] ) + mock_lsc = mocker.patch("app.endpoints.models.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() mocker.patch("app.endpoints.models.configuration", mock_config) @@ -306,14 +303,14 @@ async def test_models_endpoint_handler_model_list_retrieved( ) assert response is not None assert len(response.models) == 4 - assert response.models[0]["identifier"] == "model1" - assert response.models[0]["model_type"] == "llm" - assert response.models[1]["identifier"] == "model2" - assert response.models[1]["model_type"] == "embedding" - assert response.models[2]["identifier"] == "model3" - assert response.models[2]["model_type"] == "llm" - assert response.models[3]["identifier"] == "model4" - assert response.models[3]["model_type"] == "embedding" + assert response.models[0].identifier == "model1" + assert response.models[0].model_type == "llm" + assert response.models[1].identifier == "model2" + assert response.models[1].model_type == "embedding" + assert response.models[2].identifier == "model3" + assert response.models[2].model_type == "llm" + assert response.models[3].identifier == "model4" + assert response.models[3].model_type == "embedding" @pytest.mark.asyncio @@ -352,15 +349,15 @@ async def test_models_endpoint_handler_model_list_retrieved_with_query_parameter # Mock the LlamaStack client mock_client = mocker.AsyncMock() - mock_client.models.list.return_value = [ - Model("model1", "provider1", "llm"), - Model("model2", "provider2", "embedding"), - Model("model3", "provider3", "llm"), - Model("model4", "provider4", "embedding"), - ] - mock_lsc = mocker.patch( - "app.endpoints.models.AsyncLlamaStackClientHolder.get_client" + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + _make_model("model1", "provider1", "llm"), + _make_model("model2", "provider2", "embedding"), + _make_model("model3", "provider3", "llm"), + _make_model("model4", "provider4", "embedding"), + ] ) + mock_lsc = mocker.patch("app.endpoints.models.AsyncOgxClientHolder.get_client") mock_lsc.return_value = mock_client mock_config = mocker.Mock() mocker.patch("app.endpoints.models.configuration", mock_config) @@ -381,10 +378,10 @@ async def test_models_endpoint_handler_model_list_retrieved_with_query_parameter ) assert response is not None assert len(response.models) == 2 - assert response.models[0]["identifier"] == "model1" - assert response.models[0]["model_type"] == "llm" - assert response.models[1]["identifier"] == "model3" - assert response.models[1]["model_type"] == "llm" + assert response.models[0].identifier == "model1" + assert response.models[0].model_type == "llm" + assert response.models[1].identifier == "model3" + assert response.models[1].model_type == "llm" with subtests.test(msg="Model type = 'embedding'"): response = await models_endpoint_handler( @@ -392,10 +389,10 @@ async def test_models_endpoint_handler_model_list_retrieved_with_query_parameter ) assert response is not None assert len(response.models) == 2 - assert response.models[0]["identifier"] == "model2" - assert response.models[0]["model_type"] == "embedding" - assert response.models[1]["identifier"] == "model4" - assert response.models[1]["model_type"] == "embedding" + assert response.models[0].identifier == "model2" + assert response.models[0].model_type == "embedding" + assert response.models[1].identifier == "model4" + assert response.models[1].model_type == "embedding" with subtests.test(msg="Model type = 'xyzzy'"): response = await models_endpoint_handler( @@ -443,13 +440,11 @@ async def test_models_endpoint_llama_stack_connection_error( "authentication": {"module": "noop"}, } - # mock AsyncLlamaStackClientHolder to raise APIConnectionError + # mock AsyncOgxClientHolder to raise APIConnectionError # when models.list() method is called mock_client = mocker.AsyncMock() mock_client.models.list.side_effect = APIConnectionError(request=None) # type: ignore - mock_client_holder = mocker.patch( - "app.endpoints.models.AsyncLlamaStackClientHolder" - ) + mock_client_holder = mocker.patch("app.endpoints.models.AsyncOgxClientHolder") mock_client_holder.return_value.get_client.return_value = mock_client cfg = AppConfig() @@ -470,5 +465,5 @@ async def test_models_endpoint_llama_stack_connection_error( request=request, auth=auth, model_type=ModelFilter(model_type=None) ) assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE - assert e.value.detail["response"] == "Unable to connect to Llama Stack" # type: ignore - assert "Unable to connect to Llama Stack" in e.value.detail["cause"] # type: ignore + assert e.value.detail["response"] == "Unable to connect to OGX" # type: ignore + assert "Unable to connect to OGX" in e.value.detail["cause"] # type: ignore diff --git a/tests/unit/app/endpoints/test_prompts.py b/tests/unit/app/endpoints/test_prompts.py index 6cfb3b709..c24b7870f 100644 --- a/tests/unit/app/endpoints/test_prompts.py +++ b/tests/unit/app/endpoints/test_prompts.py @@ -4,8 +4,8 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError, BadRequestError -from llama_stack_client.types.prompt import Prompt +from ogx_client import APIConnectionError, BadRequestError +from ogx_client.types.prompt import Prompt from pytest_mock import MockerFixture from app.endpoints.prompts import ( @@ -75,7 +75,7 @@ def prompts_client_mocks_fixture( mock_client = mocker.AsyncMock() mock_client.prompts = mock_prompts mocker.patch( - "app.endpoints.prompts.AsyncLlamaStackClientHolder.get_client", + "app.endpoints.prompts.AsyncOgxClientHolder.get_client", return_value=mock_client, ) return mock_client, mock_prompts diff --git a/tests/unit/app/endpoints/test_providers.py b/tests/unit/app/endpoints/test_providers.py index 9905ed045..87237ee7b 100644 --- a/tests/unit/app/endpoints/test_providers.py +++ b/tests/unit/app/endpoints/test_providers.py @@ -2,8 +2,8 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError, BadRequestError -from llama_stack_client.types import ProviderInfo +from ogx_client import APIConnectionError, BadRequestError +from ogx_client.types import ProviderInfo from pytest_mock import MockerFixture from app.endpoints.providers import ( @@ -43,7 +43,7 @@ async def test_providers_endpoint_connection_error( mocker.patch("app.endpoints.providers.configuration", minimal_config) mocker.patch( - "app.endpoints.providers.AsyncLlamaStackClientHolder" + "app.endpoints.providers.AsyncOgxClientHolder" ).return_value.get_client.side_effect = APIConnectionError(request=mocker.Mock()) request = Request(scope={"type": "http"}) @@ -56,7 +56,7 @@ async def test_providers_endpoint_connection_error( assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE detail = e.value.detail assert isinstance(detail, dict) - assert detail["response"] == "Unable to connect to Llama Stack" # type: ignore + assert detail["response"] == "Unable to connect to OGX" # type: ignore @pytest.mark.asyncio @@ -92,7 +92,7 @@ async def test_providers_endpoint_success( mock_client = mocker.AsyncMock() mock_client.providers.list.return_value = provider_list mocker.patch( - "app.endpoints.providers.AsyncLlamaStackClientHolder" + "app.endpoints.providers.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -113,10 +113,8 @@ async def test_get_provider_not_found( """Test that /providers/{provider_id} endpoint raises HTTP 404 if the provider is not found.""" mocker.patch("app.endpoints.providers.configuration", minimal_config) - # Mock AsyncLlamaStackClientHolder to return a client that raises BadRequestError - mock_client_holder = mocker.patch( - "app.endpoints.providers.AsyncLlamaStackClientHolder" - ) + # Mock AsyncOgxClientHolder to return a client that raises BadRequestError + mock_client_holder = mocker.patch("app.endpoints.providers.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() mock_client.providers.retrieve = mocker.AsyncMock( side_effect=BadRequestError( @@ -160,7 +158,7 @@ async def test_get_provider_success( mock_client = mocker.AsyncMock() mock_client.providers.retrieve = mocker.AsyncMock(return_value=provider) mocker.patch( - "app.endpoints.providers.AsyncLlamaStackClientHolder" + "app.endpoints.providers.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -184,7 +182,7 @@ async def test_get_provider_connection_error( mock_authorization_resolvers(mocker) mocker.patch( - "app.endpoints.providers.AsyncLlamaStackClientHolder" + "app.endpoints.providers.AsyncOgxClientHolder" ).return_value.get_client.side_effect = APIConnectionError(request=mocker.Mock()) request = Request(scope={"type": "http"}) @@ -199,4 +197,4 @@ async def test_get_provider_connection_error( assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE detail = e.value.detail assert isinstance(detail, dict) - assert detail["response"] == "Unable to connect to Llama Stack" # type: ignore + assert detail["response"] == "Unable to connect to OGX" # type: ignore diff --git a/tests/unit/app/endpoints/test_query.py b/tests/unit/app/endpoints/test_query.py index 4c67fe64a..902ad1701 100644 --- a/tests/unit/app/endpoints/test_query.py +++ b/tests/unit/app/endpoints/test_query.py @@ -1,16 +1,14 @@ -# pylint: disable=redefined-outer-name, import-error,too-many-locals,too-many-lines -# pyright: reportCallIssue=false +# pylint: disable=too-many-locals """Unit tests for the /query (v2) REST API endpoint using Responses API.""" from typing import Any import pytest -from fastapi import HTTPException, Request -from llama_stack_api.openai_responses import OpenAIResponseObject -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from fastapi import Request +from ogx_client import AsyncOgxClient from pytest_mock import MockerFixture -from app.endpoints.query import query_endpoint_handler, retrieve_response +from app.endpoints.query import query_endpoint_handler from configuration import AppConfig from models.api.requests import QueryRequest from models.api.responses.successful import QueryResponse @@ -21,12 +19,9 @@ RAGChunk, RAGContext, ReferencedDocument, - ToolCallSummary, - ToolResultSummary, TurnSummary, ) from models.database.conversations import UserConversation -from utils.token_counter import TokenCounter # User ID must be proper UUID MOCK_AUTH = ( @@ -115,7 +110,7 @@ async def test_successful_query_no_conversation( mocker.patch("app.endpoints.query.check_tokens_available") mocker.patch("app.endpoints.query.validate_model_provider_override") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_response_obj = mocker.Mock() mock_response_obj.output = [] mock_client.responses = mocker.Mock() @@ -123,7 +118,7 @@ async def test_successful_query_no_conversation( mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.query.AsyncLlamaStackClientHolder", + "app.endpoints.query.AsyncOgxClientHolder", return_value=mock_client_holder, ) mocker.patch( @@ -198,7 +193,7 @@ async def test_query_merges_inline_and_tool_rag_chunks_and_documents( mocker.patch("app.endpoints.query.check_tokens_available") mocker.patch("app.endpoints.query.validate_model_provider_override") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_response_obj = mocker.Mock() mock_response_obj.output = [] mock_client.responses = mocker.Mock() @@ -206,7 +201,7 @@ async def test_query_merges_inline_and_tool_rag_chunks_and_documents( mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.query.AsyncLlamaStackClientHolder", + "app.endpoints.query.AsyncOgxClientHolder", return_value=mock_client_holder, ) mocker.patch( @@ -295,11 +290,11 @@ async def test_successful_query_with_conversation( return_value=mocker.Mock(spec=UserConversation), ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.query.AsyncLlamaStackClientHolder", + "app.endpoints.query.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -364,7 +359,7 @@ async def test_query_with_attachments( "app.endpoints.query.validate_attachments_metadata" ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_response_obj = mocker.Mock() mock_response_obj.output = [] mock_client.responses = mocker.Mock() @@ -372,7 +367,7 @@ async def test_query_with_attachments( mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.query.AsyncLlamaStackClientHolder", + "app.endpoints.query.AsyncOgxClientHolder", return_value=mock_client_holder, ) mocker.patch( @@ -439,11 +434,11 @@ async def test_query_with_topic_summary( mocker.patch("app.endpoints.query.check_tokens_available") mocker.patch("app.endpoints.query.validate_model_provider_override") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.query.AsyncLlamaStackClientHolder", + "app.endpoints.query.AsyncOgxClientHolder", return_value=mock_client_holder, ) mocker.patch( @@ -505,7 +500,7 @@ async def test_query_azure_token_refresh( mocker.patch("app.endpoints.query.check_tokens_available") mocker.patch("app.endpoints.query.validate_model_provider_override") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_response_obj = mocker.Mock() mock_response_obj.output = [] mock_client.responses = mocker.Mock() @@ -513,7 +508,7 @@ async def test_query_azure_token_refresh( mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.query.AsyncLlamaStackClientHolder", + "app.endpoints.query.AsyncOgxClientHolder", return_value=mock_client_holder, ) mocker.patch( @@ -546,7 +541,7 @@ async def test_query_azure_token_refresh( "app.endpoints.query.AzureEntraIDManager", return_value=mock_azure_manager ) - mock_updated_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_updated_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder.update_azure_token = mocker.AsyncMock( return_value=mock_updated_client ) @@ -575,241 +570,3 @@ async def mock_retrieve_agent_response( ) mock_client_holder.update_azure_token.assert_called_once() - - -class TestRetrieveResponse: - """Tests for retrieve_response function.""" - - @pytest.mark.asyncio - async def test_retrieve_response_success(self, mocker: MockerFixture) -> None: - """Test successful response retrieval.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.input = "test query" - mock_responses_params.model = "provider1/model1" - mock_responses_params.tools = None - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_output_item = mocker.Mock() - mock_output_item.type = "message" - mock_output_item.content = "Response text" - - mock_usage = mocker.Mock() - mock_usage.input_tokens = 10 - mock_usage.output_tokens = 5 - mock_response = mocker.Mock(spec=OpenAIResponseObject) - mock_response.output = [mock_output_item] - mock_response.usage = mock_usage - - mock_client.responses.create = mocker.AsyncMock(return_value=mock_response) - - mock_summary = TurnSummary() - mock_summary.llm_response = "Response text" - mock_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - mocker.patch( - "app.endpoints.query.build_turn_summary", - return_value=mock_summary, - ) - - result = await retrieve_response( - mock_client, mock_responses_params, ShieldModerationPassed() - ) - - assert isinstance(result, TurnSummary) - assert result.llm_response == "Response text" - assert result.token_usage.input_tokens == 10 - assert result.token_usage.output_tokens == 5 - - @pytest.mark.asyncio - async def test_retrieve_response_shield_blocked( - self, mocker: MockerFixture - ) -> None: - """Test response retrieval when shield moderation blocks the request.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_refusal = mocker.Mock() - mock_moderation_result = mocker.Mock() - mock_moderation_result.decision = "blocked" - mock_moderation_result.message = "Content blocked by moderation" - mock_moderation_result.moderation_id = "mod_123" - mock_moderation_result.refusal_response = mock_refusal - mock_append = mocker.patch( - "app.endpoints.query.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) - - result = await retrieve_response( - mock_client, mock_responses_params, mock_moderation_result - ) - - assert isinstance(result, TurnSummary) - assert result.llm_response == "Content blocked by moderation" - mock_append.assert_called_once() - - @pytest.mark.asyncio - async def test_retrieve_response_connection_error( - self, mocker: MockerFixture - ) -> None: - """Test response retrieval raises HTTPException on connection error.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.input = "test query" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_client.responses.create = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", request=mocker.Mock() - ) - ) - - with pytest.raises(HTTPException) as exc_info: - await retrieve_response( - mock_client, mock_responses_params, ShieldModerationPassed() - ) - - assert exc_info.value.status_code == 503 - - @pytest.mark.asyncio - async def test_retrieve_response_api_status_error( - self, mocker: MockerFixture - ) -> None: - """Test response retrieval raises HTTPException on API status error.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.input = "test query" - mock_responses_params.model = "provider1/model1" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_client.responses.create = mocker.AsyncMock( - side_effect=APIStatusError( - message="API error", response=mocker.Mock(request=None), body=None - ) - ) - mocker.patch( - "app.endpoints.query.handle_known_apistatus_errors", - return_value=mocker.Mock( - model_dump=lambda: { - "status_code": 500, - "detail": {"response": "Error", "cause": "API error"}, - } - ), - ) - - with pytest.raises(HTTPException): - await retrieve_response( - mock_client, mock_responses_params, ShieldModerationPassed() - ) - - @pytest.mark.asyncio - async def test_retrieve_response_runtime_error_context_length( - self, mocker: MockerFixture - ) -> None: - """Test retrieve_response handles RuntimeError with context_length.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_client.responses.create = mocker.AsyncMock( - side_effect=RuntimeError("context_length exceeded") - ) - - with pytest.raises(HTTPException) as exc_info: - await retrieve_response( - mock_client, mock_responses_params, ShieldModerationPassed() - ) - - assert exc_info.value.status_code == 413 - - @pytest.mark.asyncio - async def test_retrieve_response_runtime_error_other( - self, mocker: MockerFixture - ) -> None: - """Test retrieve_response re-raises RuntimeError without context_length.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_client.responses.create = mocker.AsyncMock( - side_effect=RuntimeError("Some other error") - ) - - with pytest.raises(RuntimeError): - await retrieve_response( - mock_client, mock_responses_params, ShieldModerationPassed() - ) - - @pytest.mark.asyncio - async def test_retrieve_response_with_tool_calls( - self, mocker: MockerFixture - ) -> None: - """Test response retrieval processes tool calls.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.input = "test query" - mock_responses_params.model = "provider1/model1" - mock_responses_params.tools = None - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_usage = mocker.Mock() - mock_usage.input_tokens = 10 - mock_usage.output_tokens = 5 - mock_response = mocker.Mock(spec=OpenAIResponseObject) - mock_response.output = [mocker.Mock(type="message")] - mock_response.usage = mock_usage - - mock_client.responses.create = mocker.AsyncMock(return_value=mock_response) - - mock_tool_call = ToolCallSummary(id="1", name="test", args={}) - mock_tool_result = ToolResultSummary( - id="1", status="success", content="result", round=1 - ) - mock_summary = TurnSummary() - mock_summary.llm_response = "Response text" - mock_summary.tool_calls = [mock_tool_call] - mock_summary.tool_results = [mock_tool_result] - mock_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - mocker.patch( - "app.endpoints.query.build_turn_summary", - return_value=mock_summary, - ) - - result = await retrieve_response( - mock_client, mock_responses_params, ShieldModerationPassed() - ) - - assert result.llm_response == "Response text" - assert len(result.tool_calls) == 1 - assert len(result.tool_results) == 1 - assert result.token_usage.input_tokens == 10 - assert result.token_usage.output_tokens == 5 - assert result.rag_chunks == [] - assert result.referenced_documents == [] diff --git a/tests/unit/app/endpoints/test_rags.py b/tests/unit/app/endpoints/test_rags.py index 563c223fe..e641f25f0 100644 --- a/tests/unit/app/endpoints/test_rags.py +++ b/tests/unit/app/endpoints/test_rags.py @@ -5,7 +5,7 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError, BadRequestError +from ogx_client import APIConnectionError, BadRequestError from pytest_mock import MockerFixture from app.endpoints.rags import ( @@ -46,7 +46,7 @@ async def test_rags_endpoint_connection_error( mock_client = mocker.AsyncMock() mock_client.vector_stores.list.side_effect = APIConnectionError(request=None) # type: ignore mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -60,7 +60,7 @@ async def test_rags_endpoint_connection_error( detail = e.value.detail assert isinstance(detail, dict) assert "response" in detail - assert "Unable to connect to Llama Stack" in detail["response"] # type: ignore[index] + assert "Unable to connect to OGX" in detail["response"] # type: ignore[index] @pytest.mark.asyncio @@ -103,7 +103,7 @@ def __init__(self) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.list.return_value = RagList() mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -149,7 +149,7 @@ async def test_rag_info_endpoint_rag_not_found( ) ) # type: ignore mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -177,7 +177,7 @@ async def test_rag_info_endpoint_connection_error( request=None # type: ignore ) mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -191,7 +191,7 @@ async def test_rag_info_endpoint_connection_error( detail = e.value.detail assert isinstance(detail, dict) assert "response" in detail - assert "Unable to connect to Llama Stack" in detail["response"] # type: ignore[index] + assert "Unable to connect to OGX" in detail["response"] # type: ignore[index] @pytest.mark.asyncio @@ -231,7 +231,7 @@ def __init__(self) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.retrieve.return_value = RagInfo() mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -323,7 +323,7 @@ def __init__(self) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.list.return_value = RagList() mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) @@ -359,7 +359,7 @@ def __init__(self) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.retrieve.return_value = RagInfo() mocker.patch( - "app.endpoints.rags.AsyncLlamaStackClientHolder" + "app.endpoints.rags.AsyncOgxClientHolder" ).return_value.get_client.return_value = mock_client request = Request(scope={"type": "http"}) diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index 5c0c851fa..15fefe74a 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -9,14 +9,14 @@ import pytest from fastapi import HTTPException, Request from fastapi.responses import StreamingResponse -from llama_stack_api import OpenAIResponseObject -from llama_stack_api.openai_responses import ( +from ogx_api import OpenAIResponseObject +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoiceMode as ToolChoiceMode, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMessage, ) -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient from pytest_mock import MockerFixture from app.endpoints.responses import ( @@ -129,14 +129,14 @@ def _patch_base(mocker: MockerFixture, config: AppConfig) -> None: def _patch_client(mocker: MockerFixture) -> Any: - """Patch AsyncLlamaStackClientHolder; return (mock_client, mock_holder).""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + """Patch AsyncOgxClientHolder; return (mock_client, mock_holder).""" + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_vector_stores = mocker.Mock() mock_vector_stores.list = mocker.AsyncMock(return_value=mocker.Mock(data=[])) mock_client.vector_stores = mock_vector_stores mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) return mock_client, mock_holder @@ -183,16 +183,11 @@ def _patch_moderation(mocker: MockerFixture, decision: str = "passed") -> Any: moderation_result = ShieldModerationBlocked( message="Content blocked", moderation_id="mod_blocked", - refusal_response=OpenAIResponseMessage( - role="assistant", - content="Content blocked", - type="message", - ), ) else: moderation_result = ShieldModerationPassed() mocker.patch( - f"{MODULE}.run_shield_moderation", + f"{MODULE}.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=moderation_result), ) return moderation_result @@ -360,7 +355,7 @@ async def test_responses_with_conversation_validates_and_retrieves( ) _, mock_holder = _patch_client(mocker) mocker.patch( - f"{ENDPOINTS_MODULE}.AsyncLlamaStackClientHolder", + f"{ENDPOINTS_MODULE}.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch( @@ -495,7 +490,7 @@ async def test_responses_azure_token_refresh( mock_azure.is_token_expired = True mock_azure.refresh_token.return_value = True mocker.patch(f"{MODULE}.AzureEntraIDManager", return_value=mock_azure) - updated_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + updated_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_holder.update_azure_token = mocker.AsyncMock(return_value=updated_client) _patch_rag(mocker) _patch_moderation(mocker, decision="passed") @@ -599,7 +594,7 @@ async def test_responses_blocked_with_conversation_appends_refusal( ) mock_client, mock_holder = _patch_client(mocker) mocker.patch( - f"{ENDPOINTS_MODULE}.AsyncLlamaStackClientHolder", + f"{ENDPOINTS_MODULE}.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch( @@ -622,9 +617,6 @@ async def test_responses_blocked_with_conversation_appends_refusal( mock_moderation = _patch_moderation(mocker, decision="blocked") mock_moderation.message = "Blocked" mock_moderation.moderation_id = "resp_blocked_123" - mock_moderation.refusal_response = OpenAIResponseMessage( - type="message", role="assistant", content="Blocked" - ) mock_append = mocker.patch( f"{MODULE}.append_turn_items_to_conversation", new=mocker.AsyncMock(), @@ -754,7 +746,7 @@ async def test_handle_non_streaming_blocked_returns_refusal( ) -> None: """Test that blocked moderation returns response with refusal message.""" request = _request_with_model_and_conv("Bad input") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" mock_moderation.message = "Content blocked" @@ -816,7 +808,7 @@ async def test_handle_non_streaming_success_returns_response( ) -> None: """Test successful handle_non_streaming_response returns ResponsesResponse.""" request = _request_with_model_and_conv("Hello") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -896,7 +888,7 @@ async def test_handle_non_streaming_with_previous_response_id_appends_turn( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -978,7 +970,7 @@ async def test_handle_non_streaming_context_length_raises_413( ) -> None: """Test that RuntimeError with context_length raises 413.""" request = _request_with_model_and_conv("Long input") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=RuntimeError("context_length exceeded") ) @@ -1017,7 +1009,7 @@ async def test_handle_non_streaming_connection_error_raises_503( ) -> None: """Test that APIConnectionError raises 503.""" request = _request_with_model_and_conv("Hi") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=APIConnectionError( message="Connection failed", @@ -1063,7 +1055,7 @@ async def test_handle_non_streaming_api_status_error_raises_http( ) -> None: """Test that APIStatusError is handled and re-raised as HTTPException.""" request = _request_with_model_and_conv("Hi") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=APIStatusError( message="API error", @@ -1115,7 +1107,7 @@ async def test_handle_non_streaming_runtime_error_without_context_reraises( ) -> None: """Test that RuntimeError without context_length is re-raised.""" request = _request_with_model_and_conv("Hi") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=RuntimeError("Some other error") ) @@ -1156,7 +1148,7 @@ async def test_handle_streaming_blocked_returns_sse_consumes_shield_generator( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" mock_moderation.message = "Blocked" @@ -1220,7 +1212,7 @@ async def test_handle_streaming_success_returns_sse_consumes_response_generator( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1262,7 +1254,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=request, client=mock_client, @@ -1299,7 +1291,7 @@ async def test_handle_streaming_in_progress_chunk_sets_quotas_and_output_text( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1349,7 +1341,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=request, @@ -1386,7 +1378,7 @@ async def test_handle_streaming_builds_tool_call_summary_from_output( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1437,7 +1429,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=request, @@ -1473,7 +1465,7 @@ async def test_handle_streaming_with_previous_response_id_appends_turn( request = _request_with_previous_response_id( "Hi", previous_response_id="r_prev" ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -1519,7 +1511,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=request, @@ -1556,7 +1548,7 @@ async def test_handle_streaming_context_length_raises_413( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=RuntimeError("context_length exceeded") ) @@ -1593,7 +1585,7 @@ async def test_handle_streaming_connection_error_raises_503( ) -> None: """Test streaming raises 503 when create raises APIConnectionError.""" request = _request_with_model_and_conv("Hi", model="provider/model1") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=APIConnectionError( message="Connection failed", @@ -2280,7 +2272,7 @@ async def test_non_streaming_sanitizes_mcp_output_and_model( conversation=VALID_CONV_ID_NORMALIZED, ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2436,7 +2428,7 @@ async def test_streaming_sanitizes_mcp_output_model_and_instructions( instructions=SERVER_INSTRUCTIONS, conversation=VALID_CONV_ID_NORMALIZED, ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2471,7 +2463,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=updated_request, @@ -2537,7 +2529,7 @@ async def test_mcp_events_filtered_without_merge_server_tools_header( mock_config.rag_id_mapping = {} request = _request_with_model_and_conv("Hi", model="provider/model1") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2593,7 +2585,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=request, @@ -2634,7 +2626,7 @@ async def test_mcp_events_filtered_with_no_mcp_servers_configured( are filtered. """ request = _request_with_model_and_conv("Hi", model="provider/model1") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -2689,7 +2681,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) api_params, context = build_api_params_and_context( updated_request=request, @@ -2728,7 +2720,7 @@ async def test_response_generator_records_failure_when_stream_iteration_raises( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) 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 5b449377d..bb60329a5 100644 --- a/tests/unit/app/endpoints/test_responses_splunk.py +++ b/tests/unit/app/endpoints/test_responses_splunk.py @@ -7,10 +7,10 @@ import pytest from fastapi import HTTPException from fastapi.responses import StreamingResponse -from llama_stack_api import OpenAIResponseObject -from llama_stack_api.openai_responses import OpenAIResponseMessage -from llama_stack_client import APIConnectionError, AsyncLlamaStackClient -from llama_stack_client import APIStatusError as LLSApiStatusError +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 openai._exceptions import APIStatusError as OpenAIAPIStatusError from pytest_mock import MockerFixture @@ -230,7 +230,7 @@ async def test_non_streaming_shield_blocked( ) -> None: """Blocked moderation fires responses_shield_blocked telemetry.""" request = _request_with_model_and_conv("Bad input") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" @@ -324,7 +324,7 @@ async def test_non_streaming_error_fires_telemetry( ) -> None: """Each error branch fires responses_error telemetry with fire_and_forget.""" request = _request_with_model_and_conv("Hello") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -377,7 +377,7 @@ async def test_non_streaming_success( ) -> None: """Successful non-streaming response fires responses_completed with token counts.""" request = _request_with_model_and_conv("Hello") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -467,7 +467,7 @@ async def test_streaming_shield_blocked( ) -> None: """Blocked moderation in streaming fires responses_shield_blocked telemetry.""" request = _request_with_model_and_conv("Bad", model="provider/model1") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "blocked" @@ -552,7 +552,7 @@ async def test_streaming_error_fires_telemetry( ) -> None: """Each streaming error branch fires responses_error telemetry with fire_and_forget.""" request = _request_with_model_and_conv("Hello") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -605,7 +605,7 @@ async def test_streaming_success( ) -> None: """Successful streaming fires responses_completed after consuming the stream.""" request = _request_with_model_and_conv("Hi", model="provider/model1") - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_moderation = mocker.Mock() mock_moderation.decision = "passed" @@ -656,7 +656,7 @@ async def mock_stream() -> Any: ) mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client - mocker.patch(f"{MODULE}.AsyncLlamaStackClientHolder", return_value=mock_holder) + mocker.patch(f"{MODULE}.AsyncOgxClientHolder", return_value=mock_holder) mock_queue = mocker.patch(f"{TELEMETRY_MODULE}.queue_responses_splunk_event") @@ -701,7 +701,7 @@ async def test_splunk_disabled_no_background_tasks( ) -> None: """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=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) 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 4270c5705..63a62e68d 100644 --- a/tests/unit/app/endpoints/test_rlsapi_v1.py +++ b/tests/unit/app/endpoints/test_rlsapi_v1.py @@ -13,8 +13,9 @@ import pytest from fastapi import HTTPException, status -from llama_stack_api import OpenAIResponseMessage -from llama_stack_client import APIConnectionError, APIStatusError +from ogx_client import APIConnectionError, APIStatusError +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pydantic import ValidationError from pytest_mock import MockerFixture @@ -24,6 +25,7 @@ TemplateRenderError, _build_instructions, _call_llm, + _check_shield_moderation, _compile_prompt_template, _get_default_model_id, _resolve_quota_subject, @@ -42,6 +44,13 @@ from models.api.responses.error import ServiceUnavailableResponse from models.api.responses.successful.rlsapi import RlsapiV1InferResponse from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed +from models.config import ( + QuestionValidityConfig, + QuestionValidityShieldConfiguration, + RedactionConfig, + RedactionRule, + RedactionShieldConfiguration, +) 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 @@ -69,6 +78,7 @@ def _set(prompt: str) -> None: mock_config.customization = mock_customization mock_config.rlsapi_v1 = mock_rlsapi_v1 mock_config.quota_limiters = [] + mock_config.shields = [] mocker.patch("app.endpoints.rlsapi_v1.configuration", mock_config) return _set @@ -85,7 +95,7 @@ def _setup_responses_mock(mocker: MockerFixture, create_behavior: Any) -> None: mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder", + "app.endpoints.rlsapi_v1.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -150,7 +160,7 @@ def mock_shield_passed_fixture(mocker: MockerFixture) -> None: with a different return value. """ mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=ShieldModerationPassed()), ) @@ -352,15 +362,21 @@ async def test_get_default_model_id_errors( """Test _get_default_model_id fallback failures raise 503 responses.""" mocker.patch("app.endpoints.rlsapi_v1.configuration", minimal_config) - mock_embedding_model = mocker.Mock() - mock_embedding_model.custom_metadata = {"model_type": "embedding"} - mock_embedding_model.id = "sentence-transformers/all-mpnet-base-v2" + mock_embedding_model = Model.model_construct( + id="sentence-transformers/all-mpnet-base-v2", + created=0, + owned_by="test", + object="model", + custom_metadata={"model_type": "embedding"}, + ) mock_client = mocker.Mock() mock_client.models = mocker.Mock() if failure_mode == "no_llm_models": - mock_client.models.list = mocker.AsyncMock(return_value=[mock_embedding_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct(data=[mock_embedding_model]) + ) else: mock_client.models.list = mocker.AsyncMock( side_effect=APIConnectionError(request=mocker.Mock()) @@ -369,7 +385,7 @@ async def test_get_default_model_id_errors( mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder", + "app.endpoints.rlsapi_v1.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -394,18 +410,24 @@ async def test_config_error_503_matches_llm_error_503_shape( """ mocker.patch("app.endpoints.rlsapi_v1.configuration", minimal_config) - mock_embedding_model = mocker.Mock() - mock_embedding_model.custom_metadata = {"model_type": "embedding"} - mock_embedding_model.id = "sentence-transformers/all-mpnet-base-v2" + mock_embedding_model = Model.model_construct( + id="sentence-transformers/all-mpnet-base-v2", + created=0, + owned_by="test", + object="model", + custom_metadata={"model_type": "embedding"}, + ) mock_client = mocker.Mock() mock_client.models = mocker.Mock() - mock_client.models.list = mocker.AsyncMock(return_value=[mock_embedding_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct(data=[mock_embedding_model]) + ) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder", + "app.endpoints.rlsapi_v1.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -414,7 +436,7 @@ async def test_config_error_503_matches_llm_error_503_shape( # Build an LLM connection error 503 using the same response model llm_response = ServiceUnavailableResponse( - backend_name="Llama Stack", + backend_name="OGX", cause="Unable to connect to the inference backend", ) llm_detail = llm_response.model_dump()["detail"] @@ -432,24 +454,34 @@ async def test_get_default_model_id_auto_discovery_success( """Test _get_default_model_id returns first discovered LLM model ID.""" mocker.patch("app.endpoints.rlsapi_v1.configuration", minimal_config) - mock_llm_model = mocker.Mock() - mock_llm_model.custom_metadata = {"model_type": "llm"} - mock_llm_model.id = "openai/gpt-4o-mini" + mock_llm_model = Model.model_construct( + id="openai/gpt-4o-mini", + created=0, + owned_by="test", + object="model", + custom_metadata={"model_type": "llm"}, + ) - mock_embedding_model = mocker.Mock() - mock_embedding_model.custom_metadata = {"model_type": "embedding"} - mock_embedding_model.id = "sentence-transformers/all-mpnet-base-v2" + mock_embedding_model = Model.model_construct( + id="sentence-transformers/all-mpnet-base-v2", + created=0, + owned_by="test", + object="model", + custom_metadata={"model_type": "embedding"}, + ) mock_client = mocker.Mock() mock_client.models = mocker.Mock() mock_client.models.list = mocker.AsyncMock( - return_value=[mock_embedding_model, mock_llm_model] + return_value=ListModelsResponse.model_construct( + data=[mock_embedding_model, mock_llm_model] + ) ) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.rlsapi_v1.AsyncLlamaStackClientHolder", + "app.endpoints.rlsapi_v1.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -828,6 +860,7 @@ async def test_infer_include_metadata_respects_verbose_config( config_mock.customization = mock_configuration.customization config_mock.rlsapi_v1 = rlsapi_v1_mock config_mock.quota_limiters = [] + config_mock.shields = [] mocker.patch("app.endpoints.rlsapi_v1.configuration", config_mock) mock_response = mocker.Mock() @@ -884,6 +917,7 @@ def _setup_config_mock( config_mock.customization = mock_configuration.customization config_mock.rlsapi_v1 = rlsapi_v1_mock config_mock.quota_limiters = [] + config_mock.shields = [] mocker.patch("app.endpoints.rlsapi_v1.configuration", config_mock) @@ -1082,6 +1116,7 @@ def _set(quota_subject: str) -> None: config_mock.customization = mock_configuration.customization config_mock.rlsapi_v1 = rlsapi_v1_mock config_mock.quota_limiters = [] + config_mock.shields = [] mocker.patch("app.endpoints.rlsapi_v1.configuration", config_mock) return _set @@ -1230,13 +1265,9 @@ async def test_infer_quota_shield_blocked_does_not_consume_tokens( blocked = ShieldModerationBlocked( message="Blocked by moderation", moderation_id="modr-test", - refusal_response=OpenAIResponseMessage( - role="assistant", - content="Blocked by moderation", - ), ) mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=blocked), ) @@ -1263,10 +1294,6 @@ def _create_blocked_moderation_result() -> ShieldModerationBlocked: return ShieldModerationBlocked( message="I can't answer that. Can I help with something else?", moderation_id="modr-test-123", - refusal_response=OpenAIResponseMessage( - role="assistant", - content="I can't answer that. Can I help with something else?", - ), ) @@ -1282,7 +1309,7 @@ async def test_infer_shield_blocked_returns_refusal( """Test that blocked shield moderation returns refusal text without calling LLM.""" blocked = _create_blocked_moderation_result() mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=blocked), ) @@ -1321,7 +1348,7 @@ async def test_infer_shield_blocked_skips_llm_call( """Test that blocked shield moderation prevents any LLM call.""" blocked = _create_blocked_moderation_result() mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=blocked), ) mock_call_llm = mocker.patch( @@ -1353,7 +1380,7 @@ async def test_infer_shield_blocked_queues_splunk_event( """Test that blocked shield moderation queues a Splunk event with correct sourcetype.""" blocked = _create_blocked_moderation_result() mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=blocked), ) @@ -1409,7 +1436,7 @@ async def test_infer_shield_moderation_receives_combined_input( """Test that shield moderation receives the full combined input source.""" mock_moderation = mocker.AsyncMock(return_value=ShieldModerationPassed()) mocker.patch( - "app.endpoints.rlsapi_v1.run_shield_moderation", + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", new=mock_moderation, ) @@ -1429,13 +1456,159 @@ async def test_infer_shield_moderation_receives_combined_input( ) mock_moderation.assert_called_once() - # The input_text argument should be the combined input source - input_text = mock_moderation.call_args[0][1] + # The input_text argument is the first positional arg to run_shield_moderation_v2 + input_text = mock_moderation.call_args[0][0] assert "Why did this fail?" in input_text assert "piped input" in input_text assert "permission denied" in input_text +# --- Test PII redaction behavior via _check_shield_moderation --- + + +def _redaction_shield() -> RedactionShieldConfiguration: + """Build a redaction shield that replaces email addresses.""" + return RedactionShieldConfiguration( + name="pii-redaction", + provider_id="redaction", + config=RedactionConfig( + rules=[ + RedactionRule( + pattern=r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}", + replacement="[REDACTED_EMAIL]", + ) + ], + ), + ) + + +def _question_validity_shield() -> QuestionValidityShieldConfiguration: + """Build a question-validity shield configuration.""" + return QuestionValidityShieldConfiguration( + name="question-validity", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ) + + +@pytest.mark.asyncio +async def test_check_shield_redaction_modifies_input( + mocker: MockerFixture, +) -> None: + """Test that PII redaction returns redacted text as moderated_input.""" + mocker.patch( + "app.endpoints.rlsapi_v1.configuration", + mocker.Mock(shields=[_redaction_shield()]), + ) + mocker.patch( + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", + new=mocker.AsyncMock(return_value=ShieldModerationPassed()), + ) + + refusal, moderated = await _check_shield_moderation( + "Contact user@example.com for details", + "req-1", + mocker.Mock(), + mocker.Mock(), + mocker.Mock(), + ) + + assert refusal is None + assert "[REDACTED_EMAIL]" in moderated + assert "user@example.com" not in moderated + + +@pytest.mark.asyncio +async def test_check_shield_no_pii_passes_original( + mocker: MockerFixture, +) -> None: + """Test that clean input passes through redaction unchanged.""" + mocker.patch( + "app.endpoints.rlsapi_v1.configuration", + mocker.Mock(shields=[_redaction_shield()]), + ) + mocker.patch( + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", + new=mocker.AsyncMock(return_value=ShieldModerationPassed()), + ) + + refusal, moderated = await _check_shield_moderation( + "How do I list files?", + "req-2", + mocker.Mock(), + mocker.Mock(), + mocker.Mock(), + ) + + assert refusal is None + assert moderated == "How do I list files?" + + +@pytest.mark.asyncio +async def test_check_shield_redaction_then_validity_passes_redacted_to_v2( + mocker: MockerFixture, +) -> None: + """Test that redacted text is forwarded to run_shield_moderation_v2.""" + mocker.patch( + "app.endpoints.rlsapi_v1.configuration", + mocker.Mock( + shields=[_redaction_shield(), _question_validity_shield()], + ), + ) + mock_v2 = mocker.AsyncMock(return_value=ShieldModerationPassed()) + mocker.patch( + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", + new=mock_v2, + ) + + refusal, moderated = await _check_shield_moderation( + "Contact user@example.com for help", + "req-3", + mocker.Mock(), + mocker.Mock(), + mocker.Mock(), + ) + + assert refusal is None + assert "[REDACTED_EMAIL]" in moderated + v2_input = mock_v2.call_args[0][0] + assert "[REDACTED_EMAIL]" in v2_input + assert "user@example.com" not in v2_input + + +@pytest.mark.asyncio +async def test_check_shield_validity_blocks_returns_refusal( + mocker: MockerFixture, +) -> None: + """Test that question validity block returns a refusal response.""" + mocker.patch( + "app.endpoints.rlsapi_v1.configuration", + mocker.Mock( + shields=[_redaction_shield(), _question_validity_shield()], + ), + ) + blocked = ShieldModerationBlocked( + message="Off-topic question.", + moderation_id="modr-block-123", + ) + mocker.patch( + "app.endpoints.rlsapi_v1.run_shield_moderation_v2", + new=mocker.AsyncMock(return_value=blocked), + ) + + refusal, _ = await _check_shield_moderation( + "What is the meaning of life?", + "req-4", + mocker.Mock(), + mocker.Mock(), + mocker.Mock(), + ) + + assert refusal is not None + assert isinstance(refusal, RlsapiV1InferResponse) + assert refusal.data.text == "Off-topic question." + + @pytest.mark.asyncio async def test_infer_splunk_event_includes_rh_identity_context( mocker: MockerFixture, diff --git a/tests/unit/app/endpoints/test_shields.py b/tests/unit/app/endpoints/test_shields.py index a74481eef..704478f51 100644 --- a/tests/unit/app/endpoints/test_shields.py +++ b/tests/unit/app/endpoints/test_shields.py @@ -4,7 +4,6 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError from pytest_mock import MockerFixture from app.endpoints.shields import shields_endpoint_handler @@ -14,52 +13,9 @@ from tests.unit.utils.auth_helpers import mock_authorization_resolvers -@pytest.mark.asyncio -async def test_shields_endpoint_handler_configuration_not_loaded( - mocker: MockerFixture, -) -> None: - """Test the shields endpoint handler if configuration is not loaded.""" - mock_authorization_resolvers(mocker) - - # simulate state when no configuration is loaded - mock_config = AppConfig() - mock_config._configuration = None # pylint: disable=protected-access - mocker.patch("app.endpoints.shields.configuration", mock_config) - - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) - - # Authorization tuple required by URL endpoint handler - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") - - with pytest.raises(HTTPException) as e: - await shields_endpoint_handler(request=request, auth=auth) - assert e.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR - assert e.value.detail["response"] == "Configuration is not loaded" # type: ignore - - -@pytest.mark.asyncio -async def test_shields_endpoint_handler_improper_llama_stack_configuration( - mocker: MockerFixture, -) -> None: - """Test the shields endpoint handler if Llama Stack configuration is not proper. - - Verify shields_endpoint_handler returns an empty ShieldsResponse when Llama - Stack is configured minimally and the client provides no shields. - - Patches the endpoint configuration and client holder to supply a mocked - Llama Stack client whose `shields.list` returns an empty list, then calls - the handler with a test request and authorization tuple and asserts the - response is a ShieldsResponse with an empty `shields` list. - """ - mock_authorization_resolvers(mocker) - - # configuration for tests - config_dict: dict[str, Any] = { +def _base_config_dict() -> dict[str, Any]: + """Return a minimal valid AppConfig dictionary.""" + return { "name": "test", "service": { "host": "localhost", @@ -82,157 +38,51 @@ async def test_shields_endpoint_handler_improper_llama_stack_configuration( "authorization": {"access_rules": []}, "authentication": {"module": "noop"}, } - cfg = AppConfig() - cfg.init_from_dict(config_dict) - mocker.patch("app.endpoints.shields.configuration", cfg) - # Mock client to avoid initialization - mock_client_holder = mocker.patch( - "app.endpoints.shields.AsyncLlamaStackClientHolder" - ) - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client +def _auth_request() -> tuple[Request, AuthTuple]: + """Return a dummy request and auth tuple for the shields handler.""" request = Request( scope={ "type": "http", "headers": [(b"authorization", b"Bearer invalid-token")], } ) - - # Authorization tuple required by URL endpoint handler auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") - - # Mock shields.list to return empty list - mock_client.shields.list.return_value = [] - - response = await shields_endpoint_handler(request=request, auth=auth) - assert isinstance(response, ShieldsResponse) - assert response.shields == [] + return request, auth @pytest.mark.asyncio -async def test_shields_endpoint_handler_configuration_loaded( +async def test_shields_endpoint_handler_configuration_not_loaded( mocker: MockerFixture, ) -> None: - """Test the shields endpoint handler if configuration is loaded. - - Verify shields_endpoint_handler raises an HTTP 503 with detail "Unable to - connect to Llama Stack" when configuration is loaded but the Llama Stack - client is unreachable. - - Sets up an AppConfig from a valid configuration, patches the endpoint's - configuration and AsyncLlamaStackClientHolder to return a client whose - shields.list raises APIConnectionError, and asserts the handler raises an - HTTPException with status 503 and the expected detail. - - Parameters: - ---------- - mocker (MockerFixture): pytest-mock fixture used to create patches and mocks. - """ + """Test the shields endpoint handler if configuration is not loaded.""" mock_authorization_resolvers(mocker) - # configuration for tests - config_dict: dict[str, Any] = { - "name": "foo", - "service": { - "host": "localhost", - "port": 8080, - "auth_enabled": False, - "workers": 1, - "color_log": True, - "access_log": True, - }, - "llama_stack": { - "api_key": "xyzzy", - "url": "http://x.y.com:1234", - "use_as_library_client": False, - }, - "user_data_collection": { - "feedback_enabled": False, - }, - "customization": None, - "authorization": {"access_rules": []}, - "authentication": {"module": "noop"}, - } - cfg = AppConfig() - cfg.init_from_dict(config_dict) - - mocker.patch("app.endpoints.shields.configuration", cfg) - # Mock client to raise APIConnectionError - mock_client_holder = mocker.patch( - "app.endpoints.shields.AsyncLlamaStackClientHolder" - ) - mock_client = mocker.AsyncMock() - mock_client.shields.list.side_effect = APIConnectionError(request=None) # type: ignore - mock_client_holder.return_value.get_client.return_value = mock_client - - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) + mock_config = AppConfig() + mock_config._configuration = None # pylint: disable=protected-access + mocker.patch("app.endpoints.shields.configuration", mock_config) - # Authorization tuple required by URL endpoint handler - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") + request, auth = _auth_request() with pytest.raises(HTTPException) as e: await shields_endpoint_handler(request=request, auth=auth) - assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE - assert e.value.detail["response"] == "Unable to connect to Llama Stack" # type: ignore + assert e.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR + assert e.value.detail["response"] == "Configuration is not loaded" # type: ignore @pytest.mark.asyncio -async def test_shields_endpoint_handler_unable_to_retrieve_shields_list( +async def test_shields_endpoint_handler_empty_shields( mocker: MockerFixture, ) -> None: - """Test the shields endpoint handler if configuration is loaded.""" + """Test the shields endpoint returns an empty list when none are configured.""" mock_authorization_resolvers(mocker) - # configuration for tests - config_dict: dict[str, Any] = { - "name": "foo", - "service": { - "host": "localhost", - "port": 8080, - "auth_enabled": False, - "workers": 1, - "color_log": True, - "access_log": True, - }, - "llama_stack": { - "api_key": "xyzzy", - "url": "http://x.y.com:1234", - "use_as_library_client": False, - }, - "user_data_collection": { - "feedback_enabled": False, - }, - "customization": None, - "authorization": {"access_rules": []}, - "authentication": {"module": "noop"}, - } cfg = AppConfig() - cfg.init_from_dict(config_dict) - - # Mock the LlamaStack client - mock_client = mocker.AsyncMock() - mock_client.shields.list.return_value = [] - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_client - mock_config = mocker.Mock() - mocker.patch("app.endpoints.shields.configuration", mock_config) + cfg.init_from_dict(_base_config_dict()) + mocker.patch("app.endpoints.shields.configuration", cfg) - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) - - # Authorization tuple required by URL endpoint handler - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") + request, auth = _auth_request() response = await shields_endpoint_handler(request=request, auth=auth) assert isinstance(response, ShieldsResponse) @@ -240,259 +90,52 @@ async def test_shields_endpoint_handler_unable_to_retrieve_shields_list( @pytest.mark.asyncio -async def test_shields_endpoint_llama_stack_connection_error( +async def test_shields_endpoint_handler_configured_shields( mocker: MockerFixture, ) -> None: - """Test the shields endpoint when LlamaStack connection fails. - - Verifies that the shields endpoint responds with HTTP 503 and an - appropriate cause when the Llama Stack client cannot be reached. - - Simulates the Llama Stack client raising an APIConnectionError and asserts - that calling the endpoint raises an HTTPException with status 503, a detail - response of "Unable to connect to Llama Stack", and a detail cause that - contains "Connection error". - """ + """Test the shields endpoint lists shields from LCS configuration.""" mock_authorization_resolvers(mocker) - # configuration for tests - config_dict: dict[str, Any] = { - "name": "foo", - "service": { - "host": "localhost", - "port": 8080, - "auth_enabled": False, - "workers": 1, - "color_log": True, - "access_log": True, - }, - "llama_stack": { - "api_key": "xyzzy", - "url": "http://x.y.com:1234", - "use_as_library_client": False, - }, - "user_data_collection": { - "feedback_enabled": False, - }, - "customization": None, - "authorization": {"access_rules": []}, - "authentication": {"module": "noop"}, - } - - # mock AsyncLlamaStackClientHolder to raise APIConnectionError - # when shields.list() method is called - mock_client = mocker.AsyncMock() - mock_client.shields.list.side_effect = APIConnectionError(request=None) # type: ignore - mock_client_holder = mocker.patch( - "app.endpoints.shields.AsyncLlamaStackClientHolder" - ) - mock_client_holder.return_value.get_client.return_value = mock_client - - cfg = AppConfig() - cfg.init_from_dict(config_dict) - - mocker.patch("app.endpoints.shields.configuration", cfg) - - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) - - # Authorization tuple required by URL endpoint handler - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") - - with pytest.raises(HTTPException) as e: - await shields_endpoint_handler(request=request, auth=auth) - assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE - assert e.value.detail["response"] == "Unable to connect to Llama Stack" # type: ignore - assert "Connection error" in e.value.detail["cause"] # type: ignore - - -@pytest.mark.asyncio -async def test_shields_endpoint_handler_success_with_shields_data( - mocker: MockerFixture, -) -> None: - """Test the shields endpoint handler with successful response and shields data.""" - mock_authorization_resolvers(mocker) - - # configuration for tests - config_dict: dict[str, Any] = { - "name": "foo", - "service": { - "host": "localhost", - "port": 8080, - "auth_enabled": False, - "workers": 1, - "color_log": True, - "access_log": True, - }, - "llama_stack": { - "api_key": "xyzzy", - "url": "http://x.y.com:1234", - "use_as_library_client": False, - }, - "user_data_collection": { - "feedback_enabled": False, - }, - "customization": None, - "authorization": {"access_rules": []}, - "authentication": {"module": "noop"}, - } - cfg = AppConfig() - cfg.init_from_dict(config_dict) - - # Mock the LlamaStack client with sample shields data - mock_shields_data = [ + config_dict = _base_config_dict() + config_dict["shields"] = [ { - "identifier": "lightspeed_question_validity-shield", - "provider_resource_id": "lightspeed_question_validity-shield", - "provider_id": "lightspeed_question_validity", - "type": "shield", - "params": {}, + "name": "question-validity", + "provider_id": "question_validity", + "config": { + "model_id": "openai/gpt-4o-mini", + "model_prompt": "Is this question valid?", + "invalid_question_response": "I can only answer product questions.", + }, }, { - "identifier": "content_filter-shield", - "provider_resource_id": "content_filter-shield", - "provider_id": "content_filter", - "type": "shield", - "params": {"threshold": 0.8}, + "name": "pii-redaction", + "provider_id": "redaction", + "config": { + "rules": [ + { + "pattern": r"\b\d{3}-\d{2}-\d{4}\b", + "replacement": "[REDACTED]", + } + ], + "case_sensitive": False, + }, }, ] - - mock_client = mocker.AsyncMock() - mock_client.shields.list.return_value = mock_shields_data - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_client - mock_config = mocker.Mock() - mocker.patch("app.endpoints.shields.configuration", mock_config) - - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) - - # Authorization tuple required by URL endpoint handler - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") - - response = await shields_endpoint_handler(request=request, auth=auth) - - assert response is not None - assert hasattr(response, "shields") - assert len(response.shields) == 2 - assert response.shields[0]["identifier"] == "lightspeed_question_validity-shield" - assert response.shields[1]["identifier"] == "content_filter-shield" - - -@pytest.mark.asyncio -async def test_shields_endpoint_handler_unexpected_exception( - mocker: MockerFixture, -) -> None: - """Test the shields endpoint when an unexpected exception is raised.""" - mock_authorization_resolvers(mocker) - - config_dict: dict[str, Any] = { - "name": "foo", - "service": { - "host": "localhost", - "port": 8080, - "auth_enabled": False, - "workers": 1, - "color_log": True, - "access_log": True, - }, - "llama_stack": { - "api_key": "xyzzy", - "url": "http://x.y.com:1234", - "use_as_library_client": False, - }, - "user_data_collection": { - "feedback_enabled": False, - }, - "customization": None, - "authorization": {"access_rules": []}, - "authentication": {"module": "noop"}, - } cfg = AppConfig() cfg.init_from_dict(config_dict) + mocker.patch("app.endpoints.shields.configuration", cfg) - mock_client = mocker.AsyncMock() - mock_client.shields.list.side_effect = RuntimeError("unexpected failure") - mock_client_holder = mocker.patch( - "app.endpoints.shields.AsyncLlamaStackClientHolder" - ) - mock_client_holder.return_value.get_client.return_value = mock_client - - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) - - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") - - with pytest.raises(RuntimeError, match="unexpected failure"): - await shields_endpoint_handler(request=request, auth=auth) - - -@pytest.mark.asyncio -async def test_shields_endpoint_handler_malformed_shield_objects( - mocker: MockerFixture, -) -> None: - """Test the shields endpoint handles shields that may have missing fields.""" - mock_authorization_resolvers(mocker) - - config_dict: dict[str, Any] = { - "name": "foo", - "service": { - "host": "localhost", - "port": 8080, - "auth_enabled": False, - "workers": 1, - "color_log": True, - "access_log": True, - }, - "llama_stack": { - "api_key": "xyzzy", - "url": "http://x.y.com:1234", - "use_as_library_client": False, - }, - "user_data_collection": { - "feedback_enabled": False, - }, - "customization": None, - "authorization": {"access_rules": []}, - "authentication": {"module": "noop"}, - } - cfg = AppConfig() - cfg.init_from_dict(config_dict) - - mock_shield_minimal = { - "identifier": "minimal-shield", - } - - mock_client = mocker.AsyncMock() - mock_client.shields.list.return_value = [mock_shield_minimal] - mock_client_holder = mocker.patch( - "app.endpoints.shields.AsyncLlamaStackClientHolder" - ) - mock_client_holder.return_value.get_client.return_value = mock_client - - request = Request( - scope={ - "type": "http", - "headers": [(b"authorization", b"Bearer invalid-token")], - } - ) - - auth: AuthTuple = ("test_user_id", "test_user", True, "test_token") + request, auth = _auth_request() response = await shields_endpoint_handler(request=request, auth=auth) assert isinstance(response, ShieldsResponse) - assert len(response.shields) == 1 - assert response.shields[0]["identifier"] == "minimal-shield" + assert len(response.shields) == 2 + assert response.shields[0].name == "question-validity" + assert response.shields[0].provider_id == "question_validity" + assert response.shields[0].type == "shield" + assert response.shields[0].config["model_id"] == "openai/gpt-4o-mini" + assert response.shields[1].name == "pii-redaction" + assert response.shields[1].provider_id == "redaction" + assert response.shields[1].type == "shield" + assert response.shields[1].config["rules"][0]["replacement"] == "[REDACTED]" diff --git a/tests/unit/app/endpoints/test_streaming_query.py b/tests/unit/app/endpoints/test_streaming_query.py index fb2d9027b..3cdc4855b 100644 --- a/tests/unit/app/endpoints/test_streaming_query.py +++ b/tests/unit/app/endpoints/test_streaming_query.py @@ -1,75 +1,31 @@ -# pylint: disable=redefined-outer-name,import-error, too-many-function-args """Unit tests for the /streaming_query (v2) endpoint using Responses API.""" -# pylint: disable=too-many-lines,too-many-function-args -import asyncio from collections.abc import AsyncIterator from typing import Any import pytest -from fastapi import HTTPException, Request +from fastapi import Request from fastapi.responses import StreamingResponse -from llama_stack_api.openai_responses import ( - OpenAIResponseObject, - OpenAIResponseObjectStream, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseCompleted as CompletedChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseFailed as FailedChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseIncomplete as IncompleteChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseMcpCallArgumentsDone as MCPArgsDoneChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseOutputItemAdded as OutputItemAddedChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseOutputItemDone as OutputItemDoneChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseOutputTextDelta as TextDeltaChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseObjectStreamResponseOutputTextDone as TextDoneChunk, -) -from llama_stack_api.openai_responses import ( - OpenAIResponseOutputMessageMCPCall as MCPCall, -) -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from ogx_client import AsyncOgxClient from pytest_mock import MockerFixture from app.endpoints.streaming_query import ( - generate_response, - response_generator, - retrieve_response_generator, streaming_query_endpoint_handler, ) from configuration import AppConfig from constants import ( INTERRUPTED_RESPONSE_MESSAGE, - MEDIA_TYPE_JSON, MEDIA_TYPE_TEXT, ) from models.api.requests import QueryRequest -from models.api.responses.error import InternalServerErrorResponse from models.common.moderation import ShieldModerationPassed from models.common.query import Attachment -from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams from models.common.turn_summary import ( - RAGChunk, RAGContext, - ReferencedDocument, TurnSummary, ) from models.config import Action -from utils.stream_interrupts import StreamInterruptRegistry -from utils.token_counter import TokenCounter INTERRUPTED_INDICATOR = f"\n\n*{INTERRUPTED_RESPONSE_MESSAGE}*" @@ -170,11 +126,11 @@ async def test_successful_streaming_query( new=mocker.AsyncMock(return_value=RAGContext()), ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder", + "app.endpoints.streaming_query.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -257,11 +213,11 @@ async def test_streaming_query_text_media_type_header( new=mocker.AsyncMock(return_value=RAGContext()), ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder", + "app.endpoints.streaming_query.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -355,11 +311,11 @@ async def test_streaming_query_with_conversation( return_value=mock_conversation, ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder", + "app.endpoints.streaming_query.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -451,11 +407,11 @@ async def test_streaming_query_with_attachments( "app.endpoints.streaming_query.validate_attachments_metadata" ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder", + "app.endpoints.streaming_query.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -537,15 +493,15 @@ async def test_streaming_query_azure_token_refresh( new=mocker.AsyncMock(return_value=RAGContext()), ) - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_updated_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) + mock_updated_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client_holder = mocker.Mock() mock_client_holder.get_client.return_value = mock_client mock_client_holder.update_azure_token = mocker.AsyncMock( return_value=mock_updated_client ) mocker.patch( - "app.endpoints.streaming_query.AsyncLlamaStackClientHolder", + "app.endpoints.streaming_query.AsyncOgxClientHolder", return_value=mock_client_holder, ) @@ -613,2095 +569,3 @@ async def mock_generate_agent_response( ) mock_client_holder.update_azure_token.assert_called_once() - - -class TestCreateResponseGenerator: - """Tests for retrieve_response_generator function.""" - - @pytest.mark.asyncio - async def test_retrieve_response_generator_success( - self, mocker: MockerFixture - ) -> None: - """Test successful response generator creation.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - } - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.moderation_result = ShieldModerationPassed() - - async def mock_response_gen() -> AsyncIterator[str]: - yield "test" - - mock_client.responses = mocker.Mock() - mock_client.responses.create = mocker.AsyncMock( - return_value=mock_response_gen() - ) - - async def mock_response_generator( - *_args: Any, **_kwargs: Any - ) -> AsyncIterator[str]: - async for item in mock_response_gen(): - yield item - - mocker.patch( - "app.endpoints.streaming_query.response_generator", - side_effect=mock_response_generator, - ) - - generator, turn_summary = await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - assert isinstance(turn_summary, TurnSummary) - assert hasattr(generator, "__aiter__") - - @pytest.mark.asyncio - async def test_retrieve_response_generator_shield_blocked( - self, mocker: MockerFixture - ) -> None: - """Test response generator creation when shield blocks.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.omit_conversation = False # non-compacted - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_TEXT - ) # pyright: ignore[reportCallIssue] - - mock_moderation_result = mocker.Mock() - mock_moderation_result.decision = "blocked" - mock_moderation_result.message = "Content blocked" - mock_moderation_result.moderation_id = "mod_123" - mock_moderation_result.refusal_response = mocker.Mock() - mock_context.moderation_result = mock_moderation_result - mock_append = mocker.patch( - "app.endpoints.streaming_query.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) - - _generator, turn_summary = await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - assert isinstance(turn_summary, TurnSummary) - assert turn_summary.llm_response == "Content blocked" - # Structured refusal captured for compacted-mode persistence (LCORE-1572). - assert turn_summary.output_items == [mock_moderation_result.refusal_response] - # Non-compacted: the refusal turn is stored here. - mock_append.assert_awaited_once() - - @pytest.mark.asyncio - async def test_retrieve_response_generator_shield_blocked_compacted( - self, mocker: MockerFixture - ) -> None: - """In compacted mode the shield refusal is not stored here (no double-store). - - generate_response persists the compacted turn (with the original input), - so storing it again in the shield branch would duplicate it (LCORE-1572). - """ - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "explicit input" - mock_responses_params.conversation = "conv_123" - mock_responses_params.omit_conversation = True # compacted - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_TEXT - ) # pyright: ignore[reportCallIssue] - - mock_moderation_result = mocker.Mock() - mock_moderation_result.decision = "blocked" - mock_moderation_result.message = "Content blocked" - mock_moderation_result.moderation_id = "mod_123" - mock_moderation_result.refusal_response = mocker.Mock() - mock_context.moderation_result = mock_moderation_result - mock_append = mocker.patch( - "app.endpoints.streaming_query.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) - - _generator, turn_summary = await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - assert turn_summary.output_items == [mock_moderation_result.refusal_response] - mock_append.assert_not_awaited() # compacted: generate_response stores it - - @pytest.mark.asyncio - async def test_retrieve_response_generator_connection_error( - self, mocker: MockerFixture - ) -> None: - """Test response generator creation raises HTTPException on connection error.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - "conversation": "conv_123", - } - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.moderation_result = ShieldModerationPassed() - - mock_request_obj = mocker.Mock() - mock_client.responses = mocker.Mock() - mock_client.responses.create = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", request=mock_request_obj - ) - ) - - mock_error_response = mocker.Mock() - mock_error_response.model_dump.return_value = { - "status_code": 503, - "detail": { - "response": "Unable to connect to Llama Stack", - "cause": "Connection failed", - }, - } - mocker.patch( - "app.endpoints.streaming_query.ServiceUnavailableResponse", - return_value=mock_error_response, - ) - - with pytest.raises(HTTPException) as exc_info: - await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - assert exc_info.value.status_code == 503 - - @pytest.mark.asyncio - async def test_retrieve_response_generator_api_status_error( - self, mocker: MockerFixture - ) -> None: - """Test response generator creation raises HTTPException on API status error.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - "conversation": "conv_123", - } - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.moderation_result = ShieldModerationPassed() - - mock_request_obj = mocker.Mock() - mock_client.responses = mocker.Mock() - mock_client.responses.create = mocker.AsyncMock( - side_effect=APIStatusError( - message="API error", response=mock_request_obj, body=None - ) - ) - - mock_error_response = mocker.Mock() - mock_error_response.model_dump.return_value = { - "status_code": 500, - "detail": {"response": "Error", "cause": "API error"}, - } - mocker.patch( - "app.endpoints.streaming_query.handle_known_apistatus_errors", - return_value=mock_error_response, - ) - - with pytest.raises(HTTPException) as exc_info: - await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - assert exc_info.value.status_code == 500 - - @pytest.mark.asyncio - async def test_retrieve_response_generator_runtime_error_context_length( - self, mocker: MockerFixture - ) -> None: - """Test response generator raises HTTPException on RuntimeError with context_length.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - "conversation": "conv_123", - } - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.moderation_result = ShieldModerationPassed() - - mock_client.responses = mocker.Mock() - mock_client.responses.create = mocker.AsyncMock( - side_effect=RuntimeError("context_length exceeded") - ) - - mock_error_response = mocker.Mock() - mock_error_response.model_dump.return_value = { - "status_code": 413, - "detail": {"response": "Prompt too long", "model": "provider1/model1"}, - } - mocker.patch( - "app.endpoints.streaming_query.PromptTooLongResponse", - return_value=mock_error_response, - ) - - with pytest.raises(HTTPException) as exc_info: - await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - assert exc_info.value.status_code == 413 - - @pytest.mark.asyncio - async def test_retrieve_response_generator_runtime_error_other( - self, mocker: MockerFixture - ) -> None: - """Test response generator creation re-raises RuntimeError without context_length.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.input = "test query" - mock_responses_params.conversation = "conv_123" - mock_responses_params.model_dump.return_value = { - "input": "test query", - "model": "provider1/model1", - "conversation": "conv_123", - } - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.client = mock_client - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.moderation_result = ShieldModerationPassed() - - mock_client.responses = mocker.Mock() - mock_client.responses.create = mocker.AsyncMock( - side_effect=RuntimeError("Some other error") - ) - - with pytest.raises(RuntimeError): - await retrieve_response_generator( - mock_responses_params, mock_context, endpoint_path="" - ) - - -class TestGenerateResponse: - """Tests for generate_response function.""" - - @pytest.fixture(autouse=True) - def isolate_stream_interrupt_registry(self, mocker: MockerFixture) -> Any: - """Patch registry accessor with a per-test mock registry instance.""" - test_registry = mocker.Mock(spec=StreamInterruptRegistry) - mocker.patch( - "utils.stream_interrupts.get_stream_interrupt_registry", - return_value=test_registry, - ) - return test_registry - - @pytest.mark.asyncio - async def test_generate_response_success(self, mocker: MockerFixture) -> None: - """Test successful response generation.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - yield "data: end\n\n" - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.user_id = "user_123" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.request_id = "123e4567-e89b-12d3-a456-426614174000" - - mock_response_obj = mocker.Mock() - mock_response_obj.output = [] - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_context.client.responses = mocker.Mock() - mock_context.client.responses.create = mocker.AsyncMock( - return_value=mock_response_obj - ) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - - mock_turn_summary = TurnSummary() - mock_turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - - mock_config = mocker.Mock() - mock_config.quota_limiters = [] - mocker.patch("app.endpoints.streaming_query.configuration", mock_config) - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - mocker.patch( - "app.endpoints.streaming_query.get_available_quotas", return_value={} - ) - mocker.patch("app.endpoints.streaming_query.store_query_results") - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert len(result) > 0 - assert any("start" in item for item in result) - assert any("end" in item for item in result) - - @pytest.mark.asyncio - async def test_generate_response_compacted_persists_structured_turn( - self, mocker: MockerFixture - ) -> None: - """Compacted mode persists the turn via store_compacted_turn with the - original input and structured output items, not flattened strings - (LCORE-1572).""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - - conv_id = "123e4567-e89b-12d3-a456-426614174000" - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = conv_id - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test", conversation_id=conv_id - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.request_id = "223e4567-e89b-12d3-a456-426614174000" - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = conv_id - - turn_summary = TurnSummary() - turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - output_item = mocker.Mock() - turn_summary.output_items = [output_item] - - mock_config = mocker.Mock() - mock_config.quota_limiters = [] - mocker.patch("app.endpoints.streaming_query.configuration", mock_config) - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - mocker.patch( - "app.endpoints.streaming_query.get_available_quotas", return_value={} - ) - mocker.patch("app.endpoints.streaming_query.store_query_results") - store_mock = mocker.patch( - "app.endpoints.streaming_query.store_compacted_turn", - new_callable=mocker.AsyncMock, - ) - - result = [ - item - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - turn_summary, - compacted=True, - original_input="the original query", - ) - ] - - assert any("end" in item for item in result) - store_mock.assert_awaited_once_with( - mock_context.client, - conv_id, - "the original query", - [output_item], - ) - - @pytest.mark.asyncio - async def test_generate_response_with_topic_summary( - self, mocker: MockerFixture - ) -> None: - """Test response generation with topic summary.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.user_id = "user_123" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.query_request = QueryRequest( - query="test", generate_topic_summary=True - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.request_id = "123e4567-e89b-12d3-a456-426614174000" - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - - mock_turn_summary = TurnSummary() - mock_turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - - mock_config = mocker.Mock() - mock_config.quota_limiters = [] - mocker.patch("app.endpoints.streaming_query.configuration", mock_config) - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - mocker.patch( - "app.endpoints.streaming_query.get_available_quotas", return_value={} - ) - mocker.patch( - "app.endpoints.streaming_query.get_topic_summary", - new=mocker.AsyncMock(return_value="Topic summary"), - ) - mocker.patch("app.endpoints.streaming_query.store_query_results") - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert len(result) > 0 - - @pytest.mark.asyncio - async def test_generate_response_connection_error( - self, mocker: MockerFixture - ) -> None: - """Test response generation handles connection error.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise APIConnectionError(message="Connection failed", request=mocker.Mock()) - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.request_id = "123e4567-e89b-12d3-a456-426614174000" - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - - mock_turn_summary = TurnSummary() - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_generate_response_api_status_error( - self, mocker: MockerFixture - ) -> None: - """Test response generation handles API status error.""" - mock_request_obj = mocker.Mock() - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise APIStatusError( - message="API error", response=mock_request_obj, body=None - ) - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test" - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.request_id = "123e4567-e89b-12d3-a456-426614174000" - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - - mock_turn_summary = TurnSummary() - - mock_error_response = InternalServerErrorResponse.query_failed("API error") - mocker.patch( - "app.endpoints.streaming_query.handle_known_apistatus_errors", - return_value=mock_error_response, - ) - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_generate_response_runtime_error_context_length( - self, mocker: MockerFixture - ) -> None: - """Test generate_response handles RuntimeError with context_length.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: start\n\n" - raise RuntimeError("context_length exceeded") - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.request_id = "123e4567-e89b-12d3-a456-426614174000" - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - - mock_turn_summary = TurnSummary() - - mock_error_response = mocker.Mock() - mock_error_response.status_code = 413 - mock_error_response.detail = mocker.Mock() - mock_error_response.detail.response = "Prompt too long" - mock_error_response.detail.cause = None - mocker.patch( - "app.endpoints.streaming_query.PromptTooLongResponse", - return_value=mock_error_response, - ) - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_generate_response_runtime_error_other( - self, mocker: MockerFixture - ) -> None: - """Test generate_response handles RuntimeError without context_length.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: start\n\n" - raise RuntimeError("Some other error") - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.request_id = "123e4567-e89b-12d3-a456-426614174000" - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - - mock_turn_summary = TurnSummary() - - mock_error_response = mocker.Mock() - mock_error_response.status_code = 500 - mock_error_response.detail = mocker.Mock() - mock_error_response.detail.response = "Internal server error" - mock_error_response.detail.cause = None - mocker.patch( - "app.endpoints.streaming_query.InternalServerErrorResponse.generic", - return_value=mock_error_response, - ) - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_generate_response_cancelled_persists_interrupted_turn( - self, - mocker: MockerFixture, - isolate_stream_interrupt_registry: Any, - ) -> None: - """Test cancelled stream persists user query with interrupted response.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise asyncio.CancelledError() - - existing_conv_id = "123e4567-e89b-12d3-a456-426614174000" - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = existing_conv_id - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test", - media_type=MEDIA_TYPE_JSON, - conversation_id=existing_conv_id, # Existing conversation: no topic summary - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = existing_conv_id - mock_responses_params.input = "test" - - mock_turn_summary = TurnSummary() - mock_turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - - consume_query_tokens_mock = mocker.patch( - "app.endpoints.streaming_query.consume_query_tokens" - ) - store_query_results_mock = mocker.patch( - "utils.stream_interrupts.store_query_results" - ) - append_turn_mock = mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new_callable=mocker.AsyncMock, - ) - - test_request_id = "223e4567-e89b-12d3-a456-426614174000" - mock_context.request_id = test_request_id - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert any("start" in item for item in result) - assert any('"event": "token"' in item for item in result) - assert any('"event": "interrupted"' in item for item in result) - assert not any('"event": "end"' in item for item in result) - consume_query_tokens_mock.assert_not_called() - - append_turn_mock.assert_called_once_with( - mock_context.client, - existing_conv_id, - "test", - INTERRUPTED_INDICATOR, - ) - store_query_results_mock.assert_called_once() - call_kwargs = store_query_results_mock.call_args[1] - assert call_kwargs["user_id"] == "user_123" - assert call_kwargs["conversation_id"] == existing_conv_id - assert call_kwargs["summary"].llm_response == INTERRUPTED_INDICATOR - assert call_kwargs["topic_summary"] is None - - isolate_stream_interrupt_registry.deregister_stream.assert_called_once_with( - test_request_id - ) - - @pytest.mark.asyncio - async def test_generate_response_cancelled_persists_topic_summary_for_new_conversation( - self, - mocker: MockerFixture, - isolate_stream_interrupt_registry: Any, - ) -> None: - """Test cancelled stream persists topic_summary when generate_topic_summary is True.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise asyncio.CancelledError() - - test_request_id = "123e4567-e89b-12d3-a456-426614174001" - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_new_456" - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="What is Kubernetes?", - media_type=MEDIA_TYPE_JSON, - conversation_id=None, # New conversation - generate_topic_summary=True, - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_context.request_id = test_request_id - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = "conv_new_456" - mock_responses_params.input = "What is Kubernetes?" - - mock_turn_summary = TurnSummary() - mock_turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - get_topic_summary_mock = mocker.patch( - "utils.stream_interrupts.get_topic_summary", - new=mocker.AsyncMock(return_value="Kubernetes container orchestration"), - ) - store_query_results_mock = mocker.patch( - "utils.stream_interrupts.store_query_results" - ) - update_topic_summary_mock = mocker.patch( - "utils.stream_interrupts.update_conversation_topic_summary" - ) - mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new_callable=mocker.AsyncMock, - ) - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - await asyncio.sleep(0.1) - - assert any('"event": "interrupted"' in item for item in result) - call_kwargs = store_query_results_mock.call_args[1] - assert call_kwargs["topic_summary"] is None - get_topic_summary_mock.assert_called_once_with( - "What is Kubernetes?", - mock_context.client, - "provider1/model1", - ) - update_topic_summary_mock.assert_called_once_with( - "conv_new_456", - "Kubernetes container orchestration", - user_id="user_123", - skip_userid_check=False, - ) - isolate_stream_interrupt_registry.deregister_stream.assert_called_once_with( - test_request_id - ) - - @pytest.mark.asyncio - async def test_generate_response_cancelled_topic_summary_none_when_get_fails( - self, - mocker: MockerFixture, - isolate_stream_interrupt_registry: Any, - ) -> None: - """Test cancelled stream persists with topic_summary=None when get_topic_summary raises.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise asyncio.CancelledError() - - test_request_id = "123e4567-e89b-12d3-a456-426614174001" - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_new_456" - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="What is Kubernetes?", - media_type=MEDIA_TYPE_JSON, - conversation_id=None, # New conversation - generate_topic_summary=True, - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_context.request_id = test_request_id - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = "conv_new_456" - mock_responses_params.input = "What is Kubernetes?" - - mock_turn_summary = TurnSummary() - mock_turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - mocker.patch( - "utils.stream_interrupts.get_topic_summary", - new=mocker.AsyncMock(side_effect=Exception("err")), - ) - store_query_results_mock = mocker.patch( - "utils.stream_interrupts.store_query_results" - ) - mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new_callable=mocker.AsyncMock, - ) - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - await asyncio.sleep(0.1) - - assert any('"event": "interrupted"' in item for item in result) - store_query_results_mock.assert_called_once() - call_kwargs = store_query_results_mock.call_args[1] - assert call_kwargs["topic_summary"] is None - isolate_stream_interrupt_registry.deregister_stream.assert_called_once_with( - test_request_id - ) - - @pytest.mark.asyncio - async def test_generate_response_cancelled_topic_summary_none_when_generate_disabled( - self, - mocker: MockerFixture, - isolate_stream_interrupt_registry: Any, - ) -> None: - """Test cancelled stream uses topic_summary=None when generate_topic_summary is False.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise asyncio.CancelledError() - - test_request_id = "123e4567-e89b-12d3-a456-426614174002" - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_new_789" - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="What is Docker?", - media_type=MEDIA_TYPE_JSON, - conversation_id=None, # New conversation - generate_topic_summary=False, # Explicitly disabled - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - mock_context.request_id = test_request_id - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = "conv_new_789" - mock_responses_params.input = "What is Docker?" - - mock_turn_summary = TurnSummary() - mock_turn_summary.token_usage = TokenCounter(input_tokens=10, output_tokens=5) - - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - get_topic_summary_mock = mocker.patch( - "utils.stream_interrupts.get_topic_summary", - new=mocker.AsyncMock(return_value="Docker containerization"), - ) - store_query_results_mock = mocker.patch( - "utils.stream_interrupts.store_query_results" - ) - mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new_callable=mocker.AsyncMock, - ) - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert any('"event": "interrupted"' in item for item in result) - get_topic_summary_mock.assert_not_called() - call_kwargs = store_query_results_mock.call_args[1] - assert call_kwargs["topic_summary"] is None - isolate_stream_interrupt_registry.deregister_stream.assert_called_once_with( - test_request_id - ) - - @pytest.mark.asyncio - async def test_generate_response_cancelled_stores_results_when_append_fails( - self, - mocker: MockerFixture, - isolate_stream_interrupt_registry: Any, - ) -> None: - """Test store_query_results still runs when append_turn_to_conversation fails.""" - - async def mock_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - raise asyncio.CancelledError() - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = "conv_123" - mock_responses_params.input = "test" - - mock_turn_summary = TurnSummary() - - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - store_query_results_mock = mocker.patch( - "utils.stream_interrupts.store_query_results" - ) - mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new_callable=mocker.AsyncMock, - side_effect=RuntimeError("Llama Stack unavailable"), - ) - - test_request_id = "123e4567-e89b-12d3-a456-426614174000" - mock_context.request_id = test_request_id - - result = [] - async for item in generate_response( - mock_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - assert any('"event": "interrupted"' in item for item in result) - store_query_results_mock.assert_called_once() - isolate_stream_interrupt_registry.deregister_stream.assert_called_once_with( - test_request_id - ) - - @pytest.mark.asyncio - async def test_generate_response_task_cancel_persists_results( - self, - mocker: MockerFixture, - isolate_stream_interrupt_registry: Any, - ) -> None: - """Test that real task.cancel() persists via CancelledError handler.""" - cancel_event = asyncio.Event() - - async def slow_generator() -> AsyncIterator[str]: - yield "data: token\n\n" - await cancel_event.wait() - yield "data: should not reach\n\n" - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.conversation_id = "conv_123" - mock_context.user_id = "user_123" - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.started_at = "2024-01-01T00:00:00Z" - mock_context.skip_userid_check = False - mock_context.client = mocker.AsyncMock(spec=AsyncLlamaStackClient) - - mock_responses_params = mocker.Mock(spec=ResponsesApiParams) - mock_responses_params.model = "provider1/model1" - mock_responses_params.conversation = "conv_123" - mock_responses_params.input = "test" - - mock_turn_summary = TurnSummary() - - mocker.patch("app.endpoints.streaming_query.consume_query_tokens") - store_query_results_mock = mocker.patch( - "utils.stream_interrupts.store_query_results" - ) - append_turn_mock = mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new_callable=mocker.AsyncMock, - ) - - test_request_id = "123e4567-e89b-12d3-a456-426614174000" - mock_context.request_id = test_request_id - - result: list[str] = [] - - async def consume_generator() -> None: - async for item in generate_response( - slow_generator(), - mock_context, - mock_responses_params, - mock_turn_summary, - ): - result.append(item) - - task = asyncio.create_task(consume_generator()) - await asyncio.sleep(0.05) - task.cancel() - await asyncio.sleep(0.05) - - assert any('"event": "interrupted"' in item for item in result) - append_turn_mock.assert_called_once() - store_query_results_mock.assert_called_once() - isolate_stream_interrupt_registry.deregister_stream.assert_called_once_with( - test_request_id - ) - - @pytest.mark.asyncio - async def test_cancel_stream_callback_persists_when_error_hits_outside_generator( - self, - ) -> None: - """Test on_interrupt callback runs via cancel_stream as a separate task.""" - registry = StreamInterruptRegistry() - test_request_id = "123e4567-e89b-12d3-a456-426614174099" - registry.deregister_stream(test_request_id) - - callback_ran = False - - async def mock_callback() -> None: - nonlocal callback_ran - callback_ran = True - - async def pending_stream() -> None: - await asyncio.sleep(10) - - task = asyncio.create_task(pending_stream()) - registry.register_stream( - test_request_id, "user_123", task, on_interrupt=mock_callback - ) - - result = registry.cancel_stream(test_request_id, "user_123") - assert result.value == "cancelled" - - # Let the scheduled callback task execute - await asyncio.sleep(0.01) - - assert callback_ran is True - - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - registry.deregister_stream(test_request_id) - - -class TestResponseGenerator: - """Tests for response_generator function.""" - - @pytest.mark.asyncio - async def test_response_generator_text_delta(self, mocker: MockerFixture) -> None: - """Test response generator processes text delta events.""" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=TextDeltaChunk) - chunk.type = "response.output_text.delta" - chunk.delta = "Hello" - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - - @pytest.mark.asyncio - async def test_response_generator_content_part_added( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes content part added events.""" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock() - chunk.type = "response.content_part.added" - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - - @pytest.mark.asyncio - async def test_response_generator_output_text_done( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes output text done events.""" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=TextDoneChunk) - chunk.type = "response.output_text.done" - chunk.text = "Complete response" - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - async for _ in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - pass - - assert mock_turn_summary.llm_response == "Complete response" - - @pytest.mark.asyncio - async def test_response_generator_output_item_done_message_type( - self, mocker: MockerFixture - ) -> None: - """Test response generator skips message type items.""" - mock_output_item = mocker.Mock() - mock_output_item.type = "message" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=OutputItemDoneChunk) - chunk.type = "response.output_item.done" - chunk.item = mock_output_item - chunk.output_index = 0 - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) >= 0 - - @pytest.mark.asyncio - async def test_response_generator_output_item_done( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes output item done events.""" - mock_output_item = mocker.Mock() - mock_output_item.type = "tool_call" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=OutputItemDoneChunk) - chunk.type = "response.output_item.done" - chunk.item = mock_output_item - chunk.output_index = 0 - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mock_tool_call = mocker.Mock() - mock_tool_call.model_dump.return_value = {"tool": "test"} - mocker.patch( - "app.endpoints.streaming_query.build_tool_call_summary", - return_value=(mock_tool_call, None), - ) - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - - @pytest.mark.asyncio - async def test_response_generator_output_item_done_with_tool_result( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes output item done events with tool result.""" - mock_output_item = mocker.Mock() - mock_output_item.type = "tool_call" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=OutputItemDoneChunk) - chunk.type = "response.output_item.done" - chunk.item = mock_output_item - chunk.output_index = 0 - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mock_tool_call = mocker.Mock() - mock_tool_call.model_dump.return_value = {"tool": "test"} - mock_tool_result = mocker.Mock() - mock_tool_result.model_dump.return_value = {"result": "test_result"} - mocker.patch( - "app.endpoints.streaming_query.build_tool_call_summary", - return_value=(mock_tool_call, mock_tool_result), - ) - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - assert len(mock_turn_summary.tool_results) == 1 - - @pytest.mark.asyncio - async def test_response_generator_response_completed( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes response completed events.""" - mock_response_obj = mocker.Mock(spec=OpenAIResponseObject) - mock_response_obj.usage = mocker.Mock(input_tokens=10, output_tokens=5) - mock_response_obj.output = [] - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=CompletedChunk) - chunk.type = "response.completed" - chunk.response = mock_response_obj - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - mock_turn_summary.llm_response = "Response" - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=10, output_tokens=5), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - async for _ in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - pass - - assert mock_turn_summary.token_usage.input_tokens == 10 - - @pytest.mark.asyncio - async def test_response_generator_response_completed_uses_text_parts( - self, mocker: MockerFixture - ) -> None: - """Test response generator uses text_parts when llm_response is empty.""" - mock_response_obj = mocker.Mock(spec=OpenAIResponseObject) - mock_response_obj.usage = mocker.Mock(input_tokens=10, output_tokens=5) - mock_response_obj.output = [] - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - # Add text delta first - delta_chunk = mocker.Mock(spec=TextDeltaChunk) - delta_chunk.type = "response.output_text.delta" - delta_chunk.delta = "Hello" - yield delta_chunk - - # Then completed (without output_text.done, so llm_response is empty) - completed_chunk = mocker.Mock(spec=CompletedChunk) - completed_chunk.type = "response.completed" - completed_chunk.response = mock_response_obj - yield completed_chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=10, output_tokens=5), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - # Should use text_parts for turn_complete event - assert len(result) > 0 - assert any("turn_complete" in item for item in result) - - @pytest.mark.asyncio - async def test_response_generator_response_incomplete( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes incomplete response events.""" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=IncompleteChunk) - chunk.type = "response.incomplete" - mock_response = mocker.Mock() - mock_response.output = [] - # Create a simple object with message attribute as a string - mock_error = type("Error", (), {"message": "context_length exceeded"})() - mock_response.error = mock_error - chunk.response = mock_response - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_response_generator_response_failed( - self, mocker: MockerFixture - ) -> None: - """Test response generator processes failed response events.""" - mock_error = mocker.Mock() - mock_error.message = "Error message" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=FailedChunk) - chunk.type = "response.failed" - mock_response = mocker.Mock() - mock_response.output = [] - mock_response.error = mock_error - chunk.response = mock_response - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_response_generator_response_failed_no_error( - self, mocker: MockerFixture - ) -> None: - """Test response generator handles failed response with no error object.""" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=FailedChunk) - chunk.type = "response.failed" - mock_response = mocker.Mock() - mock_response.output = [] - mock_response.error = None - chunk.response = mock_response - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - - @pytest.mark.asyncio - async def test_response_generator_response_failed_context_length( - self, mocker: MockerFixture - ) -> None: - """Test response generator handles failed response with context_length error.""" - mock_error = mocker.Mock() - mock_error.message = "context_length exceeded" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=FailedChunk) - chunk.type = "response.failed" - mock_response = mocker.Mock() - mock_response.output = [] - mock_response.error = mock_error - chunk.response = mock_response - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - - @pytest.mark.asyncio - async def test_response_generator_response_incomplete_no_error( - self, mocker: MockerFixture - ) -> None: - """Test response generator handles incomplete response with no error object.""" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=IncompleteChunk) - chunk.type = "response.incomplete" - mock_response = mocker.Mock() - mock_response.output = [] - mock_response.error = None - chunk.response = mock_response - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - assert len(result) > 0 - assert any("error" in item for item in result) - - @pytest.mark.asyncio - async def test_response_generator_merges_inline_and_tool_rag_chunks_and_documents( - self, mocker: MockerFixture - ) -> None: - """Test that inline RAG and tool-based RAG chunks/docs are correctly merged.""" - inline_chunk = RAGChunk(content="inline chunk content", source="byok") - inline_doc = ReferencedDocument(doc_title="Inline Doc") - inline_rag = RAGContext( - context_text="", - rag_chunks=[inline_chunk], - referenced_documents=[inline_doc], - ) - - tool_chunk = RAGChunk(content="tool chunk content", source="vs-1") - tool_ref_doc = ReferencedDocument(doc_title="Tool Doc") - - mock_response_obj = mocker.Mock(spec=OpenAIResponseObject) - mock_response_obj.usage = mocker.Mock() - mock_response_obj.output = [] - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - completed_chunk = mocker.Mock(spec=CompletedChunk) - completed_chunk.type = "response.completed" - completed_chunk.response = mock_response_obj - yield completed_chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = inline_rag - - mock_turn_summary = TurnSummary() - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", - return_value=[tool_ref_doc], - ) - mocker.patch( - "app.endpoints.streaming_query.parse_rag_chunks", - return_value=[tool_chunk], - ) - - async for _ in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - pass - - assert len(mock_turn_summary.rag_chunks) == 2 - assert mock_turn_summary.rag_chunks[0].content == "inline chunk content" - assert mock_turn_summary.rag_chunks[1].content == "tool chunk content" - assert len(mock_turn_summary.referenced_documents) == 2 - assert mock_turn_summary.referenced_documents[0].doc_title == "Inline Doc" - assert mock_turn_summary.referenced_documents[1].doc_title == "Tool Doc" - - -class TestResponseGeneratorMCPCalls: - """Tests for MCP call specific event handling in response_generator.""" - - @pytest.mark.asyncio - async def test_response_generator_mcp_call_output_item_added( - self, mocker: MockerFixture - ) -> None: - """Test response generator stores MCP call item info when output_item.added.""" - mock_mcp_item = mocker.Mock(spec=MCPCall) - mock_mcp_item.type = "mcp_call" - mock_mcp_item.id = "mcp_call_123" - mock_mcp_item.name = "test_mcp_tool" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=OutputItemAddedChunk) - chunk.type = "response.output_item.added" - chunk.item = mock_mcp_item - chunk.output_index = 0 - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - # Should process without error - assert True - - @pytest.mark.asyncio - async def test_response_generator_mcp_call_arguments_done( - self, mocker: MockerFixture - ) -> None: - """Test response generator emits tool call when MCP arguments.done.""" - mock_mcp_item = mocker.Mock(spec=MCPCall) - mock_mcp_item.type = "mcp_call" - mock_mcp_item.id = "mcp_call_123" - mock_mcp_item.name = "test_mcp_tool" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - # First, output_item.added - added_chunk = mocker.Mock(spec=OutputItemAddedChunk) - added_chunk.type = "response.output_item.added" - added_chunk.item = mock_mcp_item - added_chunk.output_index = 0 - yield added_chunk - - # Then, arguments.done - args_chunk = mocker.Mock(spec=MCPArgsDoneChunk) - args_chunk.type = "response.mcp_call.arguments.done" - args_chunk.output_index = 0 - args_chunk.arguments = '{"param": "value"}' - yield args_chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mock_tool_call = mocker.Mock() - mock_tool_call.model_dump.return_value = { - "id": "mcp_call_123", - "name": "test_mcp_tool", - } - mocker.patch( - "app.endpoints.streaming_query.build_mcp_tool_call_from_arguments_done", - return_value=mock_tool_call, - ) - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - # Should emit tool call event - assert len(result) > 0 - assert len(mock_turn_summary.tool_calls) == 1 - - @pytest.mark.asyncio - async def test_response_generator_mcp_call_output_item_done_with_arguments_done( - self, mocker: MockerFixture - ) -> None: - """Test response generator emits only result when MCP output_item.done after arguments.""" - mock_mcp_item = mocker.Mock(spec=MCPCall) - mock_mcp_item.type = "mcp_call" - mock_mcp_item.id = "mcp_call_123" - mock_mcp_item.name = "test_mcp_tool" - mock_mcp_item.error = None - mock_mcp_item.output = "Result output" - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - # First, output_item.added - added_chunk = mocker.Mock(spec=OutputItemAddedChunk) - added_chunk.type = "response.output_item.added" - added_chunk.item = mock_mcp_item - added_chunk.output_index = 0 - yield added_chunk - - # Then, arguments.done (removes from mcp_calls dict) - args_chunk = mocker.Mock(spec=MCPArgsDoneChunk) - args_chunk.type = "response.mcp_call.arguments.done" - args_chunk.output_index = 0 - args_chunk.arguments = '{"param": "value"}' - yield args_chunk - - # Finally, output_item.done (should only emit result) - done_chunk = mocker.Mock(spec=OutputItemDoneChunk) - done_chunk.type = "response.output_item.done" - done_chunk.item = mock_mcp_item - done_chunk.output_index = 0 - yield done_chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mock_tool_call = mocker.Mock() - mock_tool_call.model_dump.return_value = {"id": "mcp_call_123"} - - # Use side_effect to actually remove item from mcp_calls dict - def build_mcp_tool_call_side_effect( - output_index: int, - arguments: str, - mcp_call_items: dict[int, tuple[str, str]], - ) -> Any: - # Remove item from dict to simulate real behavior - # arguments parameter is required by function signature but unused here - _ = arguments - mcp_call_items.pop(output_index, None) - return mock_tool_call - - mocker.patch( - "app.endpoints.streaming_query.build_mcp_tool_call_from_arguments_done", - side_effect=build_mcp_tool_call_side_effect, - ) - - mock_tool_result = mocker.Mock() - mock_tool_result.model_dump.return_value = { - "id": "mcp_call_123", - "status": "success", - } - mocker.patch( - "app.endpoints.streaming_query.build_tool_result_from_mcp_output_item_done", - return_value=mock_tool_result, - ) - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - # Should have one tool call (from arguments.done) and one tool result - assert len(mock_turn_summary.tool_calls) == 1 - assert len(mock_turn_summary.tool_results) == 1 - - @pytest.mark.asyncio - async def test_response_generator_mcp_call_output_item_done_without_arguments_done( - self, mocker: MockerFixture - ) -> None: - """Test response generator emits both call and result when MCP output_item.done.""" - mock_mcp_item = mocker.Mock(spec=MCPCall) - mock_mcp_item.type = "mcp_call" - mock_mcp_item.id = "mcp_call_123" - mock_mcp_item.name = "test_mcp_tool" - mock_mcp_item.error = None - mock_mcp_item.output = "Result output" - mock_mcp_item.arguments = '{"param": "value"}' - mock_mcp_item.server_label = None - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - # Only output_item.added (arguments.done was missed) - added_chunk = mocker.Mock(spec=OutputItemAddedChunk) - added_chunk.type = "response.output_item.added" - added_chunk.item = mock_mcp_item - added_chunk.output_index = 0 - yield added_chunk - - # output_item.done (should emit both call and result since arguments.done didn't happen) - done_chunk = mocker.Mock(spec=OutputItemDoneChunk) - done_chunk.type = "response.output_item.done" - done_chunk.item = mock_mcp_item - done_chunk.output_index = 0 - yield done_chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - mock_turn_summary = TurnSummary() - - mock_tool_call = mocker.Mock() - mock_tool_call.model_dump.return_value = {"id": "mcp_call_123"} - mock_tool_result = mocker.Mock() - mock_tool_result.model_dump.return_value = { - "id": "mcp_call_123", - "status": "success", - } - mocker.patch( - "app.endpoints.streaming_query.build_tool_call_summary", - return_value=(mock_tool_call, mock_tool_result), - ) - - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - result = [] - async for item in response_generator( - mock_turn_response(), mock_context, mock_turn_summary, endpoint_path="" - ): - result.append(item) - - # Should have both tool call and result (fallback behavior) - assert len(mock_turn_summary.tool_calls) == 1 - assert len(mock_turn_summary.tool_results) == 1 - - -@pytest.mark.asyncio -async def test_response_generator_failed_captures_output_items( - mocker: MockerFixture, -) -> None: - """A failed terminal captures output_items for compacted persistence (LCORE-1572).""" - out_item = mocker.Mock() - - async def mock_turn_response() -> AsyncIterator[OpenAIResponseObjectStream]: - chunk = mocker.Mock(spec=FailedChunk) - chunk.type = "response.failed" - mock_response = mocker.Mock() - mock_response.output = [out_item] - mock_response.error = mocker.Mock(message="boom") - chunk.response = mock_response - yield chunk - - mock_context = mocker.Mock(spec=ResponseGeneratorContext) - mock_context.query_request = QueryRequest( - query="test", media_type=MEDIA_TYPE_JSON - ) # pyright: ignore[reportCallIssue] - mock_context.model_id = "provider1/model1" - mock_context.vector_store_ids = [] - mock_context.rag_id_mapping = {} - mock_context.inline_rag_context = RAGContext() - - turn_summary = TurnSummary() - mocker.patch( - "app.endpoints.streaming_query.extract_token_usage", - return_value=TokenCounter(input_tokens=0, output_tokens=0), - ) - mocker.patch( - "app.endpoints.streaming_query.parse_referenced_documents", return_value=[] - ) - - async for _ in response_generator( - mock_turn_response(), mock_context, turn_summary, endpoint_path="" - ): - pass - - assert turn_summary.output_items == [out_item] diff --git a/tests/unit/app/endpoints/test_tools.py b/tests/unit/app/endpoints/test_tools.py index ead6880de..3ab9e4be6 100644 --- a/tests/unit/app/endpoints/test_tools.py +++ b/tests/unit/app/endpoints/test_tools.py @@ -1,22 +1,19 @@ -# pylint: disable=protected-access,too-many-lines +# pylint: disable=protected-access,redefined-outer-name """Unit tests for tools endpoint.""" from pathlib import Path -from typing import Any +from typing import Optional import pytest -from fastapi import HTTPException -from llama_stack_client import APIConnectionError, BadRequestError from pydantic import AnyHttpUrl, SecretStr -from pytest_mock import MockerFixture, MockType +from pytest_mock import MockerFixture -# Import the function directly to bypass decorators from app.endpoints import tools -from app.endpoints.tools import _input_schema_to_parameters from authentication.interface import AuthTuple from configuration import AppConfig from models.api.responses.successful import ToolsResponse +from models.common.tools import ListedMcpTool from models.config import ( Configuration, CORSConfiguration, @@ -27,11 +24,35 @@ TLSConfiguration, UserDataCollection, ) +from utils.builtin_tools import FILE_SEARCH_CATALOG_TOOLS -# Shared mock auth tuple with 4 fields as expected by the application MOCK_AUTH: AuthTuple = ("mock_user_id", "mock_username", False, "mock_token") +def _mock_file_search_tools( + mocker: MockerFixture, file_search_tools: Optional[list] = None +) -> None: + """Patch LLS file-search discovery for tools tests.""" + mocker.patch("app.endpoints.tools.AsyncOgxClientHolder") + mocker.patch( + "app.endpoints.tools.get_file_search_tools", + new_callable=mocker.AsyncMock, + return_value=( + FILE_SEARCH_CATALOG_TOOLS + if file_search_tools is None + else file_search_tools + ), + ) + + +def _make_app_config(mocker: MockerFixture, config: Configuration) -> AppConfig: + """Create an AppConfig with the given configuration and patch it.""" + app_config = AppConfig() + app_config._configuration = config + mocker.patch("app.endpoints.tools.configuration", app_config) + return app_config + + @pytest.fixture def mock_configuration() -> Configuration: """Create a mock configuration with MCP servers.""" @@ -89,1053 +110,210 @@ def mock_configuration() -> Configuration: ) -def _make_tool_def_mock(mocker: MockerFixture, fields: dict[str, Any]) -> MockType: - """Create a mock ToolDef object matching Llama Stack's tools.list() output. - - The mock supports ``dict()`` conversion so the endpoint can do - ``tool_dict = dict(tool)`` and get a plain dict back. - """ - mock = mocker.Mock() - mock.__dict__.update(fields) - mock.keys.return_value = fields.keys() - mock.__getitem__ = lambda self, key: self.__dict__[key] - mock.__iter__ = lambda self: iter(self.__dict__) - return mock - - -@pytest.fixture -def mock_tools_response(mocker: MockerFixture) -> list[MockType]: - """Create mock tools response matching Llama Stack ToolDef format. - - Each mock uses ToolDef field names: ``name`` (not ``identifier``), - ``input_schema`` (not ``parameters``), and no ``provider_id`` or ``type`` - (those live on the toolgroup, not on individual tools). - - Returns: - list[MockType]: Two mock ToolDef objects for filesystem and git tools. - """ - tool1 = _make_tool_def_mock( - mocker, - { - "name": "filesystem_read", - "description": "Read contents of a file from the filesystem", - "input_schema": { - "type": "object", - "properties": { - "path": { - "type": "string", - "description": "Path to the file to read", - } - }, - "required": ["path"], - }, - "toolgroup_id": "filesystem-tools", - "metadata": {}, - "output_schema": None, - }, - ) - - tool2 = _make_tool_def_mock( - mocker, - { - "name": "git_status", - "description": "Get the status of a git repository", - "input_schema": { - "type": "object", - "properties": { - "repository_path": { - "type": "string", - "description": "Path to the git repository", - } - }, - "required": ["repository_path"], - }, - "toolgroup_id": "git-tools", - "metadata": {}, - "output_schema": None, - }, - ) - - return [tool1, tool2] - - @pytest.mark.asyncio -async def test_tools_endpoint_success( +async def test_tools_lists_builtin_and_mcp_tools( mocker: MockerFixture, - mock_configuration: Configuration, # pylint: disable=redefined-outer-name - mock_tools_response: list[MockType], # pylint: disable=redefined-outer-name + mock_configuration: Configuration, ) -> None: - """Test successful tools endpoint response.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list response (toolgroups carry provider_id and type) - mock_toolgroup1 = mocker.Mock() - mock_toolgroup1.identifier = "filesystem-tools" - mock_toolgroup1.provider_id = "model-context-protocol" - mock_toolgroup1.type = "tool_group" - mock_toolgroup2 = mocker.Mock() - mock_toolgroup2.identifier = "git-tools" - mock_toolgroup2.provider_id = "model-context-protocol" - mock_toolgroup2.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup1, mock_toolgroup2] - - # Mock tools.list responses for each MCP server - mock_client.tools.list.side_effect = [ - [mock_tools_response[0]], # filesystem-tools response - [mock_tools_response[1]], # git-tools response - ] - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpoint - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - # Verify response - assert isinstance(response, ToolsResponse) - assert len(response.tools) == 2 - - # Verify first tool - tool1 = response.tools[0] - assert tool1["identifier"] == "filesystem_read" - assert tool1["description"] == "Read contents of a file from the filesystem" - assert tool1["server_source"] == "http://localhost:3000" - assert tool1["toolgroup_id"] == "filesystem-tools" - assert tool1["provider_id"] == "model-context-protocol" - assert tool1["type"] == "tool_group" - assert len(tool1["parameters"]) == 1 - assert tool1["parameters"][0]["name"] == "path" - assert tool1["parameters"][0]["required"] is True - - # Verify second tool - tool2 = response.tools[1] - assert tool2["identifier"] == "git_status" - assert tool2["description"] == "Get the status of a git repository" - assert tool2["server_source"] == "http://localhost:3001" - assert tool2["toolgroup_id"] == "git-tools" - assert tool2["provider_id"] == "model-context-protocol" - - # Verify client calls - assert mock_client.tools.list.call_count == 2 - mock_client.tools.list.assert_any_call( - toolgroup_id="filesystem-tools", - extra_headers={}, - extra_query={"authorization": None}, + """Return file-search tools from LLS plus MCP tools discovered locally.""" + _make_app_config(mocker, mock_configuration) + mocker.patch( + "app.endpoints.tools.check_configuration_loaded", + return_value=None, ) - mock_client.tools.list.assert_any_call( - toolgroup_id="git-tools", - extra_headers={}, - extra_query={"authorization": None}, + mocker.patch( + "app.endpoints.tools.build_mcp_headers", + return_value={"filesystem-tools": {}, "git-tools": {}}, ) - - -@pytest.mark.asyncio -async def test_tools_endpoint_no_mcp_servers(mocker: MockerFixture) -> None: - """Test tools endpoint with no MCP servers configured.""" - # Mock configuration with no MCP servers - wrap in AppConfig - mock_config = Configuration( - name="test", - service=ServiceConfiguration( - tls_config=TLSConfiguration( - tls_certificate_path=Path("tests/configuration/server.crt"), - tls_key_path=Path("tests/configuration/server.key"), - tls_key_password=Path("tests/configuration/password"), - ), - cors=CORSConfiguration( - allow_origins=["foo_origin", "bar_origin", "baz_origin"], - allow_credentials=False, - allow_methods=["foo_method", "bar_method", "baz_method"], - allow_headers=["foo_header", "bar_header", "baz_header"], - ), - host="localhost", - port=1234, - base_url=".", - auth_enabled=False, - workers=1, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - url=AnyHttpUrl("http://localhost:8321"), - api_key=SecretStr("xyzzy"), - use_as_library_client=False, - library_client_config_path=".", - timeout=10, - ), - user_data_collection=UserDataCollection( - transcripts_enabled=False, - feedback_enabled=False, - transcripts_storage=".", - feedback_storage=".", - ), - mcp_servers=[], - customization=None, - authorization=None, - deployment_environment=".", + mocker.patch("app.endpoints.tools.check_mcp_auth", return_value=None) + mocker.patch( + "app.endpoints.tools.get_agent_capability_tools", + return_value=[], + ) + _mock_file_search_tools(mocker) + mock_list = mocker.patch( + "app.endpoints.tools.list_mcp_tools", + side_effect=[ + [ + ListedMcpTool( + name="filesystem_read", + description="Read contents of a file from the filesystem", + input_schema={ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the file to read", + } + }, + "required": ["path"], + }, + ) + ], + [ + ListedMcpTool( + name="git_status", + description="Show working tree status", + input_schema=None, + ) + ], + ], ) - app_config = AppConfig() - app_config._configuration = mock_config - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list response - empty for no MCP servers - mock_client.toolgroups.list.return_value = [] - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH + request = mocker.Mock() + request.headers = {} - # Call the endpoint - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] + response = await tools.tools_endpoint_handler( + request, + auth=MOCK_AUTH, + mcp_headers={}, + ) - # Verify response assert isinstance(response, ToolsResponse) - assert len(response.tools) == 0 + assert mock_list.call_count == 2 + identifiers = {tool.identifier for tool in response.tools} + assert "insert_into_memory" in identifiers + assert "file_search" in identifiers + assert "filesystem_read" in identifiers + assert "git_status" in identifiers @pytest.mark.asyncio -async def test_tools_endpoint_api_connection_error( - mocker: MockerFixture, # pylint: disable=redefined-outer-name - mock_configuration: Configuration, # pylint: disable=redefined-outer-name -) -> None: - """Test tools endpoint with API connection error from individual servers.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list response - mock_toolgroup1 = mocker.Mock() - mock_toolgroup1.identifier = "filesystem-tools" - mock_toolgroup1.provider_id = "model-context-protocol" - mock_toolgroup1.type = "tool_group" - mock_toolgroup2 = mocker.Mock() - mock_toolgroup2.identifier = "git-tools" - mock_toolgroup2.provider_id = "model-context-protocol" - mock_toolgroup2.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup1, mock_toolgroup2] - - # Mock API connection error - create a proper APIConnectionError - api_error = APIConnectionError(request=mocker.Mock()) - mock_client.tools.list.side_effect = api_error - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpoint - should raise HTTPException when APIConnectionError occurs - with pytest.raises(HTTPException) as exc_info: - await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - assert exc_info.value.status_code == 503 - detail = exc_info.value.detail - assert isinstance(detail, dict) - assert detail["response"] == "Unable to connect to Llama Stack" # type: ignore - - -@pytest.mark.asyncio -async def test_tools_endpoint_partial_failure( # pylint: disable=redefined-outer-name +async def test_tools_skips_server_with_unresolved_auth( mocker: MockerFixture, mock_configuration: Configuration, ) -> None: - """Test tools endpoint with one MCP server failing with APIConnectionError.""" - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - mock_toolgroup1 = mocker.Mock() - mock_toolgroup1.identifier = "filesystem-tools" - mock_toolgroup1.provider_id = "model-context-protocol" - mock_toolgroup1.type = "tool_group" - mock_toolgroup2 = mocker.Mock() - mock_toolgroup2.identifier = "git-tools" - mock_toolgroup2.provider_id = "model-context-protocol" - mock_toolgroup2.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup1, mock_toolgroup2] - - api_error = APIConnectionError(request=mocker.Mock()) - mock_client.tools.list.side_effect = api_error - - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - with pytest.raises(HTTPException) as exc_info: - await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - assert exc_info.value.status_code == 503 - detail = exc_info.value.detail - assert isinstance(detail, dict) - assert detail["response"] == "Unable to connect to Llama Stack" # type: ignore - - -@pytest.mark.asyncio -async def test_tools_endpoint_toolgroup_not_found( # pylint: disable=redefined-outer-name - mocker: MockerFixture, - mock_configuration: Configuration, - mock_tools_response: list[MockType], -) -> None: - """Test tools endpoint when a toolgroup is not found (BadRequestError).""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list response - mock_toolgroup1 = mocker.Mock() - mock_toolgroup1.identifier = "filesystem-tools" - mock_toolgroup1.provider_id = "model-context-protocol" - mock_toolgroup1.type = "tool_group" - mock_toolgroup2 = mocker.Mock() - mock_toolgroup2.identifier = "git-tools" - mock_toolgroup2.provider_id = "model-context-protocol" - mock_toolgroup2.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup1, mock_toolgroup2] - - # Mock tools.list responses - first succeeds, second raises BadRequestError - bad_request_error = BadRequestError( - message="Toolgroup not found", - response=mocker.Mock(request=None), - body=None, - ) - mock_client.tools.list.side_effect = [ - [mock_tools_response[0]], # filesystem-tools response - bad_request_error, # git-tools not found + """Skip MCP servers when required auth headers cannot be resolved.""" + mock_configuration.mcp_servers = [ + ModelContextProtocolServer( + name="secure-tools", + provider_id="model-context-protocol", + url="http://localhost:3002", + authorization_headers={"Authorization": "client"}, + ) ] - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpoint - should continue processing and return tools from successful toolgroups - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - # Verify response - should have only one tool from the first successful toolgroup - assert isinstance(response, ToolsResponse) - assert len(response.tools) == 1 - assert response.tools[0]["identifier"] == "filesystem_read" - assert response.tools[0]["server_source"] == "http://localhost:3000" - assert response.tools[0]["provider_id"] == "model-context-protocol" - - # Verify that tools.list was called for both toolgroups - assert mock_client.tools.list.call_count == 2 - mock_client.tools.list.assert_any_call( - toolgroup_id="filesystem-tools", - extra_headers={}, - extra_query={"authorization": None}, - ) - mock_client.tools.list.assert_any_call( - toolgroup_id="git-tools", - extra_headers={}, - extra_query={"authorization": None}, - ) - - -@pytest.mark.asyncio -async def test_tools_endpoint_builtin_toolgroup( - mocker: MockerFixture, - mock_configuration: Configuration, # pylint: disable=redefined-outer-name -) -> None: - """Test tools endpoint with built-in toolgroups.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list response with built-in toolgroup - mock_toolgroup = mocker.Mock() - mock_toolgroup.identifier = "builtin-tools" # Not in MCP server names - mock_toolgroup.provider_id = "rag-runtime" - mock_toolgroup.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup] - - # Mock tools.list response for built-in toolgroup (ToolDef format) - mock_tool = _make_tool_def_mock( - mocker, - { - "name": "builtin_tool", - "description": "A built-in tool", - "input_schema": None, - "toolgroup_id": "builtin-tools", - "metadata": {}, - "output_schema": None, - }, - ) - - mock_client.tools.list.return_value = [mock_tool] - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpoint - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - # Verify response — identifier mapped from name, provider_id from toolgroup - assert isinstance(response, ToolsResponse) - assert len(response.tools) == 1 - assert response.tools[0]["identifier"] == "builtin_tool" - assert response.tools[0]["server_source"] == "builtin" - assert response.tools[0]["provider_id"] == "rag-runtime" - assert response.tools[0]["type"] == "tool_group" - assert response.tools[0]["parameters"] == [] - - -@pytest.mark.asyncio -async def test_tools_endpoint_mixed_toolgroups(mocker: MockerFixture) -> None: - """Test tools endpoint with both MCP and built-in toolgroups.""" - # Mock configuration with MCP servers - wrap in AppConfig - mock_config = Configuration( - name="test", - service=ServiceConfiguration( - tls_config=TLSConfiguration( - tls_certificate_path=Path("tests/configuration/server.crt"), - tls_key_path=Path("tests/configuration/server.key"), - tls_key_password=Path("tests/configuration/password"), - ), - cors=CORSConfiguration( - allow_origins=["foo_origin", "bar_origin", "baz_origin"], - allow_credentials=False, - allow_methods=["foo_method", "bar_method", "baz_method"], - allow_headers=["foo_header", "bar_header", "baz_header"], - ), - host="localhost", - port=1234, - base_url=".", - auth_enabled=False, - workers=1, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - url=AnyHttpUrl("http://localhost:8321"), - api_key=SecretStr("xyzzy"), - use_as_library_client=False, - library_client_config_path=".", - timeout=10, - ), - user_data_collection=UserDataCollection( - transcripts_enabled=False, - feedback_enabled=False, - transcripts_storage=".", - feedback_storage=".", - ), - mcp_servers=[ - ModelContextProtocolServer( - name="filesystem-tools", - provider_id="model-context-protocol", - url="http://localhost:3000", - ), - ], - customization=None, - authorization=None, - deployment_environment=".", + _make_app_config(mocker, mock_configuration) + mocker.patch( + "app.endpoints.tools.check_configuration_loaded", + return_value=None, ) - app_config = AppConfig() - app_config._configuration = mock_config - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list response with both MCP and built-in toolgroups - mock_toolgroup1 = mocker.Mock() - mock_toolgroup1.identifier = "filesystem-tools" # MCP server toolgroup - mock_toolgroup1.provider_id = "model-context-protocol" - mock_toolgroup1.type = "tool_group" - mock_toolgroup2 = mocker.Mock() - mock_toolgroup2.identifier = "builtin-tools" # Built-in toolgroup - mock_toolgroup2.provider_id = "rag-runtime" - mock_toolgroup2.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup1, mock_toolgroup2] - - # Mock tools.list responses (ToolDef format) - mock_tool1 = _make_tool_def_mock( - mocker, - { - "name": "filesystem_read", - "description": "Read file", - "input_schema": None, - "toolgroup_id": "filesystem-tools", - "metadata": {}, - "output_schema": None, - }, + mocker.patch( + "app.endpoints.tools.build_mcp_headers", + return_value={}, ) - - mock_tool2 = _make_tool_def_mock( - mocker, - { - "name": "builtin_tool", - "description": "Built-in tool", - "input_schema": None, - "toolgroup_id": "builtin-tools", - "metadata": {}, - "output_schema": None, - }, + mocker.patch("app.endpoints.tools.check_mcp_auth", return_value=None) + mocker.patch( + "app.endpoints.tools.get_agent_capability_tools", + return_value=[], ) + _mock_file_search_tools(mocker) + mock_list = mocker.patch("app.endpoints.tools.list_mcp_tools") - mock_client.tools.list.side_effect = [[mock_tool1], [mock_tool2]] - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpoint - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - # Verify response - should have both tools with correct server sources - assert isinstance(response, ToolsResponse) - assert len(response.tools) == 2 - - # Find tools by identifier to avoid order dependency - mcp_tool = next(t for t in response.tools if t["identifier"] == "filesystem_read") - builtin_tool = next(t for t in response.tools if t["identifier"] == "builtin_tool") - - assert mcp_tool["server_source"] == "http://localhost:3000" - assert mcp_tool["provider_id"] == "model-context-protocol" - assert builtin_tool["server_source"] == "builtin" - assert builtin_tool["provider_id"] == "rag-runtime" - - -@pytest.mark.asyncio -async def test_tools_endpoint_value_attribute_error( - mocker: MockerFixture, - mock_configuration: Configuration, # pylint: disable=redefined-outer-name -) -> None: - """Test tools endpoint with ValueError/AttributeError in toolgroups.list.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list to raise ValueError - mock_client.toolgroups.list.side_effect = ValueError("Invalid response format") - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpointt - should raise exception since toolgroups.list failed - with pytest.raises(ValueError, match="Invalid response format"): - await tools.tools_endpoint_handler.__wrapped__(mock_request, mock_auth, {}) # type: ignore - - -@pytest.mark.asyncio -async def test_tools_endpoint_apiconnection_error_toolgroups( # pylint: disable=redefined-outer-name - mocker: MockerFixture, mock_configuration: Configuration -) -> None: - """Test tools endpoint with APIConnectionError in toolgroups.list.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder and clien - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - # Mock toolgroups.list to raise APIConnectionError - api_error = APIConnectionError(request=mocker.Mock()) - mock_client.toolgroups.list.side_effect = api_error + request = mocker.Mock() + request.headers = {} - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpointt and expect HTTPException - with pytest.raises(HTTPException) as exc_info: - await tools.tools_endpoint_handler.__wrapped__(mock_request, mock_auth, {}) # type: ignore - - assert exc_info.value.status_code == 503 - - detail = exc_info.value.detail - assert isinstance(detail, dict) - assert detail["response"] == "Unable to connect to Llama Stack" # type: ignore - - -@pytest.mark.asyncio -async def test_tools_endpoint_client_holder_apiconnection_error( # pylint: disable=redefined-outer-name - mocker: MockerFixture, mock_configuration: Configuration -) -> None: - """Test tools endpoint with APIConnectionError in client holder.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder to raise APIConnectionError - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - api_error = APIConnectionError(request=None) # type: ignore - mock_client_holder.return_value.get_client.side_effect = api_error - - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpointt and expect HTTPException - with pytest.raises(HTTPException) as exc_info: - await tools.tools_endpoint_handler.__wrapped__(mock_request, mock_auth, {}) # type: ignore - - assert exc_info.value.status_code == 503 - - detail = exc_info.value.detail - assert isinstance(detail, dict) - assert detail["response"] == "Unable to connect to Llama Stack" # type: ignore - - -@pytest.mark.asyncio -async def test_tools_endpoint_general_exception( - mocker: MockerFixture, - mock_configuration: Configuration, # pylint: disable=redefined-outer-name -) -> None: - """Test tools endpoint with general exception.""" - # Mock configuration - wrap in AppConfig - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - - # Mock authorization decorator to bypass i - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - # Mock client holder to raise exception - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client_holder.return_value.get_client.side_effect = Exception( - "Unexpected error" + response = await tools.tools_endpoint_handler( + request, + auth=MOCK_AUTH, + mcp_headers={}, ) - # Mock request and auth - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - # Call the endpointt and expect the exception to propagate (not caught) - with pytest.raises(Exception, match="Unexpected error"): - await tools.tools_endpoint_handler.__wrapped__(mock_request, mock_auth, {}) # type: ignore + mock_list.assert_not_called() + identifiers = {tool.identifier for tool in response.tools} + assert identifiers == {"insert_into_memory", "file_search"} @pytest.mark.asyncio -async def test_tools_endpoint_authentication_error_with_mcp_endpoint( +async def test_tools_continues_when_mcp_list_returns_empty( mocker: MockerFixture, - mock_configuration: Configuration, # pylint: disable=redefined-outer-name + mock_configuration: Configuration, ) -> None: - """Test tools endpoint raises 401 with WWW-Authenticate when check_mcp_auth requires OAuth.""" - app_config = AppConfig() - app_config._configuration = mock_configuration - mocker.patch("app.endpoints.tools.configuration", app_config) - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - - expected_headers = {"WWW-Authenticate": 'Bearer error="invalid_token"'} - probe_exception = HTTPException( - status_code=401, - detail={"cause": "MCP server at http://localhost:3000 requires OAuth"}, - headers=expected_headers, + """Skip MCP servers that fail discovery without failing the request.""" + mock_configuration.mcp_servers = [ + ModelContextProtocolServer( + name="broken-tools", + provider_id="model-context-protocol", + url="http://localhost:3999", + ) + ] + _make_app_config(mocker, mock_configuration) + mocker.patch( + "app.endpoints.tools.check_configuration_loaded", + return_value=None, ) mocker.patch( - "app.endpoints.tools.check_mcp_auth", - new_callable=mocker.AsyncMock, - side_effect=probe_exception, + "app.endpoints.tools.build_mcp_headers", + return_value={"broken-tools": {}}, ) - - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - with pytest.raises(HTTPException) as exc_info: - await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - assert exc_info.value.status_code == 401 - assert exc_info.value.headers is not None - assert ( - exc_info.value.headers.get("WWW-Authenticate") == 'Bearer error="invalid_token"' + mocker.patch("app.endpoints.tools.check_mcp_auth", return_value=None) + mocker.patch( + "app.endpoints.tools.get_agent_capability_tools", + return_value=[], ) - - -class TestInputSchemaToParameters: - """Tests for _input_schema_to_parameters conversion.""" - - def test_none_schema(self) -> None: - """Test that None schema returns empty list.""" - assert _input_schema_to_parameters(None) == [] - - def test_empty_schema(self) -> None: - """Test that empty dict returns empty list.""" - assert _input_schema_to_parameters({}) == [] - - def test_schema_without_properties(self) -> None: - """Test that schema without properties returns empty list.""" - assert _input_schema_to_parameters({"type": "object"}) == [] - - def test_single_required_param(self) -> None: - """Test conversion of a single required parameter.""" - schema = { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "The search query", - } - }, - "required": ["query"], - } - result = _input_schema_to_parameters(schema) - assert len(result) == 1 - assert result[0]["name"] == "query" - assert result[0]["description"] == "The search query" - assert result[0]["parameter_type"] == "string" - assert result[0]["required"] is True - assert result[0]["default"] is None - - def test_optional_param_with_default(self) -> None: - """Test conversion of an optional parameter with a default value.""" - schema = { - "type": "object", - "properties": { - "limit": { - "type": "integer", - "description": "Max results", - "default": 10, - } - }, - "required": [], - } - result = _input_schema_to_parameters(schema) - assert len(result) == 1 - assert result[0]["name"] == "limit" - assert result[0]["parameter_type"] == "integer" - assert result[0]["required"] is False - assert result[0]["default"] == 10 - - def test_multiple_params_mixed_required(self) -> None: - """Test conversion with a mix of required and optional parameters.""" - schema = { - "type": "object", - "properties": { - "query": {"type": "string", "description": "Search query"}, - "limit": {"type": "integer", "description": "Max results"}, - }, - "required": ["query"], - } - result = _input_schema_to_parameters(schema) - assert len(result) == 2 - by_name = {p["name"]: p for p in result} - assert by_name["query"]["required"] is True - assert by_name["limit"]["required"] is False - - -@pytest.mark.asyncio -async def test_tools_endpoint_rag_builtin_toolgroup(mocker: MockerFixture) -> None: - """Test that builtin::rag tools have correct fields (LCORE-1211 regression). - - Reproduces the exact scenario from LCORE-1211: Llama Stack returns RAG - tools via the builtin::rag toolgroup using the ToolDef format. - Previously, identifier, provider_id, parameters, and type were all - returned as empty strings/lists. - """ - mock_config = Configuration( - name="test", - service=ServiceConfiguration( - tls_config=TLSConfiguration( - tls_certificate_path=Path("tests/configuration/server.crt"), - tls_key_path=Path("tests/configuration/server.key"), - tls_key_password=Path("tests/configuration/password"), - ), - cors=CORSConfiguration( - allow_origins=["*"], - allow_credentials=False, - allow_methods=["*"], - allow_headers=["*"], - ), - host="localhost", - port=8080, - base_url=".", - auth_enabled=False, - workers=1, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - url=AnyHttpUrl("http://localhost:8321"), - api_key=SecretStr("xyzzy"), - use_as_library_client=False, - library_client_config_path=".", - timeout=10, - ), - user_data_collection=UserDataCollection( - transcripts_enabled=False, - feedback_enabled=False, - transcripts_storage=".", - feedback_storage=".", - ), - mcp_servers=[], - customization=None, - authorization=None, - deployment_environment=".", + _mock_file_search_tools(mocker) + mocker.patch( + "app.endpoints.tools.list_mcp_tools", + return_value=[], ) - app_config = AppConfig() - app_config._configuration = mock_config - mocker.patch("app.endpoints.tools.configuration", app_config) - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client + request = mocker.Mock() + request.headers = {} - # Toolgroup matching the real Llama Stack builtin::rag - mock_toolgroup = mocker.Mock() - mock_toolgroup.identifier = "builtin::rag" - mock_toolgroup.provider_id = "rag-runtime" - mock_toolgroup.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup] - - # Tools matching real Llama Stack ToolDef output - rag_tool = _make_tool_def_mock( - mocker, - { - "name": "knowledge_search", - "description": "Search for information in a database.", - "input_schema": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "The query to search for.", - } - }, - "required": ["query"], - }, - "toolgroup_id": "builtin::rag", - "metadata": None, - "output_schema": None, - }, + response = await tools.tools_endpoint_handler( + request, + auth=MOCK_AUTH, + mcp_headers={}, ) - mock_client.tools.list.return_value = [rag_tool] - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - - assert isinstance(response, ToolsResponse) - assert len(response.tools) == 1 - - tool = response.tools[0] - assert tool["identifier"] == "knowledge_search" - assert tool["provider_id"] == "rag-runtime" - assert tool["type"] == "tool_group" - assert tool["server_source"] == "builtin" - assert tool["toolgroup_id"] == "builtin::rag" - - # Parameters converted from input_schema - assert len(tool["parameters"]) == 1 - assert tool["parameters"][0]["name"] == "query" - assert tool["parameters"][0]["parameter_type"] == "string" - assert tool["parameters"][0]["required"] is True + identifiers = {tool.identifier for tool in response.tools} + assert "filesystem_read" not in identifiers + assert "insert_into_memory" in identifiers @pytest.mark.asyncio -async def test_tools_endpoint_empty_legacy_fields_overridden( +async def test_tools_skips_mcp_server_when_discovery_returns_no_tools( mocker: MockerFixture, + mock_configuration: Configuration, ) -> None: - """Test that empty legacy fields are overridden by ToolDef fields. - - Regression variant: when a tool dict contains both new fields (name, - input_schema) AND empty legacy fields (identifier="", parameters=[], - provider_id="", type=""), the endpoint must populate from the new sources. - """ - mock_config = Configuration( - name="test", - service=ServiceConfiguration( - tls_config=TLSConfiguration( - tls_certificate_path=Path("tests/configuration/server.crt"), - tls_key_path=Path("tests/configuration/server.key"), - tls_key_password=Path("tests/configuration/password"), - ), - cors=CORSConfiguration( - allow_origins=["*"], - allow_credentials=False, - allow_methods=["*"], - allow_headers=["*"], - ), - host="localhost", - port=8080, - base_url=".", - auth_enabled=False, - workers=1, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - url=AnyHttpUrl("http://localhost:8321"), - api_key=SecretStr("xyzzy"), - use_as_library_client=False, - library_client_config_path=".", - timeout=10, - ), - user_data_collection=UserDataCollection( - transcripts_enabled=False, - feedback_enabled=False, - transcripts_storage=".", - feedback_storage=".", - ), - mcp_servers=[], - customization=None, - authorization=None, - deployment_environment=".", + """Skip MCP servers that return no tools without failing the request.""" + mock_configuration.mcp_servers = [ + ModelContextProtocolServer( + name="oauth-tools", + provider_id="model-context-protocol", + url="http://localhost:3003", + ) + ] + _make_app_config(mocker, mock_configuration) + mocker.patch( + "app.endpoints.tools.check_configuration_loaded", + return_value=None, ) - app_config = AppConfig() - app_config._configuration = mock_config - mocker.patch("app.endpoints.tools.configuration", app_config) - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") - mock_client = mocker.AsyncMock() - mock_client_holder.return_value.get_client.return_value = mock_client - - mock_toolgroup = mocker.Mock() - mock_toolgroup.identifier = "builtin::rag" - mock_toolgroup.provider_id = "rag-runtime" - mock_toolgroup.type = "tool_group" - mock_client.toolgroups.list.return_value = [mock_toolgroup] - - # Tool with both new fields AND empty legacy fields - rag_tool = _make_tool_def_mock( - mocker, - { - "name": "knowledge_search", - "identifier": "", - "description": "Search for information in a database.", - "input_schema": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "The query to search for.", - } - }, - "required": ["query"], - }, - "parameters": [], - "provider_id": "", - "type": "", - "toolgroup_id": "builtin::rag", - "metadata": None, - "output_schema": None, - }, + mocker.patch( + "app.endpoints.tools.build_mcp_headers", + return_value={"oauth-tools": {}}, + ) + mocker.patch("app.endpoints.tools.check_mcp_auth", return_value=None) + mocker.patch( + "app.endpoints.tools.get_agent_capability_tools", + return_value=[], + ) + _mock_file_search_tools(mocker) + mocker.patch( + "app.endpoints.tools.list_mcp_tools", + return_value=[], ) - mock_client.tools.list.return_value = [rag_tool] - - mock_request = mocker.Mock() - mock_auth = MOCK_AUTH - - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] - assert isinstance(response, ToolsResponse) - assert len(response.tools) == 1 + request = mocker.Mock() + request.headers = {} - tool = response.tools[0] - # Empty legacy fields must be overridden by new sources - assert tool["identifier"] == "knowledge_search" - assert tool["provider_id"] == "rag-runtime" - assert tool["type"] == "tool_group" - assert tool["server_source"] == "builtin" - assert tool["toolgroup_id"] == "builtin::rag" + response = await tools.tools_endpoint_handler( + request, + auth=MOCK_AUTH, + mcp_headers={}, + ) - # Parameters populated from input_schema, not empty legacy list - assert len(tool["parameters"]) == 1 - assert tool["parameters"][0]["name"] == "query" - assert tool["parameters"][0]["parameter_type"] == "string" - assert tool["parameters"][0]["required"] is True + identifiers = {tool.identifier for tool in response.tools} + assert identifiers == {"insert_into_memory", "file_search"} @pytest.mark.asyncio @@ -1151,29 +329,30 @@ async def test_tools_endpoint_includes_agent_capability_tools( app_config = AppConfig() app_config._configuration = config_with_skills mocker.patch("app.endpoints.tools.configuration", app_config) - mocker.patch("app.endpoints.tools.authorize", lambda _: lambda func: func) - mock_client_holder = mocker.patch("app.endpoints.tools.AsyncLlamaStackClientHolder") + mock_client_holder = mocker.patch("app.endpoints.tools.AsyncOgxClientHolder") mock_client = mocker.AsyncMock() mock_client_holder.return_value.get_client.return_value = mock_client mock_client.toolgroups.list.return_value = [] mock_request = mocker.Mock() - mock_auth = MOCK_AUTH + mock_request.headers = {} - response = await tools.tools_endpoint_handler.__wrapped__( - mock_request, mock_auth, {} - ) # pyright: ignore[reportFunctionMemberAccess] + response = await tools.tools_endpoint_handler( + request=mock_request, + auth=MOCK_AUTH, + mcp_headers={}, + ) - tool_ids = [tool["identifier"] for tool in response.tools] + tool_ids = [tool.identifier for tool in response.tools] assert "list_skills" in tool_ids assert "load_skill" in tool_ids assert "read_skill_resource" in tool_ids assert "run_skill_script" in tool_ids list_skills = next( - tool for tool in response.tools if tool["identifier"] == "list_skills" + tool for tool in response.tools if tool.identifier == "list_skills" ) - assert list_skills["provider_id"] == "agent-skills" - assert list_skills["toolgroup_id"] == "builtin::agent-skills" - assert list_skills["server_source"] == "builtin" + assert list_skills.provider_id == "agent-skills" + assert list_skills.toolgroup_id == "builtin::agent-skills" + assert list_skills.server_source == "builtin" diff --git a/tests/unit/app/endpoints/test_vector_stores.py b/tests/unit/app/endpoints/test_vector_stores.py index 2c83584a9..ffd9e0e6a 100644 --- a/tests/unit/app/endpoints/test_vector_stores.py +++ b/tests/unit/app/endpoints/test_vector_stores.py @@ -6,7 +6,7 @@ import pytest from fastapi import HTTPException, Request, status -from llama_stack_client import APIConnectionError, BadRequestError +from ogx_client import APIConnectionError, BadRequestError from pytest_mock import MockerFixture from app.endpoints.vector_stores import ( @@ -189,7 +189,7 @@ async def test_create_vector_store_success(mocker: MockerFixture) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.create.return_value = VectorStore("vs_123", "test_store") mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -217,7 +217,7 @@ async def test_create_vector_store_connection_error(mocker: MockerFixture) -> No mock_client = mocker.AsyncMock() mock_client.vector_stores.create.side_effect = APIConnectionError(request=None) # type: ignore mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -230,7 +230,7 @@ async def test_create_vector_store_connection_error(mocker: MockerFixture) -> No await create_vector_store(request=request, auth=auth, body=body) assert e.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE - assert e.value.detail["response"] == "Unable to connect to Llama Stack" # type: ignore + assert e.value.detail["response"] == "Unable to connect to OGX" # type: ignore @pytest.mark.asyncio @@ -247,7 +247,7 @@ async def test_list_vector_stores_success(mocker: MockerFixture) -> None: [VectorStore("vs_1", "store1"), VectorStore("vs_2", "store2")] ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -276,7 +276,7 @@ async def test_get_vector_store_success(mocker: MockerFixture) -> None: "vs_123", "test_store" ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -309,7 +309,7 @@ async def test_get_vector_store_not_found(mocker: MockerFixture) -> None: message="Not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -336,7 +336,7 @@ async def test_update_vector_store_success(mocker: MockerFixture) -> None: "vs_123", "updated_store" ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -365,7 +365,7 @@ async def test_delete_vector_store_success(mocker: MockerFixture) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.delete.return_value = None mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -392,7 +392,7 @@ async def test_create_file_success(mocker: MockerFixture) -> None: mock_client = mocker.AsyncMock() mock_client.files.create.return_value = File("file_123", "test.txt", 1024) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -427,7 +427,7 @@ async def test_add_file_to_vector_store_success(mocker: MockerFixture) -> None: "file_123", "vs_123" ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -463,7 +463,7 @@ async def test_add_file_to_vector_store_retry_on_database_lock( VectorStoreFile("file_123", "vs_123"), ] mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -504,7 +504,7 @@ async def test_add_file_to_vector_store_max_retries_exceeded( # All attempts fail with database lock error mock_client.vector_stores.files.create.side_effect = Exception("database is locked") mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -548,7 +548,7 @@ async def test_add_file_to_vector_store_non_lock_error_no_retry( # Raise a non-lock error mock_client.vector_stores.files.create.side_effect = Exception("Some other error") mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -589,7 +589,7 @@ async def test_list_vector_store_files_success(mocker: MockerFixture) -> None: ] ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -620,7 +620,7 @@ async def test_get_vector_store_file_success(mocker: MockerFixture) -> None: "file_123", "vs_123" ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -648,7 +648,7 @@ async def test_delete_vector_store_file_success(mocker: MockerFixture) -> None: mock_client = mocker.AsyncMock() mock_client.vector_stores.files.delete.return_value = None mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -678,7 +678,7 @@ async def test_list_vector_stores_connection_error(mocker: MockerFixture) -> Non mock_client = mocker.AsyncMock() mock_client.vector_stores.list.side_effect = APIConnectionError(request=None) # type: ignore mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -703,7 +703,7 @@ async def test_update_vector_store_connection_error(mocker: MockerFixture) -> No mock_client = mocker.AsyncMock() mock_client.vector_stores.update.side_effect = APIConnectionError(request=None) # type: ignore mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -735,7 +735,7 @@ async def test_update_vector_store_not_found(mocker: MockerFixture) -> None: message="Not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -763,7 +763,7 @@ async def test_delete_vector_store_connection_error(mocker: MockerFixture) -> No mock_client = mocker.AsyncMock() mock_client.vector_stores.delete.side_effect = APIConnectionError(request=None) # type: ignore mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -792,7 +792,7 @@ async def test_delete_vector_store_not_found(mocker: MockerFixture) -> None: message="Not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -819,7 +819,7 @@ async def test_create_file_connection_error(mocker: MockerFixture) -> None: mock_client = mocker.AsyncMock() mock_client.files.create.side_effect = APIConnectionError(request=None) # type: ignore mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -853,7 +853,7 @@ async def test_create_file_bad_request(mocker: MockerFixture) -> None: message="File too large", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -950,7 +950,7 @@ async def test_add_file_to_vector_store_connection_error( request=None # type: ignore ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -982,7 +982,7 @@ async def test_add_file_to_vector_store_not_found(mocker: MockerFixture) -> None message="File not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1014,7 +1014,7 @@ async def test_list_vector_store_files_connection_error( request=None # type: ignore ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1045,7 +1045,7 @@ async def test_list_vector_store_files_not_found(mocker: MockerFixture) -> None: message="Vector store not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1075,7 +1075,7 @@ async def test_get_vector_store_file_connection_error(mocker: MockerFixture) -> request=None # type: ignore ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1106,7 +1106,7 @@ async def test_get_vector_store_file_not_found(mocker: MockerFixture) -> None: message="File not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1137,7 +1137,7 @@ async def test_delete_vector_store_file_connection_error( request=None # type: ignore ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1168,7 +1168,7 @@ async def test_delete_vector_store_file_not_found(mocker: MockerFixture) -> None message="File not found", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1197,7 +1197,7 @@ async def test_get_vector_store_connection_error(mocker: MockerFixture) -> None: request=None # type: ignore ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1224,7 +1224,7 @@ async def test_create_file_adds_txt_extension_when_missing( mock_client = mocker.AsyncMock() mock_client.files.create.return_value = File("file_123", "uploaded_file.txt", 12) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) @@ -1263,7 +1263,7 @@ async def test_create_file_non_size_bad_request_returns_400( message="Invalid file format", response=mock_response, body=None ) mock_lsc = mocker.patch( - "app.endpoints.vector_stores.AsyncLlamaStackClientHolder.get_client" + "app.endpoints.vector_stores.AsyncOgxClientHolder.get_client" ) mock_lsc.return_value = mock_client mocker.patch("app.endpoints.vector_stores.configuration", cfg) diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index b654cd1e7..b16ea951e 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -3,18 +3,19 @@ from __future__ import annotations import logging -from collections.abc import Generator +from collections.abc import Callable, Generator from pathlib import Path +from typing import Optional import httpx import pytest -from llama_stack_client import AsyncLlamaStackClient +from ogx_client import AsyncOgxClient from pytest_mock import AsyncMockType, MockerFixture from configuration import AppConfig from constants import DEFAULT_LOGGER_NAME from models.common.responses.responses_api_params import ResponsesApiParams -from models.config import SkillsConfiguration +from models.config import ShieldConfiguration, SkillsConfiguration type AgentFixtures = Generator[ tuple[ @@ -57,7 +58,7 @@ def prepare_agent_mocks_fixture( ) -> AgentFixtures: """Prepare for mock for the LLM agent. - Provides common mocks for AsyncLlamaStackClient and AsyncAgent + Provides common mocks for AsyncOgxClient and AsyncAgent with proper agent_id setup to avoid initialization errors. Yields: @@ -109,9 +110,9 @@ def minimal_config_fixture() -> AppConfig: @pytest.fixture(name="mock_client") def mock_client_fixture( # pylint: disable=protected-access mocker: MockerFixture, -) -> AsyncLlamaStackClient: +) -> AsyncOgxClient: """Remote Llama Stack client mock for build_agent tests.""" - client = mocker.Mock(spec=AsyncLlamaStackClient) + client = mocker.Mock(spec=AsyncOgxClient) client.base_url = "http://localhost:8321" client.api_key = "test-key" client._client = mocker.Mock(spec=httpx.AsyncClient) @@ -143,3 +144,26 @@ def mock_skills_configuration_fixture(tmp_path: Path) -> SkillsConfiguration: encoding="utf-8", ) return SkillsConfiguration(paths=[skills_root]) + + +@pytest.fixture(name="make_agent_config") +def make_agent_config_fixture( + mocker: MockerFixture, +) -> Callable[..., AppConfig]: + """Return a factory building a duck-typed AppConfig stand-in for build_agent. + + ``build_agent`` only reads ``config.skills`` and ``config.shields`` off the + config object it receives, so tests can pass a lightweight mock instead of + a fully-initialized ``AppConfig``. + """ + + def _make( + skills: Optional[SkillsConfiguration] = None, + shields: Optional[list[ShieldConfiguration]] = None, + ) -> AppConfig: + config = mocker.Mock() + config.skills = skills + config.shields = shields or [] + return config + + return _make diff --git a/tests/unit/metrics/test_utis.py b/tests/unit/metrics/test_utis.py index 7d4490c2d..3310ea66a 100644 --- a/tests/unit/metrics/test_utis.py +++ b/tests/unit/metrics/test_utis.py @@ -1,23 +1,32 @@ """Unit tests for functions defined in metrics/utils.py""" import pytest +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pytest_mock import MockerFixture from metrics.utils import setup_model_metrics +def _make_model(model_id: str, provider_id: str, model_type: str) -> Model: + """Build an OGX Model for metrics tests.""" + return Model.model_construct( + id=model_id, + created=0, + owned_by="test", + object="model", + custom_metadata={"provider_id": provider_id, "model_type": model_type}, + ) + + @pytest.mark.asyncio async def test_setup_model_metrics(mocker: MockerFixture) -> None: """Test the setup_model_metrics function.""" - # Mock the LlamaStackAsLibraryClient - mock_client = mocker.patch( - "client.AsyncLlamaStackClientHolder.get_client" - ).return_value + # Mock the OGXAsLibraryClient + mock_client = mocker.patch("client.AsyncOgxClientHolder.get_client").return_value # Make sure the client is an AsyncMock for async methods mock_client = mocker.AsyncMock() - mocker.patch( - "client.AsyncLlamaStackClientHolder.get_client", return_value=mock_client - ) + mocker.patch("client.AsyncOgxClientHolder.get_client", return_value=mock_client) mocker.patch( "metrics.utils.configuration.inference.default_provider", "default_provider", @@ -28,34 +37,20 @@ async def test_setup_model_metrics(mocker: MockerFixture) -> None: ) mock_metric = mocker.patch("metrics.provider_model_configuration") - # Mock a model that is the default - model_default = mocker.Mock( - id="default_model", - custom_metadata={"provider_id": "default_provider", "model_type": "llm"}, - ) - # Mock a model that is not the default - model_0 = mocker.Mock( - id="test_model-0", - custom_metadata={"provider_id": "test_provider-0", "model_type": "llm"}, - ) - # Mock a second model which is not default - model_1 = mocker.Mock( - id="test_model-1", - custom_metadata={"provider_id": "test_provider-1", "model_type": "llm"}, - ) - # Mock a model that is not an LLM type, should be ignored - not_llm_model = mocker.Mock( - id="not-llm-model", - custom_metadata={"provider_id": "not-llm-provider", "model_type": "not-llm"}, - ) + model_default = _make_model("default_model", "default_provider", "llm") + model_0 = _make_model("test_model-0", "test_provider-0", "llm") + model_1 = _make_model("test_model-1", "test_provider-1", "llm") + not_llm_model = _make_model("not-llm-model", "not-llm-provider", "not-llm") # Mock the list of models returned by the client - mock_client.models.list.return_value = [ - model_0, - model_default, - not_llm_model, - model_1, - ] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[ + model_0, + model_default, + not_llm_model, + model_1, + ] + ) await setup_model_metrics() diff --git a/tests/unit/models/config/test_dump_configuration.py b/tests/unit/models/config/test_dump_configuration.py index 56970e82b..36dc1011e 100644 --- a/tests/unit/models/config/test_dump_configuration.py +++ b/tests/unit/models/config/test_dump_configuration.py @@ -262,6 +262,7 @@ def test_dump_configuration_minimal_cfg(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -490,6 +491,7 @@ def test_dump_configuration_valid_values(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -869,6 +871,7 @@ def test_dump_configuration_with_quota_limiters(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1132,6 +1135,7 @@ def test_dump_configuration_with_quota_limiters_different_values( }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1428,6 +1432,7 @@ def test_dump_configuration_byok(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1651,6 +1656,7 @@ def test_dump_configuration_pg_namespace(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2034,6 +2040,7 @@ def test_dump_configuration_allow_degraded_mode(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2263,6 +2270,7 @@ def test_dump_configuration_max_retries_settings(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2492,6 +2500,7 @@ def test_dump_configuration_retry_count_settings(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2728,4 +2737,5 @@ def test_dump_configuration_specific_compaction_values(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } diff --git a/tests/unit/models/config/test_llama_stack_configuration.py b/tests/unit/models/config/test_llama_stack_configuration.py index 239815d8f..c22139c35 100644 --- a/tests/unit/models/config/test_llama_stack_configuration.py +++ b/tests/unit/models/config/test_llama_stack_configuration.py @@ -113,7 +113,7 @@ def test_llama_stack_wrong_configuration_constructor_no_url() -> None: """ with pytest.raises( ValueError, - match="Llama stack URL is not specified and library client mode is not specified", + match="Llama Stack URL is not specified and library client mode is not specified", ): LlamaStackConfiguration() # pyright: ignore[reportCallIssue] @@ -122,7 +122,7 @@ def test_llama_stack_wrong_configuration_constructor_library_mode_off() -> None: """Test the LlamaStackConfiguration constructor.""" with pytest.raises( ValueError, - match="Llama stack URL is not specified and library client mode is not enabled", + match="Llama Stack URL is not specified and library client mode is not enabled", ): LlamaStackConfiguration( use_as_library_client=False diff --git a/tests/unit/models/config/test_shields_configuration.py b/tests/unit/models/config/test_shields_configuration.py new file mode 100644 index 000000000..e5a80d743 --- /dev/null +++ b/tests/unit/models/config/test_shields_configuration.py @@ -0,0 +1,196 @@ +"""Unit tests for ShieldConfiguration model and the Configuration.shields list.""" + +# pylint: disable=no-member + +import pytest +from pydantic import ValidationError + +from models.config import ( + CompactionConfiguration, + Configuration, + LlamaStackConfiguration, + QuestionValidityConfig, + QuestionValidityShieldConfiguration, + RedactionConfig, + RedactionRule, + RedactionShieldConfiguration, + ServiceConfiguration, + UserDataCollection, +) + + +class TestShieldConfiguration: + """Tests for the ShieldConfiguration discriminated union variants.""" + + def test_question_validity_shield(self) -> None: + """A question_validity shield parses config into QuestionValidityConfig.""" + shield = QuestionValidityShieldConfiguration.model_validate( + { + "name": "topic-guard", + "provider_id": "question_validity", + "config": {"model_id": "test-model"}, + } + ) + assert shield.name == "topic-guard" + assert shield.provider_id == "question_validity" + assert isinstance(shield.config, QuestionValidityConfig) + assert shield.config.model_id == "test-model" + + def test_redaction_shield(self) -> None: + """A redaction shield parses config into RedactionConfig.""" + shield = RedactionShieldConfiguration.model_validate( + { + "name": "pii-guard", + "provider_id": "redaction", + "config": {"rules": [{"pattern": r"\d+", "replacement": "[NUM]"}]}, + } + ) + assert shield.name == "pii-guard" + assert shield.provider_id == "redaction" + assert isinstance(shield.config, RedactionConfig) + assert len(shield.config.compiled_patterns) == 1 + + def test_accepts_already_constructed_config_instance(self) -> None: + """config may be passed as an already-constructed model instance.""" + shield = QuestionValidityShieldConfiguration( + name="topic-guard", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ) + assert isinstance(shield.config, QuestionValidityConfig) + + def test_rejects_config_mismatched_with_provider_id(self) -> None: + """A redaction provider_id with question_validity-shaped config is rejected.""" + with pytest.raises(ValidationError, match="model_id"): + RedactionShieldConfiguration.model_validate( + { + "name": "bad", + "provider_id": "redaction", + "config": {"model_id": "oops"}, + } + ) + + def test_rejects_unknown_provider_id(self) -> None: + """An unrecognized shield provider_id is rejected by the root Configuration.""" + with pytest.raises(ValidationError): + Configuration.model_validate( + { + **_minimal_configuration_kwargs(), + "shields": [ + { + "name": "bad", + "provider_id": "unknown_type", + "config": {"model_id": "test-model"}, + } + ], + } + ) + + def test_rejects_unknown_fields(self) -> None: + """Unknown fields are forbidden on shield configuration variants.""" + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + QuestionValidityShieldConfiguration.model_validate( + { + "name": "topic-guard", + "provider_id": "question_validity", + "config": {"model_id": "test-model"}, + "unknown_field": "value", + } + ) + + +def _minimal_configuration_kwargs() -> dict: + return { + "name": "test", + "service": ServiceConfiguration(), + "llama_stack": LlamaStackConfiguration( + use_as_library_client=True, + library_client_config_path="tests/configuration/run.yaml", + ), + "user_data_collection": UserDataCollection( + feedback_enabled=False, feedback_storage=None + ), + "compaction": CompactionConfiguration(), + } + + +def test_root_configuration_has_shields_field() -> None: + """The root Configuration declares a shields list field, empty by default.""" + field_info = Configuration.model_fields.get("shields") + assert field_info is not None + + factory = field_info.default_factory + assert factory is not None + assert factory() == [] # type: ignore[call-arg] + + +def test_root_configuration_default_shields_is_empty() -> None: + """Configuration constructed without shields defaults to an empty list.""" + cfg = Configuration(**_minimal_configuration_kwargs()) + assert cfg.shields == [] + + +def test_root_configuration_accepts_multiple_shields_of_same_type() -> None: + """Multiple shields of the same provider may be configured with distinct ids.""" + cfg = Configuration( + **_minimal_configuration_kwargs(), + shields=[ + QuestionValidityShieldConfiguration( + name="topic-guard-a", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="model-a"), + ), + QuestionValidityShieldConfiguration( + name="topic-guard-b", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="model-b"), + ), + ], + ) + assert len(cfg.shields) == 2 + assert cfg.shields[0].name == "topic-guard-a" + assert cfg.shields[1].name == "topic-guard-b" + + +def test_root_configuration_accepts_mixed_shield_types() -> None: + """Shields of different providers may be mixed in the same list.""" + cfg = Configuration( + **_minimal_configuration_kwargs(), + shields=[ + QuestionValidityShieldConfiguration( + name="topic-guard", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig( + rules=[RedactionRule(pattern=r"\d+", replacement="[NUM]")] + ), + ), + ], + ) + assert len(cfg.shields) == 2 + assert cfg.shields[0].provider_id == "question_validity" + assert cfg.shields[1].provider_id == "redaction" + + +def test_root_configuration_rejects_duplicate_shield_names() -> None: + """Shield names must be unique across the shields list.""" + with pytest.raises(ValidationError, match="Shield names must be unique"): + Configuration( + **_minimal_configuration_kwargs(), + shields=[ + QuestionValidityShieldConfiguration( + name="dup", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="model-a"), + ), + RedactionShieldConfiguration( + name="dup", + provider_id="redaction", + config=RedactionConfig(), + ), + ], + ) diff --git a/tests/unit/models/responses/test_error_responses.py b/tests/unit/models/responses/test_error_responses.py index 1675cc844..0da44aa1d 100644 --- a/tests/unit/models/responses/test_error_responses.py +++ b/tests/unit/models/responses/test_error_responses.py @@ -712,12 +712,12 @@ class TestServiceUnavailableResponse: def test_constructor(self) -> None: """Test ServiceUnavailableResponse with valid parameters.""" response = ServiceUnavailableResponse( - backend_name="Llama Stack", cause="Connection timeout" + backend_name="OGX", cause="Connection timeout" ) assert isinstance(response, AbstractErrorResponse) assert response.status_code == status.HTTP_503_SERVICE_UNAVAILABLE assert isinstance(response.detail, DetailModel) - assert response.detail.response == "Unable to connect to Llama Stack" + assert response.detail.response == "Unable to connect to OGX" assert response.detail.cause == "Connection timeout" def test_different_backend_names(self) -> None: @@ -746,24 +746,21 @@ def test_openapi_response(self) -> None: assert expected_count == 2 # Verify example structure - assert "llama stack" in examples + assert "ogx" in examples assert "kubernetes api" in examples - llama_example = examples["llama stack"] - assert "value" in llama_example - assert "detail" in llama_example["value"] - assert ( - llama_example["value"]["detail"]["response"] - == "Unable to connect to Llama Stack" - ) + ogx_example = examples["ogx"] + assert "value" in ogx_example + assert "detail" in ogx_example["value"] + assert ogx_example["value"]["detail"]["response"] == "Unable to connect to OGX" def test_openapi_response_with_explicit_examples(self) -> None: """Test ServiceUnavailableResponse.openapi_response() with explicit examples.""" - result = ServiceUnavailableResponse.openapi_response(examples=["llama stack"]) + result = ServiceUnavailableResponse.openapi_response(examples=["ogx"]) examples = result["content"]["application/json"]["examples"] # Verify only 1 example is returned when explicitly specified assert len(examples) == 1 - assert "llama stack" in examples + assert "ogx" in examples class TestPromptTooLongResponse: diff --git a/tests/unit/models/responses/test_successful_responses.py b/tests/unit/models/responses/test_successful_responses.py index d1e4864c3..fe27f789e 100644 --- a/tests/unit/models/responses/test_successful_responses.py +++ b/tests/unit/models/responses/test_successful_responses.py @@ -70,7 +70,7 @@ def test_constructor(self) -> None: ] response = ModelsResponse(models=models) assert isinstance(response, AbstractSuccessfulResponse) - assert response.models == models + assert [model.model_dump() for model in response.models] == models assert len(response.models) == 1 def test_empty_models_list(self) -> None: @@ -82,8 +82,18 @@ def test_empty_models_list(self) -> None: def test_multiple_models(self) -> None: """Test ModelsResponse with multiple models.""" models = [ - {"identifier": "model1", "provider_id": "provider1"}, - {"identifier": "model2", "provider_id": "provider2"}, + { + "identifier": "model1", + "provider_id": "provider1", + "api_model_type": "llm", + "model_type": "llm", + }, + { + "identifier": "model2", + "provider_id": "provider2", + "api_model_type": "embedding", + "model_type": "embedding", + }, ] response = ModelsResponse(models=models) assert len(response.models) == 2 @@ -135,13 +145,15 @@ def test_constructor(self) -> None: "identifier": "filesystem_read", "description": "Read contents of a file", "parameters": [], - "provider_id": "mcp", + "provider_id": "model-context-protocol", + "toolgroup_id": "filesystem-tools", + "server_source": "http://localhost:3000", "type": "tool", } ] response = ToolsResponse(tools=tools) assert isinstance(response, AbstractSuccessfulResponse) - assert response.tools == tools + assert [tool.model_dump() for tool in response.tools] == tools def test_empty_tools_list(self) -> None: """Test ToolsResponse with empty tools list.""" @@ -174,10 +186,23 @@ class TestShieldsResponse: def test_constructor(self) -> None: """Test ShieldsResponse with valid shields list.""" - shields = [{"name": "shield1", "status": "active"}] + shields = [ + { + "name": "question-validity", + "provider_id": "question_validity", + "type": "shield", + "config": { + "model_id": "openai/gpt-4o-mini", + "model_prompt": "Is this question valid?", + "invalid_question_response": ( + "I can only answer questions about the product." + ), + }, + } + ] response = ShieldsResponse(shields=shields) assert isinstance(response, AbstractSuccessfulResponse) - assert response.shields == shields + assert [shield.model_dump() for shield in response.shields] == shields def test_missing_required_parameter(self) -> None: """Test ShieldsResponse raises ValidationError when shields is missing.""" diff --git a/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py b/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py index 2e6037fc9..9658e917a 100644 --- a/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py +++ b/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py @@ -14,6 +14,7 @@ DEFAULT_INVALID_QUESTION_RESPONSE, DEFAULT_MODEL_PROMPT, ) +from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed from models.config import ( QuestionValidityConfig, ) @@ -21,6 +22,7 @@ SUBJECT_ALLOWED, SUBJECT_REJECTED, QuestionValidity, + _extract_conversation_id, _extract_message_str_from_user_content, ) @@ -64,6 +66,47 @@ def test_sequence_with_non_text_content(self) -> None: assert result == "keep" +class TestExtractConversationId: + """Tests for _extract_conversation_id helper.""" + + def test_extracts_conversation_id(self, mocker: MockerFixture) -> None: + """Test extraction when extra_body.conversation is set.""" + model = mocker.Mock() + model.settings = {"extra_body": {"conversation": "conv_123"}} + + assert _extract_conversation_id(model) == "conv_123" + + def test_returns_none_when_settings_missing(self, mocker: MockerFixture) -> None: + """Test that None settings yields None.""" + model = mocker.Mock() + model.settings = None + + assert _extract_conversation_id(model) is None + + def test_returns_none_when_extra_body_missing(self, mocker: MockerFixture) -> None: + """Test that missing extra_body yields None.""" + model = mocker.Mock() + model.settings = {} + + assert _extract_conversation_id(model) is None + + def test_returns_none_when_extra_body_not_dict(self, mocker: MockerFixture) -> None: + """Test that a non-dict extra_body yields None instead of raising.""" + model = mocker.Mock() + model.settings = {"extra_body": "not-a-dict"} + + assert _extract_conversation_id(model) is None + + def test_returns_none_when_conversation_not_string( + self, mocker: MockerFixture + ) -> None: + """Test that a non-string conversation value yields None.""" + model = mocker.Mock() + model.settings = {"extra_body": {"conversation": 123}} + + assert _extract_conversation_id(model) is None + + class TestQuestionValidityConfigInit: """Tests for QuestionValidityConfig initialization.""" @@ -110,13 +153,13 @@ class TestQuestionValidityInit: """Tests for QuestionValidity dataclass initialization.""" def test_post_init_wires_client_and_model(self, mocker: MockerFixture) -> None: - """Test that __post_init__ obtains the client and passes it to from_llama_stack_client.""" + """Test that __post_init__ obtains the client and passes it to from_ogx_client.""" mock_client = mocker.Mock() - mock_holder = mocker.patch(f"{_MODULE}.AsyncLlamaStackClientHolder") + mock_holder = mocker.patch(f"{_MODULE}.AsyncOgxClientHolder") mock_holder.return_value.get_client.return_value = mock_client mock_from_client = mocker.patch( - f"{_MODULE}.LlamaStackResponsesModel.from_llama_stack_client", + f"{_MODULE}.OgxResponsesModel.from_ogx_client", ) config = QuestionValidityConfig(model_id="test-model") @@ -130,11 +173,11 @@ def test_post_init_wires_client_and_model(self, mocker: MockerFixture) -> None: ) def test_model_is_assigned_from_factory(self, mocker: MockerFixture) -> None: - """Test that the model returned by from_llama_stack_client is stored.""" + """Test that the model returned by from_ogx_client is stored.""" mock_model = mocker.Mock() - mocker.patch(f"{_MODULE}.AsyncLlamaStackClientHolder") + mocker.patch(f"{_MODULE}.AsyncOgxClientHolder") mocker.patch( - f"{_MODULE}.LlamaStackResponsesModel.from_llama_stack_client", + f"{_MODULE}.OgxResponsesModel.from_ogx_client", return_value=mock_model, ) config = QuestionValidityConfig(model_id="test") @@ -150,8 +193,8 @@ class TestBuildPrompt: @pytest.fixture(autouse=True) def _mock_create_model(self, mocker: MockerFixture) -> None: """Mock model creation for all tests.""" - mocker.patch(f"{_MODULE}.AsyncLlamaStackClientHolder") - mocker.patch(f"{_MODULE}.LlamaStackResponsesModel.from_llama_stack_client") + mocker.patch(f"{_MODULE}.AsyncOgxClientHolder") + mocker.patch(f"{_MODULE}.OgxResponsesModel.from_ogx_client") @pytest.fixture(name="question_validity") def question_validity_fixture(self) -> QuestionValidity: @@ -212,15 +255,24 @@ class TestWrapRun: @pytest.fixture(autouse=True) def _mock_create_model(self, mocker: MockerFixture) -> None: """Mock model creation for all tests.""" - mocker.patch(f"{_MODULE}.AsyncLlamaStackClientHolder") - mocker.patch(f"{_MODULE}.LlamaStackResponsesModel.from_llama_stack_client") + mocker.patch(f"{_MODULE}.AsyncOgxClientHolder") + mocker.patch(f"{_MODULE}.OgxResponsesModel.from_ogx_client") + + @pytest.fixture(name="mock_append_turn", autouse=True) + def mock_append_turn_fixture(self, mocker: MockerFixture) -> MockType: + """Mock the conversation-persistence call used on rejection.""" + return mocker.patch( + f"{_MODULE}.append_turn_to_conversation", new_callable=mocker.AsyncMock + ) @pytest.fixture(name="mock_ctx") def mock_ctx_fixture(self, mocker: MockerFixture) -> RunContext: - """Create a mock RunContext.""" + """Create a mock RunContext bound to a model with a conversation ID.""" ctx = mocker.Mock(spec=RunContext) ctx.prompt = "How do I create a pod?" ctx.usage = RunUsage() + ctx.model = mocker.Mock() + ctx.model.settings = {"extra_body": {"conversation": "conv_test"}} return ctx @pytest.fixture(name="mock_handler") @@ -279,6 +331,89 @@ async def test_rejected_question_returns_rejection( assert isinstance(result, AgentRunResult) assert result.output == DEFAULT_INVALID_QUESTION_RESPONSE + @pytest.mark.asyncio + async def test_rejected_question_persists_turn_to_conversation( + self, + mocker: MockerFixture, + mock_ctx: RunContext, + mock_handler: MockType, + mock_append_turn: MockType, + ) -> None: + """Test that a rejection appends the user question and refusal to the conversation.""" + mock_client = mocker.Mock() + mocker.patch( + f"{_MODULE}.AsyncOgxClientHolder" + ).return_value.get_client.return_value = mock_client + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_REJECTED)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch( + "pydantic_ai_lightspeed.capabilities.question_validity._capability.model_request", + return_value=mock_response, + ) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + await qv.wrap_run(mock_ctx, handler=mock_handler) + + mock_append_turn.assert_awaited_once_with( + mock_client, + "conv_test", + "How do I create a pod?", + DEFAULT_INVALID_QUESTION_RESPONSE, + ) + + @pytest.mark.asyncio + async def test_rejection_skips_persistence_when_conversation_id_missing( + self, + mocker: MockerFixture, + mock_ctx: RunContext, + mock_handler: MockType, + mock_append_turn: MockType, + ) -> None: + """Test that persistence is skipped (not crashed) without a conversation ID.""" + mock_ctx.model = mocker.Mock(settings={}) + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_REJECTED)], + usage=RequestUsage(), + ) + mocker.patch( + "pydantic_ai_lightspeed.capabilities.question_validity._capability.model_request", + return_value=mock_response, + ) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.wrap_run(mock_ctx, handler=mock_handler) + + mock_append_turn.assert_not_awaited() + assert result.output == DEFAULT_INVALID_QUESTION_RESPONSE + + @pytest.mark.asyncio + async def test_allowed_question_does_not_persist_turn( + self, + mocker: MockerFixture, + mock_ctx: RunContext, + mock_handler: MockType, + mock_append_turn: MockType, + ) -> None: + """Test that an allowed question does not touch the conversation.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_ALLOWED)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch( + "pydantic_ai_lightspeed.capabilities.question_validity._capability.model_request", + return_value=mock_response, + ) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + await qv.wrap_run(mock_ctx, handler=mock_handler) + + mock_append_turn.assert_not_awaited() + @pytest.mark.asyncio async def test_unexpected_response_treated_as_rejected( self, @@ -468,6 +603,8 @@ async def test_wrap_run_with_none_prompt( ctx = mocker.Mock(spec=RunContext) ctx.prompt = None ctx.usage = RunUsage() + ctx.model = mocker.Mock() + ctx.model.settings = {"extra_body": {"conversation": "conv_test"}} mock_response = ModelResponse( parts=[TextPart(content=SUBJECT_REJECTED)], @@ -532,3 +669,108 @@ async def test_wrap_run_with_sequence_prompt( prompt_str = str(messages[0]) assert "How to" in prompt_str assert "scale a deployment?" in prompt_str + + +class TestQuestionValidityRun: + """Tests for QuestionValidity.run method.""" + + @pytest.fixture(autouse=True) + def _mock_create_model(self, mocker: MockerFixture) -> None: + """Mock model creation for all tests.""" + mocker.patch(f"{_MODULE}.AsyncOgxClientHolder") + mocker.patch(f"{_MODULE}.OgxResponsesModel.from_ogx_client") + + @pytest.mark.asyncio + async def test_allowed_returns_passed(self, mocker: MockerFixture) -> None: + """Test that an allowed response returns ShieldModerationPassed.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_ALLOWED)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("How do I create a pod?") + + assert isinstance(result, ShieldModerationPassed) + assert result.decision == "passed" + + @pytest.mark.asyncio + async def test_rejected_returns_blocked(self, mocker: MockerFixture) -> None: + """Test that a rejected response returns ShieldModerationBlocked.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_REJECTED)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("What is the meaning of life?") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == DEFAULT_INVALID_QUESTION_RESPONSE + assert result.moderation_id.startswith("modr-") + assert result.refusal_response.role == "assistant" + assert result.refusal_response.content == DEFAULT_INVALID_QUESTION_RESPONSE + + @pytest.mark.asyncio + async def test_unexpected_response_returns_blocked( + self, mocker: MockerFixture + ) -> None: + """Test that an unexpected model response is treated as blocked.""" + mock_response = ModelResponse( + parts=[TextPart(content="I don't understand")], + usage=RequestUsage(input_tokens=10, output_tokens=5), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("some input") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == DEFAULT_INVALID_QUESTION_RESPONSE + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "response_text", + [" ALLOWED", "ALLOWED ", " ALLOWED ", "ALLOWED\n"], + ids=["leading-space", "trailing-space", "both-spaces", "trailing-newline"], + ) + async def test_allowed_with_whitespace_returns_passed( + self, mocker: MockerFixture, response_text: str + ) -> None: + """Test that ALLOWED with surrounding whitespace still returns passed.""" + mock_response = ModelResponse( + parts=[TextPart(content=response_text)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("How do I scale pods?") + + assert isinstance(result, ShieldModerationPassed) + + @pytest.mark.asyncio + async def test_custom_invalid_response_message(self, mocker: MockerFixture) -> None: + """Test that a custom rejection message is used in the blocked result.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_REJECTED)], + usage=RequestUsage(), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig( + model_id="test", invalid_question_response="Custom rejection." + ) + qv = QuestionValidity(config=config) + result = await qv.run("off-topic question") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == "Custom rejection." + assert result.refusal_response.role == "assistant" + assert result.refusal_response.content == "Custom rejection." diff --git a/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py b/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py index c112caba5..cc7f6fe68 100644 --- a/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py +++ b/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py @@ -13,6 +13,7 @@ from pydantic_ai.models import ModelRequestContext from pytest_mock import MockerFixture +from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed from models.config import ( RedactionConfig, RedactionRule, @@ -314,3 +315,46 @@ async def test_after_model_request_no_match( ) assert result is resp assert resp.parts[0].content == "clean response" + + +class TestPiiRedactionCapabilityRun: + """Tests for PiiRedactionCapability.run method.""" + + @pytest.fixture(name="capability") + def capability_fixture(self) -> PiiRedactionCapability: + """Create a PiiRedactionCapability with an email redaction rule. + + Returns: + A configured PiiRedactionCapability instance. + """ + config = RedactionConfig( + rules=[ + RedactionRule( + pattern=r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}", + replacement="[REDACTED_EMAIL]", + ) + ], + case_sensitive=True, + ) + return PiiRedactionCapability(config=config) + + @pytest.mark.asyncio() + async def test_clean_text_returns_passed( + self, capability: PiiRedactionCapability + ) -> None: + """Test that clean text returns ShieldModerationPassed.""" + result = await capability.run("no sensitive content here") + + assert isinstance(result, ShieldModerationPassed) + assert result.decision == "passed" + + @pytest.mark.asyncio() + async def test_pii_text_returns_blocked( + self, capability: PiiRedactionCapability + ) -> None: + """Test that text with PII returns ShieldModerationBlocked.""" + result = await capability.run("contact user@example.com for details") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == "Sensitive content detected." + assert result.moderation_id.startswith("modr-") diff --git a/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py b/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py index d0e881f4f..962f7625a 100644 --- a/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py +++ b/tests/unit/pydantic_ai_lightspeed/llamastack/test_model.py @@ -20,7 +20,7 @@ from models.common.responses.responses_api_params import ResponsesApiParams from pydantic_ai_lightspeed.llamastack._model import ( _LLS_RESPONSES_EXTRA_FIELDS, - LlamaStackResponsesModel, + OgxResponsesModel, _FilteredResponseStream, _model_settings_from_responses_params, ) @@ -106,7 +106,9 @@ def test_tools_maps_to_openai_native_tools(self) -> None: assert "openai_native_tools" in settings assert settings["openai_native_tools"] is params.tools - assert "tools" not in settings.get("extra_body", {}) + extra_body = settings.get("extra_body", {}) + assert isinstance(extra_body, dict) + assert "tools" not in extra_body def test_none_fields_excluded(self) -> None: """Test that None optional fields do not appear in the result.""" @@ -119,27 +121,26 @@ def test_none_fields_excluded(self) -> None: assert "openai_previous_response_id" not in settings -class TestFromLlamaStackClient: - """Tests for LlamaStackResponsesModel.from_llama_stack_client factory.""" +class TestFromOgxClient: + """Tests for OgxResponsesModel.from_ogx_client factory.""" def test_with_responses_params(self, mocker: MockerFixture) -> None: """Test that responses_params is converted and forwarded.""" mock_provider = mocker.Mock() mocker.patch( - "pydantic_ai_lightspeed.llamastack._model.LlamaStackProvider" - ".from_llama_stack_client", + "pydantic_ai_lightspeed.llamastack._model.OgxProvider.from_ogx_client", return_value=mock_provider, ) mock_init = mocker.patch.object( - LlamaStackResponsesModel, "__init__", return_value=None + OgxResponsesModel, "__init__", return_value=None ) params = _make_params(temperature=0.5) client = mocker.Mock() - result = LlamaStackResponsesModel.from_llama_stack_client( + result = OgxResponsesModel.from_ogx_client( "test-model", client, responses_params=params ) - assert isinstance(result, LlamaStackResponsesModel) + assert isinstance(result, OgxResponsesModel) args, kwargs = mock_init.call_args assert kwargs["settings"]["temperature"] == 0.5 assert kwargs["provider"] is mock_provider @@ -150,21 +151,20 @@ def test_with_model_settings(self, mocker: MockerFixture) -> None: """Test that model_settings is forwarded directly.""" mock_provider = mocker.Mock() mocker.patch( - "pydantic_ai_lightspeed.llamastack._model.LlamaStackProvider" - ".from_llama_stack_client", + "pydantic_ai_lightspeed.llamastack._model.OgxProvider.from_ogx_client", return_value=mock_provider, ) mock_init = mocker.patch.object( - LlamaStackResponsesModel, "__init__", return_value=None + OgxResponsesModel, "__init__", return_value=None ) settings: ModelSettings = {"temperature": 0.9} client = mocker.Mock() - result = LlamaStackResponsesModel.from_llama_stack_client( + result = OgxResponsesModel.from_ogx_client( "test-model", client, model_settings=settings ) - assert isinstance(result, LlamaStackResponsesModel) + assert isinstance(result, OgxResponsesModel) args, kwargs = mock_init.call_args assert kwargs["settings"] is settings assert kwargs["provider"] is mock_provider @@ -175,18 +175,17 @@ def test_with_neither(self, mocker: MockerFixture) -> None: """Test that settings is None when neither param is provided.""" mock_provider = mocker.Mock() mocker.patch( - "pydantic_ai_lightspeed.llamastack._model.LlamaStackProvider" - ".from_llama_stack_client", + "pydantic_ai_lightspeed.llamastack._model.OgxProvider.from_ogx_client", return_value=mock_provider, ) mock_init = mocker.patch.object( - LlamaStackResponsesModel, "__init__", return_value=None + OgxResponsesModel, "__init__", return_value=None ) client = mocker.Mock() - result = LlamaStackResponsesModel.from_llama_stack_client("test-model", client) + result = OgxResponsesModel.from_ogx_client("test-model", client) - assert isinstance(result, LlamaStackResponsesModel) + assert isinstance(result, OgxResponsesModel) args, kwargs = mock_init.call_args assert kwargs["settings"] is None assert kwargs["provider"] is mock_provider @@ -196,8 +195,7 @@ def test_with_neither(self, mocker: MockerFixture) -> None: def test_both_raises_value_error(self, mocker: MockerFixture) -> None: """Test that providing both raises ValueError.""" mocker.patch( - "pydantic_ai_lightspeed.llamastack._model.LlamaStackProvider" - ".from_llama_stack_client", + "pydantic_ai_lightspeed.llamastack._model.OgxProvider.from_ogx_client", return_value=mocker.Mock(), ) @@ -206,7 +204,7 @@ def test_both_raises_value_error(self, mocker: MockerFixture) -> None: client = mocker.Mock() with pytest.raises(ValueError, match="ResponsesApiParams or ModelSetting"): - LlamaStackResponsesModel.from_llama_stack_client( + OgxResponsesModel.from_ogx_client( "test-model", client, responses_params=params, @@ -215,16 +213,16 @@ def test_both_raises_value_error(self, mocker: MockerFixture) -> None: class TestPrepareConversationContinuation: - """Tests for LlamaStackResponsesModel._prepare_conversation_continuation.""" + """Tests for OgxResponsesModel._prepare_conversation_continuation.""" @pytest.fixture(name="model") - def model_fixture(self, mocker: MockerFixture) -> LlamaStackResponsesModel: - """Create a LlamaStackResponsesModel with mocked __init__.""" - mocker.patch.object(LlamaStackResponsesModel, "__init__", return_value=None) - return LlamaStackResponsesModel("test-model") + def model_fixture(self, mocker: MockerFixture) -> OgxResponsesModel: + """Create a OgxResponsesModel with mocked __init__.""" + mocker.patch.object(OgxResponsesModel, "__init__", return_value=None) + return OgxResponsesModel("test-model") def test_none_settings_returns_unchanged( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that None model_settings returns messages and settings unchanged.""" messages = [mocker.Mock()] @@ -234,9 +232,7 @@ def test_none_settings_returns_unchanged( assert result_msgs is messages assert result_settings is None - def test_empty_settings_returns_unchanged( - self, model: LlamaStackResponsesModel - ) -> None: + def test_empty_settings_returns_unchanged(self, model: OgxResponsesModel) -> None: """Test that empty dict model_settings returns unchanged.""" messages: list = [] settings: ModelSettings = {} @@ -247,7 +243,7 @@ def test_empty_settings_returns_unchanged( assert result_settings is settings def test_no_extra_body_returns_unchanged( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that settings without extra_body returns unchanged.""" messages = [mocker.Mock()] @@ -259,7 +255,7 @@ def test_no_extra_body_returns_unchanged( assert result_settings is settings def test_extra_body_without_conversation_returns_unchanged( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that extra_body without 'conversation' key returns unchanged.""" messages = [mocker.Mock()] @@ -271,7 +267,7 @@ def test_extra_body_without_conversation_returns_unchanged( assert result_settings is settings def test_no_model_response_returns_unchanged( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that messages without ModelResponse returns unchanged.""" messages = [mocker.Mock(), mocker.Mock()] @@ -283,7 +279,7 @@ def test_no_model_response_returns_unchanged( assert result_settings is settings def test_model_response_without_provider_id_returns_unchanged( - self, model: LlamaStackResponsesModel + self, model: OgxResponsesModel ) -> None: """Test that ModelResponse without provider_response_id is ignored.""" response_msg = ModelResponse(parts=[], provider_response_id=None) @@ -296,7 +292,7 @@ def test_model_response_without_provider_id_returns_unchanged( assert result_settings is settings def test_trims_messages_and_strips_previous_response_id( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that messages are trimmed and previous_response_id is removed.""" msg_before = mocker.Mock() @@ -317,7 +313,7 @@ def test_trims_messages_and_strips_previous_response_id( assert result_settings["extra_body"] == {"conversation": "conv-1"} def test_trims_without_previous_response_id_in_settings( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test trimming works when settings lacks previous_response_id.""" response_msg = ModelResponse(parts=[], provider_response_id="resp-1") @@ -332,7 +328,7 @@ def test_trims_without_previous_response_id_in_settings( assert "openai_previous_response_id" not in result_settings def test_uses_last_model_response_when_multiple( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that the last ModelResponse with provider_response_id is used.""" msg1 = mocker.Mock() @@ -353,7 +349,7 @@ def test_uses_last_model_response_when_multiple( assert "openai_previous_response_id" not in result_settings def test_only_skip_model_response_with_provider_response_id( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that the last ModelResponse with provider_response_id is used.""" msg1 = mocker.Mock() @@ -374,7 +370,7 @@ def test_only_skip_model_response_with_provider_response_id( assert "openai_previous_response_id" not in result_settings def test_does_not_mutate_original_settings( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that the original settings dict is not modified.""" response_msg = ModelResponse(parts=[], provider_response_id="resp-1") @@ -388,15 +384,15 @@ def test_does_not_mutate_original_settings( class TestRequest: - """Tests for LlamaStackResponsesModel.request.""" + """Tests for OgxResponsesModel.request.""" @pytest.mark.asyncio async def test_calls_prepare_and_delegates_to_super( self, mocker: MockerFixture ) -> None: """Test that request calls _prepare_conversation_continuation and delegates.""" - mocker.patch.object(LlamaStackResponsesModel, "__init__", return_value=None) - model = LlamaStackResponsesModel("test-model") + mocker.patch.object(OgxResponsesModel, "__init__", return_value=None) + model = OgxResponsesModel("test-model") original_msgs = [mocker.Mock()] original_settings: OpenAIResponsesModelSettings = { @@ -462,13 +458,13 @@ def _make_response_created_event() -> responses.ResponseCreatedEvent: class TestRequestStream: - """Tests for LlamaStackResponsesModel.request_stream.""" + """Tests for OgxResponsesModel.request_stream.""" @pytest.fixture(name="model") - def model_fixture(self, mocker: MockerFixture) -> LlamaStackResponsesModel: - """Create a LlamaStackResponsesModel with stream-related attributes set.""" - mocker.patch.object(LlamaStackResponsesModel, "__init__", return_value=None) - model = LlamaStackResponsesModel("test-model") + def model_fixture(self, mocker: MockerFixture) -> OgxResponsesModel: + """Create a OgxResponsesModel with stream-related attributes set.""" + mocker.patch.object(OgxResponsesModel, "__init__", return_value=None) + model = OgxResponsesModel("test-model") mocker.patch.object( type(model), "model_name", @@ -485,7 +481,7 @@ def model_fixture(self, mocker: MockerFixture) -> LlamaStackResponsesModel: @pytest.mark.asyncio async def test_calls_prepare_continuation( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that request_stream calls _prepare_conversation_continuation.""" original_msgs = [mocker.Mock()] @@ -510,7 +506,7 @@ async def test_calls_prepare_continuation( @pytest.mark.asyncio async def test_empty_stream_raises( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that an empty stream raises UnexpectedModelBehavior.""" mocker.patch.object( @@ -527,7 +523,7 @@ async def test_empty_stream_raises( @pytest.mark.asyncio async def test_wrong_first_event_raises( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that a non-ResponseCreatedEvent first event raises.""" mocker.patch.object( @@ -561,7 +557,7 @@ async def test_wrong_first_event_raises( @pytest.mark.asyncio async def test_happy_path_yields_streamed_response( - self, model: LlamaStackResponsesModel, mocker: MockerFixture + self, model: OgxResponsesModel, mocker: MockerFixture ) -> None: """Test that a valid stream yields an OpenAIResponsesStreamedResponse.""" mocker.patch.object( diff --git a/tests/unit/pydantic_ai_lightspeed/llamastack/test_provider.py b/tests/unit/pydantic_ai_lightspeed/llamastack/test_provider.py index 3e02ad3a6..67b385a95 100644 --- a/tests/unit/pydantic_ai_lightspeed/llamastack/test_provider.py +++ b/tests/unit/pydantic_ai_lightspeed/llamastack/test_provider.py @@ -4,88 +4,88 @@ import httpx import pytest -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack_client import AsyncLlamaStackClient +from ogx.core.library_client import AsyncOGXAsLibraryClient +from ogx_client import AsyncOgxClient from openai import AsyncOpenAI from pytest_mock import MockerFixture from pydantic_ai_lightspeed.llamastack._provider import ( DEFAULT_BASE_URL, - LlamaStackProvider, + OgxProvider, ) -from pydantic_ai_lightspeed.llamastack._transport import LlamaStackServerTransport +from pydantic_ai_lightspeed.llamastack._transport import OgxServerTransport -class TestLlamaStackProviderProperties: - """Tests for LlamaStackProvider basic properties.""" +class TestOgxProviderProperties: + """Tests for OgxProvider basic properties.""" def test_name(self) -> None: """Test that the provider name is 'llama-stack'.""" - provider = LlamaStackProvider() + provider = OgxProvider() assert provider.name == "llama-stack" def test_base_url_default(self) -> None: """Test that the default base URL matches the expected default.""" - provider = LlamaStackProvider() + provider = OgxProvider() assert DEFAULT_BASE_URL in provider.base_url def test_client_returns_async_openai(self) -> None: """Test that the client property returns an AsyncOpenAI instance.""" - provider = LlamaStackProvider() + provider = OgxProvider() assert isinstance(provider.client, AsyncOpenAI) def test_repr(self) -> None: """Test the string representation of the provider.""" - provider = LlamaStackProvider() + provider = OgxProvider() result = repr(provider) - assert "LlamaStackProvider" in result + assert "OgxProvider" in result assert "llama-stack" in result def test_model_profile_known_model(self) -> None: """Test model_profile returns a profile for a known OpenAI model.""" - profile = LlamaStackProvider.model_profile("gpt-4o") + profile = OgxProvider.model_profile("gpt-4o") assert profile is not None def test_model_profile_unknown_model(self) -> None: """Test model_profile returns a default profile for an unrecognized model.""" - profile = LlamaStackProvider.model_profile("totally-unknown-model-xyz") + profile = OgxProvider.model_profile("totally-unknown-model-xyz") assert profile is not None -class TestLlamaStackProviderServerMode: - """Tests for LlamaStackProvider server mode initialization.""" +class TestOgxProviderServerMode: + """Tests for OgxProvider server mode initialization.""" def test_explicit_base_url(self) -> None: """Test that an explicit base_url is used.""" - provider = LlamaStackProvider(base_url="http://my-server:9999/v1") + provider = OgxProvider(base_url="http://my-server:9999/v1") assert "my-server:9999" in provider.base_url def test_explicit_api_key(self) -> None: """Test that an explicit api_key is used.""" - provider = LlamaStackProvider(api_key="my-secret-key") + provider = OgxProvider(api_key="my-secret-key") assert provider.client.api_key == "my-secret-key" def test_default_api_key_is_not_needed(self) -> None: """Test that the default API key is 'not-needed'.""" - provider = LlamaStackProvider() + provider = OgxProvider() assert provider.client.api_key == "not-needed" def test_custom_http_client(self, mocker: MockerFixture) -> None: """Test that a provided http_client is wired into the provider.""" custom_client = mocker.Mock(spec=httpx.AsyncClient) - provider = LlamaStackProvider(http_client=custom_client) + provider = OgxProvider(http_client=custom_client) assert provider._client._client is custom_client -class TestLlamaStackProviderLibraryMode: - """Tests for LlamaStackProvider library mode initialization.""" +class TestOgxProviderLibraryMode: + """Tests for OgxProvider library mode initialization.""" def test_library_client_creates_transport(self, mocker: MockerFixture) -> None: """Test that providing a library_client sets up the transport-based client.""" mock_lib_client = mocker.Mock() mock_lib_client.provider_data = None - provider = LlamaStackProvider(library_client=mock_lib_client) + provider = OgxProvider(library_client=mock_lib_client) assert provider._library_client is mock_lib_client assert "llama-stack-library" in provider.base_url @@ -95,12 +95,12 @@ def test_library_client_api_key_is_not_needed(self, mocker: MockerFixture) -> No mock_lib_client = mocker.Mock() mock_lib_client.provider_data = None - provider = LlamaStackProvider(library_client=mock_lib_client) + provider = OgxProvider(library_client=mock_lib_client) assert provider.client.api_key == "not-needed" -class TestLlamaStackProviderMutualExclusion: +class TestOgxProviderMutualExclusion: """Tests for mutual exclusion between library_client and server mode options.""" def test_library_client_and_base_url_raises(self, mocker: MockerFixture) -> None: @@ -112,7 +112,7 @@ def test_library_client_and_base_url_raises(self, mocker: MockerFixture) -> None ValueError, match="Cannot provide both `library_client` and `base_url`", ): - LlamaStackProvider( + OgxProvider( library_client=mock_lib_client, base_url="http://localhost:8321/v1", ) @@ -126,7 +126,7 @@ def test_library_client_and_api_key_raises(self, mocker: MockerFixture) -> None: ValueError, match="Cannot provide both `library_client` and `api_key`", ): - LlamaStackProvider( + OgxProvider( library_client=mock_lib_client, api_key="my-key", ) @@ -140,23 +140,23 @@ def test_library_client_and_http_client_raises(self, mocker: MockerFixture) -> N ValueError, match="Cannot provide both `library_client` and `http_client`", ): - LlamaStackProvider( + OgxProvider( library_client=mock_lib_client, http_client=mocker.Mock(spec=httpx.AsyncClient), ) -class TestFromLlamaStackClient: - """Tests for LlamaStackProvider.from_llama_stack_client.""" +class TestFromOgxClient: + """Tests for OgxProvider.from_ogx_client.""" def test_library_client_dispatches_to_library_mode( self, mocker: MockerFixture ) -> None: - """Test that an AsyncLlamaStackAsLibraryClient creates a library-mode provider.""" - mock_lib_client = mocker.Mock(spec=AsyncLlamaStackAsLibraryClient) + """Test that an AsyncOGXAsLibraryClient creates a library-mode provider.""" + mock_lib_client = mocker.Mock(spec=AsyncOGXAsLibraryClient) mock_lib_client.provider_data = None - provider = LlamaStackProvider.from_llama_stack_client(mock_lib_client) + provider = OgxProvider.from_ogx_client(mock_lib_client) assert provider._library_client is mock_lib_client assert "llama-stack-library" in provider.base_url @@ -165,26 +165,26 @@ def test_server_client_extracts_base_url_with_v1( self, mocker: MockerFixture ) -> None: """Test that a server client whose base_url already ends with /v1 is used as-is.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321/v1" mock_client.api_key = "test-key" mock_client._client = mocker.Mock(spec=httpx.AsyncClient) mock_client.default_headers = {} - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert "my-server:8321/v1" in provider.base_url assert provider.base_url.count("/v1") == 1 def test_server_client_appends_v1_when_missing(self, mocker: MockerFixture) -> None: """Test that /v1 is appended when the server client's base_url lacks it.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321" mock_client.api_key = "test-key" mock_client._client = mocker.Mock(spec=httpx.AsyncClient) mock_client.default_headers = {} - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert provider.base_url.rstrip("/").endswith("/v1") @@ -192,26 +192,26 @@ def test_server_client_strips_trailing_slash_before_appending_v1( self, mocker: MockerFixture ) -> None: """Test that a trailing slash is stripped before appending /v1.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321/" mock_client.api_key = "test-key" mock_client._client = mocker.Mock(spec=httpx.AsyncClient) mock_client.default_headers = {} - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert "//v1" not in provider.base_url assert provider.base_url.rstrip("/").endswith("/v1") def test_server_client_uses_provided_api_key(self, mocker: MockerFixture) -> None: """Test that the server client's api_key is forwarded to the provider.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321/v1" mock_client.api_key = "my-secret" mock_client._client = mocker.Mock(spec=httpx.AsyncClient) mock_client.default_headers = {} - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert provider.client.api_key == "my-secret" @@ -219,26 +219,26 @@ def test_server_client_defaults_api_key_when_none( self, mocker: MockerFixture ) -> None: """Test that a None api_key falls back to 'not-needed'.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321/v1" mock_client.api_key = None mock_client._client = mocker.Mock(spec=httpx.AsyncClient) mock_client.default_headers = {} - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert provider.client.api_key == "not-needed" def test_server_client_passes_http_client(self, mocker: MockerFixture) -> None: """Test that the server client's internal httpx client is reused when no provider data.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321/v1" mock_client.api_key = "test-key" inner_http = mocker.Mock(spec=httpx.AsyncClient) mock_client._client = inner_http mock_client.default_headers = {} - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert provider._client._client is inner_http @@ -246,29 +246,29 @@ def test_server_client_wraps_transport_with_provider_data( self, mocker: MockerFixture ) -> None: """Test provider data from default_headers is forwarded in server mode.""" - mock_client = mocker.Mock(spec=AsyncLlamaStackClient) + mock_client = mocker.Mock(spec=AsyncOgxClient) mock_client.base_url = "http://my-server:8321/v1" mock_client.api_key = "test-key" inner_http = httpx.AsyncClient() mock_client._client = inner_http mock_client.default_headers = { - "X-LlamaStack-Provider-Data": '{"azure_api_key": "token"}' + "X-OGX-Provider-Data": '{"azure_api_key": "token"}' } - provider = LlamaStackProvider.from_llama_stack_client(mock_client) + provider = OgxProvider.from_ogx_client(mock_client) assert isinstance( provider._client._client._transport, # pylint: disable=protected-access - LlamaStackServerTransport, + OgxServerTransport, ) class TestSetHttpClient: # pylint: disable=too-few-public-methods - """Tests for LlamaStackProvider._set_http_client.""" + """Tests for OgxProvider._set_http_client.""" def test_replaces_internal_http_client(self, mocker: MockerFixture) -> None: """Test that _set_http_client replaces the underlying httpx client.""" - provider = LlamaStackProvider() + provider = OgxProvider() new_client = mocker.Mock(spec=httpx.AsyncClient) provider._set_http_client(new_client) diff --git a/tests/unit/pydantic_ai_lightspeed/llamastack/test_transport.py b/tests/unit/pydantic_ai_lightspeed/llamastack/test_transport.py index 9a01c4454..c66a0704e 100644 --- a/tests/unit/pydantic_ai_lightspeed/llamastack/test_transport.py +++ b/tests/unit/pydantic_ai_lightspeed/llamastack/test_transport.py @@ -11,8 +11,8 @@ from pytest_mock import MockerFixture from pydantic_ai_lightspeed.llamastack._transport import ( - LlamaStackLibraryTransport, - LlamaStackServerTransport, + OgxLibraryTransport, + OgxServerTransport, _AsyncByteStream, decode_request_headers, inject_provider_data_into_headers, @@ -22,7 +22,7 @@ @pytest.fixture(name="mock_library_client") def mock_library_client_fixture(mocker: MockerFixture) -> Any: - """Create a mock AsyncLlamaStackAsLibraryClient. + """Create a mock AsyncOGXAsLibraryClient. Returns: A mocked library client with route_impls set to an empty dict. @@ -34,13 +34,13 @@ def mock_library_client_fixture(mocker: MockerFixture) -> Any: @pytest.fixture(name="transport") -def transport_fixture(mock_library_client: Any) -> LlamaStackLibraryTransport: - """Create a LlamaStackLibraryTransport with a mocked library client. +def transport_fixture(mock_library_client: Any) -> OgxLibraryTransport: + """Create a OgxLibraryTransport with a mocked library client. Returns: - An initialized LlamaStackLibraryTransport. + An initialized OgxLibraryTransport. """ - return LlamaStackLibraryTransport(mock_library_client) + return OgxLibraryTransport(mock_library_client) class TestAsyncByteStream: @@ -81,16 +81,16 @@ def test_injects_when_absent(self) -> None: """Test provider data is serialized into the canonical header.""" headers = inject_provider_data_into_headers({}, {"api_key": "test-key"}) - assert headers["X-LlamaStack-Provider-Data"] == '{"api_key": "test-key"}' + assert headers["X-OGX-Provider-Data"] == '{"api_key": "test-key"}' def test_does_not_override_existing_header(self) -> None: """Test an existing provider data header is left unchanged.""" headers = inject_provider_data_into_headers( - {"X-LlamaStack-Provider-Data": '{"existing": true}'}, + {"X-OGX-Provider-Data": '{"existing": true}'}, {"api_key": "ignored"}, ) - assert headers["X-LlamaStack-Provider-Data"] == '{"existing": true}' + assert headers["X-OGX-Provider-Data"] == '{"existing": true}' def test_request_with_provider_data_headers_returns_copy(self) -> None: """Test request helper returns a new request when headers are injected.""" @@ -106,13 +106,11 @@ def test_request_with_provider_data_headers_returns_copy(self) -> None: ) assert updated is not request - assert ( - updated.headers["X-LlamaStack-Provider-Data"] == '{"api_key": "test-key"}' - ) + assert updated.headers["X-OGX-Provider-Data"] == '{"api_key": "test-key"}' -class TestLlamaStackServerTransport: - """Tests for LlamaStackServerTransport.""" +class TestOgxServerTransport: + """Tests for OgxServerTransport.""" @pytest.mark.asyncio async def test_injects_provider_data_before_delegating( @@ -121,7 +119,7 @@ async def test_injects_provider_data_before_delegating( """Test provider data is added to outbound HTTP requests.""" wrapped = mocker.AsyncMock() wrapped.handle_async_request.return_value = httpx.Response(200) - transport = LlamaStackServerTransport( + transport = OgxServerTransport( wrapped, provider_data={"api_key": "test-key"}, ) @@ -135,7 +133,7 @@ async def test_injects_provider_data_before_delegating( delegated_request = wrapped.handle_async_request.await_args.args[0] assert ( - delegated_request.headers["X-LlamaStack-Provider-Data"] + delegated_request.headers["X-OGX-Provider-Data"] == '{"api_key": "test-key"}' ) @@ -146,7 +144,7 @@ async def test_preserves_existing_provider_data_header( """Test per-request provider data headers are not overwritten.""" wrapped = mocker.AsyncMock() wrapped.handle_async_request.return_value = httpx.Response(200) - transport = LlamaStackServerTransport( + transport = OgxServerTransport( wrapped, provider_data={"api_key": "ignored"}, ) @@ -155,15 +153,12 @@ async def test_preserves_existing_provider_data_header( "POST", "http://localhost/v1/responses", content=b"{}", - headers={"X-LlamaStack-Provider-Data": '{"existing": true}'}, + headers={"X-OGX-Provider-Data": '{"existing": true}'}, ) await transport.handle_async_request(request) delegated_request = wrapped.handle_async_request.await_args.args[0] - assert ( - delegated_request.headers["X-LlamaStack-Provider-Data"] - == '{"existing": true}' - ) + assert delegated_request.headers["X-OGX-Provider-Data"] == '{"existing": true}' class TestDecodeRequestHeaders: # pylint: disable=too-few-public-methods @@ -180,24 +175,24 @@ def test_decodes_raw_headers(self) -> None: assert decode_request_headers(request)["X-Test"] == "value" -class TestLlamaStackLibraryTransportInit: # pylint: disable=too-few-public-methods - """Tests for LlamaStackLibraryTransport initialization.""" +class TestOgxLibraryTransportInit: # pylint: disable=too-few-public-methods + """Tests for OgxLibraryTransport initialization.""" def test_stores_client(self, mock_library_client: Any) -> None: """Test that the transport stores the provided library client.""" - transport = LlamaStackLibraryTransport(mock_library_client) + transport = OgxLibraryTransport(mock_library_client) assert transport._client is mock_library_client class TestHandleAsyncRequest: - """Tests for LlamaStackLibraryTransport.handle_async_request.""" + """Tests for OgxLibraryTransport.handle_async_request.""" @pytest.mark.asyncio async def test_raises_when_route_impls_is_none(self, mocker: MockerFixture) -> None: """Test RuntimeError is raised when the library client is not initialized.""" client = mocker.Mock() client.route_impls = None - transport = LlamaStackLibraryTransport(client) + transport = OgxLibraryTransport(client) request = httpx.Request("POST", "http://localhost/v1/responses") @@ -209,7 +204,7 @@ async def test_raises_when_route_impls_is_none(self, mocker: MockerFixture) -> N @pytest.mark.asyncio async def test_non_streaming_request( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test a non-streaming request is dispatched correctly.""" body = {"model": "test-model", "messages": []} @@ -235,7 +230,7 @@ async def test_non_streaming_request( @pytest.mark.asyncio async def test_streaming_request( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test a streaming request returns an event-stream response.""" body = {"model": "test-model", "stream": True} @@ -263,7 +258,7 @@ async def mock_stream_result() -> AsyncGenerator[dict[str, int], None]: @pytest.mark.asyncio async def test_empty_body_request( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test that a request with no content body passes an empty dict.""" request = httpx.Request("GET", "http://localhost/v1/models") @@ -286,7 +281,7 @@ async def test_provider_data_header_injection(self, mocker: MockerFixture) -> No client = mocker.Mock() client.route_impls = {} client.provider_data = {"api_key": "test-key"} - transport = LlamaStackLibraryTransport(client) + transport = OgxLibraryTransport(client) body = {"model": "test-model"} request = httpx.Request( @@ -309,10 +304,8 @@ async def test_provider_data_header_injection(self, mocker: MockerFixture) -> No await transport.handle_async_request(request) call_args = mock_ctx.call_args[0][0] - assert "X-LlamaStack-Provider-Data" in call_args - assert json.loads(call_args["X-LlamaStack-Provider-Data"]) == { - "api_key": "test-key" - } + assert "X-OGX-Provider-Data" in call_args + assert json.loads(call_args["X-OGX-Provider-Data"]) == {"api_key": "test-key"} @pytest.mark.asyncio async def test_provider_data_header_not_injected_when_present( @@ -322,14 +315,14 @@ async def test_provider_data_header_not_injected_when_present( client = mocker.Mock() client.route_impls = {} client.provider_data = {"api_key": "should-not-override"} - transport = LlamaStackLibraryTransport(client) + transport = OgxLibraryTransport(client) body = {"model": "test-model"} request = httpx.Request( "POST", "http://localhost/v1/responses", content=json.dumps(body).encode("utf-8"), - headers={"X-LlamaStack-Provider-Data": '{"existing": true}'}, + headers={"X-OGX-Provider-Data": '{"existing": true}'}, ) mock_func = mocker.AsyncMock(return_value={"id": "resp-1"}) @@ -346,15 +339,15 @@ async def test_provider_data_header_not_injected_when_present( await transport.handle_async_request(request) call_args = mock_ctx.call_args[0][0] - assert json.loads(call_args["X-LlamaStack-Provider-Data"]) == {"existing": True} + assert json.loads(call_args["X-OGX-Provider-Data"]) == {"existing": True} class TestHandleNonStreaming: - """Tests for LlamaStackLibraryTransport._handle_non_streaming.""" + """Tests for OgxLibraryTransport._handle_non_streaming.""" @pytest.mark.asyncio async def test_merges_path_params( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test that path parameters are merged into the request body.""" body: dict[str, Any] = {"model": "test"} @@ -379,7 +372,7 @@ async def test_merges_path_params( @pytest.mark.asyncio async def test_delete_returns_no_content( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test that DELETE with None result returns 204 No Content.""" request = httpx.Request("DELETE", "http://localhost/v1/resource/123") @@ -400,7 +393,7 @@ async def test_delete_returns_no_content( @pytest.mark.asyncio async def test_delete_with_result_returns_ok( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test that DELETE with a non-None result returns 200 OK.""" request = httpx.Request("DELETE", "http://localhost/v1/resource/123") @@ -420,11 +413,11 @@ async def test_delete_with_result_returns_ok( class TestHandleStreaming: # pylint: disable=too-few-public-methods - """Tests for LlamaStackLibraryTransport._handle_streaming.""" + """Tests for OgxLibraryTransport._handle_streaming.""" @pytest.mark.asyncio async def test_produces_sse_format( - self, mocker: MockerFixture, transport: LlamaStackLibraryTransport + self, mocker: MockerFixture, transport: OgxLibraryTransport ) -> None: """Test that streaming responses produce SSE-formatted byte chunks.""" diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py index 760391c3c..a2b68f663 100644 --- a/tests/unit/test_client.py +++ b/tests/unit/test_client.py @@ -9,12 +9,14 @@ import pytest from fastapi import HTTPException -from llama_stack_client import APIConnectionError, APIStatusError +from ogx_client import APIConnectionError, APIStatusError +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pydantic import AnyHttpUrl, SecretStr from pytest_mock import MockerFixture from authorization.azure_token_manager import AzureEntraIDManager -from client import AsyncLlamaStackClientHolder +from client import AsyncOgxClientHolder from configuration import AzureEntraIdConfiguration from models.config import LlamaStackConfiguration from utils.types import Singleton @@ -28,12 +30,12 @@ def reset_singleton() -> None: def test_async_client_get_client_method() -> None: """Test how get_client method works for uninitialized client.""" - client = AsyncLlamaStackClientHolder() + client = AsyncOgxClientHolder() with pytest.raises( RuntimeError, match=( - "AsyncLlamaStackClient has not been initialised. " + "AsyncOgxClient has not been initialised. " "Ensure 'load\\(..\\)' has been called." ), ): @@ -50,7 +52,7 @@ async def test_get_async_llama_stack_library_client() -> None: library_client_config_path="./tests/configuration/minimal-stack.yaml", timeout=60, ) - client = AsyncLlamaStackClientHolder() + client = AsyncOgxClientHolder() await client.load(cfg) assert client is not None @@ -71,7 +73,7 @@ async def test_get_async_llama_stack_remote_client() -> None: library_client_config_path="./tests/configuration/minimal-stack.yaml", timeout=60, ) - client = AsyncLlamaStackClientHolder() + client = AsyncOgxClientHolder() await client.load(cfg) assert client is not None @@ -102,9 +104,9 @@ async def test_get_async_llama_stack_wrong_configuration( cfg.library_client_config_path = None with pytest.raises( ValueError, - match="Cannot synthesize Llama Stack config", + match="Cannot synthesize OGX config", ): - client = AsyncLlamaStackClientHolder() + client = AsyncOgxClientHolder() await client.load(cfg) @@ -131,7 +133,7 @@ async def test_update_azure_token_service_client() -> None: library_client_config_path=None, timeout=60, ) - holder = AsyncLlamaStackClientHolder() + holder = AsyncOgxClientHolder() await holder.load(cfg) original_client = holder.get_client() @@ -139,9 +141,7 @@ async def test_update_azure_token_service_client() -> None: assert updated_client is not original_client assert holder.get_client() is updated_client - provider_data_json = updated_client.default_headers.get( - "X-LlamaStack-Provider-Data" - ) + provider_data_json = updated_client.default_headers.get("X-OGX-Provider-Data") provider_data = json.loads(provider_data_json) assert provider_data["azure_api_key"] == "fresh-token" assert provider_data["azure_api_base"] == "https://api.example.com" @@ -170,16 +170,14 @@ async def test_load_service_client_defers_azure_provider_data() -> None: library_client_config_path=None, timeout=60, ) - holder = AsyncLlamaStackClientHolder() + holder = AsyncOgxClientHolder() await holder.load(cfg) default_headers = holder.get_client().default_headers or {} - assert "X-LlamaStack-Provider-Data" not in default_headers + assert "X-OGX-Provider-Data" not in default_headers updated_client = await holder.update_azure_token() - provider_data_json = updated_client.default_headers.get( - "X-LlamaStack-Provider-Data" - ) + provider_data_json = updated_client.default_headers.get("X-OGX-Provider-Data") assert provider_data_json is not None provider_data = json.loads(provider_data_json) assert provider_data["azure_api_key"] == "startup-token" @@ -198,7 +196,7 @@ async def test_reload_library_client() -> None: library_client_config_path="./tests/configuration/minimal-stack.yaml", timeout=60, ) - holder = AsyncLlamaStackClientHolder() + holder = AsyncOgxClientHolder() await holder.load(cfg) original_client = holder.get_client() @@ -219,30 +217,35 @@ class TestCheckModelAvailable: @pytest.fixture def holder_with_mock_client( self, mocker: MockerFixture - ) -> tuple[AsyncLlamaStackClientHolder, Any]: + ) -> tuple[AsyncOgxClientHolder, Any]: """Create a holder with a mocked async client.""" - holder = AsyncLlamaStackClientHolder() + holder = AsyncOgxClientHolder() mock_client = mocker.AsyncMock() holder._lsc = mock_client return holder, mock_client - def _make_model(self, mocker: MockerFixture, model_id: str) -> Any: - """Create a mock model with the given ID.""" - model = mocker.Mock() - model.id = model_id - return model + def _make_model(self, mocker: MockerFixture, model_id: str) -> Model: + """Create an OGX Model with the given ID.""" + _ = mocker + return Model.model_construct( + id=model_id, + created=0, + owned_by="test", + object="model", + custom_metadata={}, + ) @pytest.mark.asyncio async def test_model_available( self, mocker: MockerFixture, - holder_with_mock_client: tuple[AsyncLlamaStackClientHolder, Any], + holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns True when the model is found in the registry.""" holder, mock_client = holder_with_mock_client - mock_client.models.list.return_value = [ - self._make_model(mocker, self.EXPECTED_MODEL_ID) - ] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[self._make_model(mocker, self.EXPECTED_MODEL_ID)] + ) available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) @@ -253,11 +256,13 @@ async def test_model_available( async def test_model_not_found_service_client( self, mocker: MockerFixture, - holder_with_mock_client: tuple[AsyncLlamaStackClientHolder, Any], + holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns False and skips reload for non-library (service) clients.""" holder, mock_client = holder_with_mock_client - mock_client.models.list.return_value = [self._make_model(mocker, "other/model")] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[self._make_model(mocker, "other/model")] + ) available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) @@ -267,7 +272,7 @@ async def test_model_not_found_service_client( @pytest.mark.asyncio async def test_client_not_initialized(self) -> None: """Test returns False when the client has not been initialized.""" - holder = AsyncLlamaStackClientHolder() + holder = AsyncOgxClientHolder() available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) @@ -295,14 +300,14 @@ async def test_client_not_initialized(self) -> None: async def test_api_error( self, mocker: MockerFixture, - holder_with_mock_client: tuple[AsyncLlamaStackClientHolder, Any], + holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], exception_factory: Callable, ) -> None: """Test returns False when model list fails with API errors.""" _, mock_client = holder_with_mock_client mock_client.models.list.side_effect = exception_factory(mocker) - holder = AsyncLlamaStackClientHolder() + holder = AsyncOgxClientHolder() available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) assert available is False @@ -312,12 +317,12 @@ async def test_api_error( async def test_model_found_after_reload( self, mocker: MockerFixture, - holder_with_mock_client: tuple[AsyncLlamaStackClientHolder, Any], + holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns True when model is missing initially but found after reload.""" holder, mock_client = holder_with_mock_client mocker.patch.object( - AsyncLlamaStackClientHolder, + AsyncOgxClientHolder, "is_library_client", new_callable=mocker.PropertyMock, return_value=True, @@ -326,7 +331,10 @@ async def test_model_found_after_reload( wrong_model = self._make_model(mocker, "other/model") correct_model = self._make_model(mocker, self.EXPECTED_MODEL_ID) - mock_client.models.list.side_effect = [[wrong_model], [correct_model]] + mock_client.models.list.side_effect = [ + ListModelsResponse.model_construct(data=[wrong_model]), + ListModelsResponse.model_construct(data=[correct_model]), + ] available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) @@ -338,12 +346,12 @@ async def test_model_found_after_reload( async def test_reload_fails_returns_not_found( self, mocker: MockerFixture, - holder_with_mock_client: tuple[AsyncLlamaStackClientHolder, Any], + holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns False when model is missing and client reload fails.""" holder, mock_client = holder_with_mock_client mocker.patch.object( - AsyncLlamaStackClientHolder, + AsyncOgxClientHolder, "is_library_client", new_callable=mocker.PropertyMock, return_value=True, @@ -351,7 +359,9 @@ async def test_reload_fails_returns_not_found( holder.reload_library_client = mocker.AsyncMock( side_effect=RuntimeError("Cannot reload: config path not set") ) - mock_client.models.list.return_value = [self._make_model(mocker, "other/model")] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[self._make_model(mocker, "other/model")] + ) available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) @@ -362,12 +372,12 @@ async def test_reload_fails_returns_not_found( async def test_reload_http_exception_returns_not_found( self, mocker: MockerFixture, - holder_with_mock_client: tuple[AsyncLlamaStackClientHolder, Any], + holder_with_mock_client: tuple[AsyncOgxClientHolder, Any], ) -> None: """Test returns False when reload raises HTTPException.""" holder, mock_client = holder_with_mock_client mocker.patch.object( - AsyncLlamaStackClientHolder, + AsyncOgxClientHolder, "is_library_client", new_callable=mocker.PropertyMock, return_value=True, @@ -375,7 +385,9 @@ async def test_reload_http_exception_returns_not_found( holder.reload_library_client = mocker.AsyncMock( side_effect=HTTPException(status_code=503, detail="Llama Stack unavailable") ) - mock_client.models.list.return_value = [self._make_model(mocker, "other/model")] + mock_client.models.list.return_value = ListModelsResponse.model_construct( + data=[self._make_model(mocker, "other/model")] + ) available, reason = await holder.check_model_available(self.EXPECTED_MODEL_ID) diff --git a/tests/unit/test_configuration.py b/tests/unit/test_configuration.py index fe1d1fc20..23f77a4ca 100644 --- a/tests/unit/test_configuration.py +++ b/tests/unit/test_configuration.py @@ -130,6 +130,10 @@ def test_default_configuration() -> None: # try to read property _ = cfg.deployment_environment # pylint: disable=pointless-statement + with pytest.raises(LogicError, match="logic error: configuration is not loaded"): + # try to read property + _ = cfg.shields # pylint: disable=pointless-statement + def test_configuration_is_singleton() -> None: """Test that configuration is singleton.""" @@ -242,6 +246,70 @@ def test_init_from_dict() -> None: # check token usage history assert cfg.token_usage_history is None + # check shields - not configured in config_dict, defaults to empty list + assert cfg.shields == [] + + +def test_init_from_dict_with_shields() -> None: + """Test initialization with guardrail shields configuration.""" + config_dict: dict[str, Any] = { + "name": "foo", + "service": { + "host": "localhost", + "port": 8080, + "auth_enabled": False, + "workers": 1, + "color_log": True, + "access_log": True, + }, + "llama_stack": { + "api_key": "xyzzy", + "url": "http://x.y.com:1234", + "use_as_library_client": False, + }, + "user_data_collection": { + "feedback_enabled": False, + }, + "mcp_servers": [], + "customization": None, + "authentication": { + "module": "noop", + }, + "shields": [ + { + "name": "topic-guard-a", + "provider_id": "question_validity", + "config": {"model_id": "test-model"}, + }, + { + "name": "topic-guard-b", + "provider_id": "question_validity", + "config": {"model_id": "test-model-2"}, + }, + { + "name": "pii-guard", + "provider_id": "redaction", + "config": { + "rules": [ + {"pattern": r"\d+", "replacement": "[NUM]"}, + ], + }, + }, + ], + } + cfg = AppConfig() + cfg.init_from_dict(config_dict) + + assert len(cfg.shields) == 3 + assert cfg.shields[0].name == "topic-guard-a" + assert cfg.shields[0].provider_id == "question_validity" + assert cfg.shields[0].config.model_id == "test-model" # type: ignore[union-attr] + assert cfg.shields[1].name == "topic-guard-b" + assert cfg.shields[1].config.model_id == "test-model-2" # type: ignore[union-attr] + assert cfg.shields[2].name == "pii-guard" + assert cfg.shields[2].provider_id == "redaction" + assert len(cfg.shields[2].config.compiled_patterns) == 1 # type: ignore[union-attr] + def test_init_from_dict_with_mcp_servers() -> None: """Test initialization with MCP servers configuration.""" diff --git a/tests/unit/test_llama_stack_synthesize.py b/tests/unit/test_llama_stack_synthesize.py index dddb6f006..c03fb604c 100644 --- a/tests/unit/test_llama_stack_synthesize.py +++ b/tests/unit/test_llama_stack_synthesize.py @@ -116,10 +116,10 @@ def test_load_default_baseline_returns_usable_dict() -> None: def test_load_default_baseline_includes_mcp_tool_runtime() -> None: - """Default stack ships MCP beside rag-runtime (same rationale as RAG).""" + """Default stack ships MCP beside file-search.""" baseline = load_default_baseline() ids = _tool_runtime_ids(baseline) - assert "rag-runtime" in ids + assert "file-search" in ids assert "model-context-protocol" in ids diff --git a/tests/unit/utils/agents/test_query.py b/tests/unit/utils/agents/test_query.py index 4fe83c807..97aa40649 100644 --- a/tests/unit/utils/agents/test_query.py +++ b/tests/unit/utils/agents/test_query.py @@ -6,10 +6,7 @@ import pytest from fastapi import HTTPException -from llama_stack_api.openai_responses import ( - OpenAIResponseMessage as ResponseMessage, -) -from llama_stack_client import APIConnectionError, APIStatusError +from ogx_client import APIConnectionError, APIStatusError from pydantic_ai.messages import ( FinishReason, ImageUrl, @@ -112,10 +109,6 @@ def blocked_moderation_fixture() -> ShieldModerationBlocked: return ShieldModerationBlocked( message="Content blocked by shield.", moderation_id="modr-test-456", - refusal_response=ResponseMessage( - role="assistant", - content="Content blocked by shield.", - ), ) @@ -530,7 +523,7 @@ async def test_api_status_error_raises_http_exception( "detail": {"response": "Quota exceeded", "cause": "quota exceeded"}, } mocker.patch( - "utils.agents.query.handle_known_apistatus_errors", + "utils.agents.error_handler.handle_known_apistatus_errors", return_value=mock_error, ) diff --git a/tests/unit/utils/agents/test_streaming.py b/tests/unit/utils/agents/test_streaming.py index 8421b2584..f452f8106 100644 --- a/tests/unit/utils/agents/test_streaming.py +++ b/tests/unit/utils/agents/test_streaming.py @@ -10,10 +10,7 @@ import pytest from fastapi import HTTPException -from llama_stack_api.openai_responses import ( - OpenAIResponseMessage as ResponseMessage, -) -from llama_stack_client import APIStatusError +from ogx_client import APIStatusError from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError from pydantic_ai.messages import ( @@ -121,10 +118,6 @@ def blocked_moderation_fixture() -> ShieldModerationBlocked: return ShieldModerationBlocked( message="Content blocked by shield.", moderation_id="modr-test-456", - refusal_response=ResponseMessage( - role="assistant", - content="Content blocked by shield.", - ), ) @@ -616,7 +609,7 @@ async def test_agent_error_raises_http_exception( "detail": {"response": "Error", "cause": "agent failed"}, } mocker.patch( - "utils.agents.streaming.map_agent_inference_error", + "utils.agents.error_handler.map_agent_inference_error", return_value=mock_error, ) @@ -753,7 +746,7 @@ async def inner() -> AsyncIterator[str]: mock_error.detail.response = "Quota exceeded" mock_error.detail.cause = "quota exceeded" mocker.patch( - "utils.agents.streaming.map_agent_inference_error", + "utils.agents.error_handler.map_agent_inference_error", return_value=mock_error, ) mocker.patch( diff --git a/tests/unit/utils/test_builtin_tools.py b/tests/unit/utils/test_builtin_tools.py new file mode 100644 index 000000000..1d54b901c --- /dev/null +++ b/tests/unit/utils/test_builtin_tools.py @@ -0,0 +1,111 @@ +# pylint: disable=protected-access + +"""Unit tests for builtin file-search tool discovery.""" + +import pytest +from fastapi import HTTPException +from ogx_client import APIConnectionError +from ogx_client.types.shared.provider_info import ProviderInfo +from pytest_mock import MockerFixture + +from utils.builtin_tools import get_file_search_tools + + +def _provider( + *, + provider_id: str, + provider_type: str, + api: str = "tool_runtime", +) -> ProviderInfo: + """Build a ProviderInfo test fixture.""" + return ProviderInfo( + api=api, + provider_id=provider_id, + provider_type=provider_type, + config={}, + health={}, + ) + + +@pytest.mark.asyncio +async def test_get_file_search_tools_returns_empty_when_not_configured( + mocker: MockerFixture, +) -> None: + """Return no tools when Llama Stack has no file-search provider.""" + client = mocker.AsyncMock() + client.providers.list = mocker.AsyncMock( + return_value=[ + _provider( + provider_id="model-context-protocol", + provider_type="remote::model-context-protocol", + ) + ] + ) + + tools = await get_file_search_tools(client) + + assert tools == [] + client.get.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_file_search_tools_returns_static_catalog_when_provider_present( + mocker: MockerFixture, +) -> None: + """Return the known file-search catalog when the provider is configured.""" + client = mocker.AsyncMock() + client.providers.list = mocker.AsyncMock( + return_value=[ + _provider( + provider_id="file-search", + provider_type="inline::file-search", + ) + ] + ) + + tools = await get_file_search_tools(client) + + assert len(tools) == 2 + by_id = {tool.identifier: tool for tool in tools} + assert set(by_id) == {"insert_into_memory", "file_search"} + + memory_tool = by_id["insert_into_memory"] + assert memory_tool.description == "Insert documents into memory" + assert memory_tool.parameters == [] + assert memory_tool.provider_id == "file-search" + assert memory_tool.toolgroup_id == "builtin::file_search" + assert memory_tool.server_source == "builtin" + assert memory_tool.type == "tool" + + search_tool = by_id["file_search"] + assert search_tool.description == "Search files for relevant information" + assert len(search_tool.parameters) == 1 + query_param = search_tool.parameters[0] + assert query_param.name == "query" + assert query_param.parameter_type == "string" + assert query_param.required is True + assert query_param.default is None + assert all(tool.provider_id == "file-search" for tool in tools) + assert all(tool.toolgroup_id == "builtin::file_search" for tool in tools) + assert all(tool.server_source == "builtin" for tool in tools) + + client.get.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_file_search_tools_raises_503_on_provider_connection_error( + mocker: MockerFixture, +) -> None: + """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()) + ) + + with pytest.raises(HTTPException) as exc_info: + await get_file_search_tools(client) + + assert exc_info.value.status_code == 503 + detail = exc_info.value.detail + assert isinstance(detail, dict) + assert detail["response"] == "Unable to connect to OGX" diff --git a/tests/unit/utils/test_common.py b/tests/unit/utils/test_common.py deleted file mode 100644 index f902b541f..000000000 --- a/tests/unit/utils/test_common.py +++ /dev/null @@ -1,414 +0,0 @@ -"""Test module for utils/common.py.""" - -from logging import Logger - -import pytest -from pydantic import AnyHttpUrl -from pytest_mock import MockerFixture - -from models.config import ( - Configuration, - LlamaStackConfiguration, - ModelContextProtocolServer, - ServiceConfiguration, - UserDataCollection, -) -from utils.common import ( - register_mcp_servers_async, -) - - -@pytest.mark.asyncio -async def test_register_mcp_servers_empty_list(mocker: MockerFixture) -> None: - """Test register_mcp_servers with empty MCP servers list.""" - mock_logger = mocker.Mock(spec=Logger) - - # Mock the LlamaStack client (shouldn't be called since no MCP servers) - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - - # Create configuration with empty MCP servers - config = Configuration( - name="test", - service=ServiceConfiguration( - host="localhost", - port=1234, - base_url=None, - auth_enabled=True, - workers=10, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - use_as_library_client=False, - url=AnyHttpUrl("http://localhost:8321"), - library_client_config_path=None, - api_key=None, - timeout=60, - ), - user_data_collection=UserDataCollection( - feedback_enabled=False, - feedback_storage=None, - transcripts_enabled=False, - transcripts_storage=None, - ), - mcp_servers=[], - customization=None, - ) # pyright: ignore[reportCallIssue] - # Call the function - await register_mcp_servers_async(mock_logger, config) - - # Verify get_llama_stack_client was NOT called since no MCP servers - mock_lsc.assert_not_called() - # Verify debug message was logged - mock_logger.debug.assert_called_with( - "No MCP servers configured, skipping registration" - ) - - -@pytest.mark.asyncio -async def test_register_mcp_servers_single_server_not_registered( - mocker: MockerFixture, -) -> None: - """Test register_mcp_servers with single MCP server that is not yet registered.""" - # Mock the logger - mock_logger = mocker.Mock(spec=Logger) - - # Mock the LlamaStack client - mock_client = mocker.AsyncMock() - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_client - mock_tool = mocker.Mock() - mock_tool.provider_resource_id = "existing-server" - mock_client.toolgroups.list.return_value = [mock_tool] - mock_client.toolgroups.register.return_value = None - - # Create configuration with one MCP server - mcp_server = ModelContextProtocolServer( - name="new-server", - url="http://localhost:8080", - provider_id="model-context-protocol", - ) - config = Configuration( - name="test", - service=ServiceConfiguration( - host="localhost", - port=1234, - base_url=None, - auth_enabled=True, - workers=10, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - use_as_library_client=False, - url=AnyHttpUrl("http://localhost:8321"), - library_client_config_path=None, - api_key=None, - timeout=60, - ), - user_data_collection=UserDataCollection( - feedback_enabled=False, - feedback_storage=None, - transcripts_enabled=False, - transcripts_storage=None, - ), - mcp_servers=[mcp_server], - customization=None, - ) # pyright: ignore[reportCallIssue] - - # Call the function - await register_mcp_servers_async(mock_logger, config) - - # Verify client.toolgroups.list was called - mock_client.toolgroups.list.assert_called_once() - # Verify client.toolgroups.register was called with correct parameters - mock_client.toolgroups.register.assert_called_once_with( - toolgroup_id="new-server", - provider_id="model-context-protocol", - mcp_endpoint={"uri": "http://localhost:8080"}, - ) - # Verify debug logging was called - mock_logger.debug.assert_called() - - -@pytest.mark.asyncio -async def test_register_mcp_servers_single_server_already_registered( - mocker: MockerFixture, -) -> None: - """Test register_mcp_servers with single MCP server that is already registered.""" - # Mock the logger - mock_logger = mocker.Mock(spec=Logger) - - # Mock the LlamaStack client - mock_client = mocker.AsyncMock() - mock_tool = mocker.Mock() - mock_tool.provider_resource_id = "existing-server" - mock_client.toolgroups.list.return_value = [mock_tool] - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_client - - # Create configuration with MCP server that matches existing toolgroup - mcp_server = ModelContextProtocolServer( - name="existing-server", url="http://localhost:8080", provider_id="qwe" - ) - config = Configuration( - name="test", - service=ServiceConfiguration( - host="localhost", - port=1234, - base_url=None, - auth_enabled=True, - workers=10, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - use_as_library_client=False, - url=AnyHttpUrl("http://localhost:8321"), - library_client_config_path=None, - api_key=None, - timeout=60, - ), - user_data_collection=UserDataCollection( - feedback_enabled=False, - feedback_storage=None, - transcripts_enabled=False, - transcripts_storage=None, - ), - mcp_servers=[mcp_server], - customization=None, - ) # pyright: ignore[reportCallIssue] - - # Call the function - await register_mcp_servers_async(mock_logger, config) - - # Verify client.tools.list was called - mock_client.toolgroups.list.assert_called_once() - # Verify client.toolgroups.register was NOT called since server already registered - assert not mock_client.toolgroups.register.called - - -@pytest.mark.asyncio -async def test_register_mcp_servers_multiple_servers_mixed_registration( - mocker: MockerFixture, -) -> None: - """Test register_mcp_servers with multiple MCP servers - some registered, some not.""" - # Mock the logger - mock_logger = mocker.Mock(spec=Logger) - - # Mock the LlamaStack client - mock_client = mocker.AsyncMock() - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_client - mock_tool1 = mocker.Mock() - mock_tool1.provider_resource_id = "existing-server" - mock_tool2 = mocker.Mock() - mock_tool2.provider_resource_id = "another-existing" - mock_client.toolgroups.list.return_value = [mock_tool1, mock_tool2] - mock_client.toolgroups.register.return_value = None - - # Create configuration with multiple MCP servers - mcp_servers = [ - ModelContextProtocolServer( - name="existing-server", - url="http://localhost:8080", - ), # pyright: ignore[reportCallIssue] - ModelContextProtocolServer( - name="new-server", - url="http://localhost:8081", - ), # pyright: ignore[reportCallIssue] - ModelContextProtocolServer( - name="another-new-server", - provider_id="custom-provider", - url="https://api.example.com", - ), - ] - config = Configuration( - name="test", - service=ServiceConfiguration( - host="localhost", - port=1234, - base_url=None, - auth_enabled=True, - workers=10, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - use_as_library_client=False, - url=AnyHttpUrl("http://localhost:8321"), - library_client_config_path=None, - api_key=None, - timeout=60, - ), - user_data_collection=UserDataCollection( - feedback_enabled=False, - feedback_storage=None, - transcripts_enabled=False, - transcripts_storage=None, - ), - mcp_servers=mcp_servers, - customization=None, - ) # pyright: ignore[reportCallIssue] - - # Call the function - await register_mcp_servers_async(mock_logger, config) - - # Verify client.tools.list was called - mock_client.toolgroups.list.assert_called_once() - # Verify client.toolgroups.register was called twice (for the two new servers) - assert mock_client.toolgroups.register.call_count == 2 - - # Check the specific calls - expected_calls = [ - mocker.call( - toolgroup_id="new-server", - provider_id="model-context-protocol", - mcp_endpoint={"uri": "http://localhost:8081"}, - ), - mocker.call( - toolgroup_id="another-new-server", - provider_id="custom-provider", - mcp_endpoint={"uri": "https://api.example.com"}, - ), - ] - mock_client.toolgroups.register.assert_has_calls(expected_calls, any_order=True) - - -@pytest.mark.asyncio -async def test_register_mcp_servers_with_custom_provider(mocker: MockerFixture) -> None: - """Test register_mcp_servers with MCP server using custom provider.""" - # Mock the logger - mock_logger = mocker.Mock(spec=Logger) - - # Mock the LlamaStack client - mock_client = mocker.AsyncMock() - mock_client.toolgroups.list.return_value = [] - mock_client.toolgroups.register.return_value = None - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_client - - # Create configuration with MCP server using custom provider - mcp_server = ModelContextProtocolServer( - name="custom-server", - provider_id="my-custom-provider", - url="https://custom.example.com/mcp", - ) - config = Configuration( - name="test", - service=ServiceConfiguration( - host="localhost", - port=1234, - base_url=None, - auth_enabled=True, - workers=10, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - use_as_library_client=False, - url=AnyHttpUrl("http://localhost:8321"), - library_client_config_path=None, - api_key=None, - timeout=60, - ), - user_data_collection=UserDataCollection( - feedback_enabled=False, - feedback_storage=None, - transcripts_enabled=False, - transcripts_storage=None, - ), - mcp_servers=[mcp_server], - customization=None, - ) # pyright: ignore[reportCallIssue] - - # Call the function - await register_mcp_servers_async(mock_logger, config) - - # Verify client.toolgroups.register was called with custom provider - mock_client.toolgroups.register.assert_called_once_with( - toolgroup_id="custom-server", - provider_id="my-custom-provider", - mcp_endpoint={"uri": "https://custom.example.com/mcp"}, - ) - - -@pytest.mark.asyncio -async def test_register_mcp_servers_async_with_library_client( - mocker: MockerFixture, -) -> None: - """ - Test that `register_mcp_servers_async` correctly registers MCP - servers when using the library client configuration. - - This test verifies that the function initializes the async - client, checks for existing toolgroups, and registers new MCP - servers as needed when the configuration specifies the use of a - library client. - """ - # Mock the logger - mock_logger = mocker.Mock(spec=Logger) - - # Mock the LlamaStackAsLibraryClient - mock_async_client = mocker.AsyncMock() - mock_async_client.initialize = mocker.AsyncMock() - mock_lsc = mocker.patch("client.AsyncLlamaStackClientHolder.get_client") - mock_lsc.return_value = mock_async_client - - # Mock tools.list to return empty list - mock_tool = mocker.Mock() - mock_tool.provider_resource_id = "existing-tool" - mock_async_client.toolgroups.list = mocker.AsyncMock(return_value=[mock_tool]) - mock_async_client.toolgroups.register = mocker.AsyncMock() - - # Create configuration with library client enabled - mcp_server = ModelContextProtocolServer( - name="test-server", url="http://localhost:8080" - ) # pyright: ignore[reportCallIssue] - config = Configuration( - name="test", - service=ServiceConfiguration( - host="localhost", - port=1234, - base_url=None, - auth_enabled=True, - workers=10, - color_log=True, - access_log=True, - root_path="/.", - ), - llama_stack=LlamaStackConfiguration( - use_as_library_client=True, - library_client_config_path="tests/configuration/run.yaml", - url=None, - api_key=None, - timeout=60, - ), - user_data_collection=UserDataCollection( - feedback_enabled=False, - feedback_storage=None, - transcripts_enabled=False, - transcripts_storage=None, - ), - mcp_servers=[mcp_server], - customization=None, - ) # pyright: ignore[reportCallIssue] - - # Call the async function - await register_mcp_servers_async(mock_logger, config) - - # Verify initialization was called - mock_async_client.initialize.assert_called_once() - # Verify tools.list was called - mock_async_client.toolgroups.list.assert_called_once() - # Verify toolgroups.register was called for the new server - mock_async_client.toolgroups.register.assert_called_once_with( - toolgroup_id="test-server", - provider_id="model-context-protocol", - mcp_endpoint={"uri": "http://localhost:8080"}, - ) diff --git a/tests/unit/utils/test_conversation_compaction.py b/tests/unit/utils/test_conversation_compaction.py index 475b3797c..29678d691 100644 --- a/tests/unit/utils/test_conversation_compaction.py +++ b/tests/unit/utils/test_conversation_compaction.py @@ -8,7 +8,7 @@ from typing import Any, Optional, cast import pytest -from llama_stack_api.openai_responses import OpenAIResponseMessage +from ogx_api.openai_responses import OpenAIResponseMessage from pytest_mock import MockerFixture from models.common.responses.responses_api_params import ResponsesApiParams diff --git a/tests/unit/utils/test_conversations.py b/tests/unit/utils/test_conversations.py index 3003e2e35..50c670aee 100644 --- a/tests/unit/utils/test_conversations.py +++ b/tests/unit/utils/test_conversations.py @@ -5,8 +5,17 @@ import pytest from fastapi import HTTPException -from llama_stack_api import OpenAIResponseMessage -from llama_stack_client import APIConnectionError, APIStatusError +from ogx_api import OpenAIResponseMessage +from ogx_client import APIConnectionError, APIStatusError +from ogx_client.types.conversations.item_list_response import ( + OpenAIResponseInputFunctionToolCallOutputOutputListOpenAIResponseInputMessageContentTextOpenAIResponseInputMessageContentImageOpenAIResponseInputMessageContentFileOpenAIResponseInputMessageContentFile as FunctionCallOutputFile, # pylint: disable=line-too-long +) +from ogx_client.types.conversations.item_list_response import ( + OpenAIResponseInputFunctionToolCallOutputOutputListOpenAIResponseInputMessageContentTextOpenAIResponseInputMessageContentImageOpenAIResponseInputMessageContentFileOpenAIResponseInputMessageContentImage as FunctionCallOutputImage, # pylint: disable=line-too-long +) +from ogx_client.types.conversations.item_list_response import ( + OpenAIResponseInputFunctionToolCallOutputOutputListOpenAIResponseInputMessageContentTextOpenAIResponseInputMessageContentImageOpenAIResponseInputMessageContentFileOpenAIResponseInputMessageContentText as FunctionCallOutputText, # pylint: disable=line-too-long +) from pytest_mock import MockerFixture from constants import DEFAULT_RAG_TOOL @@ -15,7 +24,9 @@ from utils.conversations import ( _build_tool_call_summary_from_item, _extract_text_from_content, + _function_call_output_to_str, append_turn_items_to_conversation, + append_turn_to_conversation, build_conversation_turns_from_items, get_all_conversation_items, ) @@ -354,6 +365,60 @@ 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: + """Test parsing function_call_output with mixed content parts.""" + mock_item = mocker.Mock() + mock_item.type = "function_call_output" + mock_item.call_id = "call_456" + mock_item.status = "success" + mock_item.output = [ + FunctionCallOutputText(type="input_text", text="result text"), + FunctionCallOutputImage( + type="input_image", + image_url="https://example.com/image.png", + ), + FunctionCallOutputFile( + type="input_file", + file_id="file_123", + filename="report.pdf", + ), + ] + + _, tool_result = _build_tool_call_summary_from_item(mock_item) + + assert tool_result is not None + content = tool_result.model_dump()["content"] + assert isinstance(content, str) + assert content.startswith("result text") + assert '"type":"input_image"' in content + assert '"file_id":"file_123"' in content + + +class TestFunctionCallOutputToStr: + """Test cases for _function_call_output_to_str helper.""" + + def test_returns_string_output_unchanged(self) -> None: + """Return plain string output as-is.""" + assert _function_call_output_to_str("plain result") == "plain result" + + def test_extracts_text_and_serializes_other_parts(self) -> None: + """Extract text parts and JSON-serialize image/file parts.""" + output = [ + FunctionCallOutputText(type="input_text", text="hello"), + FunctionCallOutputImage(type="input_image", file_id="img_1"), + FunctionCallOutputFile(type="input_file", filename="data.csv"), + ] + + content = _function_call_output_to_str(output) + + assert content.startswith("hello") + assert '"type":"input_image"' in content + assert '"file_id":"img_1"' in content + assert '"type":"input_file"' in content + assert '"filename":"data.csv"' in content + def test_unknown_item_type(self, mocker: MockerFixture) -> None: """Test parsing an unknown item type.""" mock_item = mocker.Mock() @@ -761,6 +826,37 @@ async def test_appends_user_input_and_llm_output( assert items[1]["content"] == "I cannot help with that" +class TestAppendTurnToConversation: # pylint: disable=too-few-public-methods + """Tests for append_turn_to_conversation function.""" + + @pytest.mark.asyncio + async def test_appends_user_and_assistant_messages( + self, mocker: MockerFixture + ) -> None: + """Test that append_turn_to_conversation creates conversation items correctly.""" + mock_client = mocker.Mock() + mock_client.conversations.items.create = mocker.AsyncMock(return_value=None) + + await append_turn_to_conversation( + mock_client, + conversation_id="conv-123", + user_message="Hello", + assistant_message="I cannot help with that", + ) + + mock_client.conversations.items.create.assert_called_once_with( + "conv-123", + items=[ + {"type": "message", "role": "user", "content": "Hello"}, + { + "type": "message", + "role": "assistant", + "content": "I cannot help with that", + }, + ], + ) + + class TestGetAllConversationItems: """Tests for get_all_conversation_items function.""" @@ -837,7 +933,7 @@ async def test_handles_connection_error(self, mocker: MockerFixture) -> None: await get_all_conversation_items(mock_client, "conv_xyz") assert exc_info.value.status_code == 503 - assert "Llama Stack" in str(exc_info.value.detail) + assert "OGX" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_handles_api_status_error(self, mocker: MockerFixture) -> None: diff --git a/tests/unit/utils/test_endpoints.py b/tests/unit/utils/test_endpoints.py index 2220d25d2..3da8b1774 100644 --- a/tests/unit/utils/test_endpoints.py +++ b/tests/unit/utils/test_endpoints.py @@ -296,7 +296,7 @@ async def test_conversation_id_returns_context_with_existing_conversation( mock_client = mocker.Mock() mock_holder.get_client.return_value = mock_client mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) @@ -336,7 +336,7 @@ async def test_previous_response_id_turn_not_found_raises_404( mock_holder = mocker.Mock() mock_holder.get_client.return_value = mocker.Mock() mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch("utils.endpoints.check_turn_existence", return_value=False) @@ -362,7 +362,7 @@ async def test_previous_response_id_same_as_last_returns_existing_conversation( mock_holder = mocker.Mock() mock_holder.get_client.return_value = mocker.Mock() mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch("utils.endpoints.check_turn_existence", return_value=True) @@ -412,7 +412,7 @@ async def test_previous_response_id_fork_creates_new_conversation( mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch("utils.endpoints.check_turn_existence", return_value=True) @@ -457,7 +457,7 @@ async def test_previous_response_id_fork_respects_generate_topic_summary( mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch("utils.endpoints.check_turn_existence", return_value=True) @@ -500,7 +500,7 @@ async def test_no_context_creates_new_conversation( mock_holder = mocker.Mock() mock_holder.get_client.return_value = mock_client mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch( @@ -528,7 +528,7 @@ async def test_no_context_respects_generate_topic_summary( mock_holder = mocker.Mock() mock_holder.get_client.return_value = mocker.Mock() mocker.patch( - "utils.endpoints.AsyncLlamaStackClientHolder", + "utils.endpoints.AsyncOgxClientHolder", return_value=mock_holder, ) mocker.patch( diff --git a/tests/unit/utils/test_llama_stack_version.py b/tests/unit/utils/test_llama_stack_version.py index f1b3d00bd..3a86be959 100644 --- a/tests/unit/utils/test_llama_stack_version.py +++ b/tests/unit/utils/test_llama_stack_version.py @@ -3,8 +3,8 @@ from typing import Any import pytest -from llama_stack_client import APIConnectionError -from llama_stack_client.types import VersionInfo +from ogx_client import APIConnectionError +from ogx_client.types import VersionInfo from pytest_mock import MockerFixture from pytest_subtests import SubTests from semver import Version diff --git a/tests/unit/utils/test_mcp_tools.py b/tests/unit/utils/test_mcp_tools.py new file mode 100644 index 000000000..491784e08 --- /dev/null +++ b/tests/unit/utils/test_mcp_tools.py @@ -0,0 +1,94 @@ +"""Unit tests for MCP tool discovery utilities.""" + +import httpx +import pytest +from pytest_mock import MockerFixture + +from utils.mcp_tools import _MCP_HTTP_TIMEOUT, list_mcp_tools + + +@pytest.mark.asyncio +async def test_list_mcp_tools_forwards_headers_to_transport( + mocker: MockerFixture, +) -> None: + """Forward headers to the MCP HTTP client, adding Bearer when missing.""" + mock_http_client = mocker.AsyncMock() + mock_http_client.__aenter__ = mocker.AsyncMock(return_value=mock_http_client) + mock_http_client.__aexit__ = mocker.AsyncMock(return_value=None) + + mock_async_client = mocker.patch( + "utils.mcp_tools.httpx.AsyncClient", + return_value=mock_http_client, + ) + mock_streamable = mocker.patch("utils.mcp_tools.streamable_http_client") + mock_streamable.return_value.__aenter__ = mocker.AsyncMock( + return_value=(mocker.Mock(), mocker.Mock(), mocker.Mock()) + ) + mock_streamable.return_value.__aexit__ = mocker.AsyncMock(return_value=None) + mocker.patch( + "utils.mcp_tools._list_tools_from_session", + new=mocker.AsyncMock(return_value=[]), + ) + + await list_mcp_tools( + "http://localhost:3000/mcp", + headers={ + "Authorization": "client-token", + "X-Custom": "value", + }, + ) + + mock_async_client.assert_called_once_with( + headers={ + "Authorization": "Bearer client-token", + "X-Custom": "value", + }, + timeout=_MCP_HTTP_TIMEOUT, + follow_redirects=True, + ) + + +@pytest.mark.asyncio +async def test_list_mcp_tools_returns_empty_list_when_all_transports_fail( + mocker: MockerFixture, +) -> None: + """Skip unavailable MCP servers by returning an empty tool list.""" + mocker.patch( + "utils.mcp_tools._MCP_TRANSPORTS", + ( + ( + "streamable HTTP", + mocker.AsyncMock(side_effect=httpx.ConnectError("boom")), + ), + ("SSE", mocker.AsyncMock(side_effect=httpx.TimeoutException("timeout"))), + ), + ) + + tools = await list_mcp_tools("http://localhost:3000/mcp", headers={}) + + assert tools == [] + + +@pytest.mark.asyncio +async def test_list_mcp_tools_returns_empty_list_on_http_error( + mocker: MockerFixture, +) -> None: + """Skip MCP servers that return HTTP errors.""" + request = httpx.Request("GET", "http://localhost:3000/mcp") + response = httpx.Response(401, request=request) + http_error = httpx.HTTPStatusError( + "HTTP 401", + request=request, + response=response, + ) + mocker.patch( + "utils.mcp_tools._MCP_TRANSPORTS", + ( + ("streamable HTTP", mocker.AsyncMock(side_effect=http_error)), + ("SSE", mocker.AsyncMock(side_effect=http_error)), + ), + ) + + tools = await list_mcp_tools("http://localhost:3000/mcp", headers={}) + + assert tools == [] diff --git a/tests/unit/utils/test_model_list.py b/tests/unit/utils/test_model_list.py new file mode 100644 index 000000000..9fc58875b --- /dev/null +++ b/tests/unit/utils/test_model_list.py @@ -0,0 +1,145 @@ +"""Unit tests for utils/model_list.py helpers.""" + +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model +from ogx_client.types.model_list_response import ( + AnthropicListModelsResponse, + AnthropicListModelsResponseData, + GoogleListModelsResponse, + GoogleListModelsResponseModel, +) + +from models.common.models import CatalogModel +from utils.model_list import ( + parse_anthropic_model, + parse_google_model, + parse_model_list_response, + parse_openai_style_model, +) + + +def test_parse_openai_style_model_with_custom_metadata() -> None: + """OpenAI-style models map custom_metadata into CatalogModel.""" + model = Model.model_construct( + id="provider/model", + created=1, + owned_by="org", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider", + "provider_resource_id": "model", + "extra": "value", + }, + ) + + parsed = parse_openai_style_model(model) + + assert isinstance(parsed, CatalogModel) + assert parsed.identifier == "provider/model" + assert parsed.model_type == "llm" + assert parsed.provider_id == "provider" + assert parsed.provider_resource_id == "model" + assert parsed.metadata == {"extra": "value"} + assert parsed.api_model_type == "llm" + assert parsed.type == "model" + + +def test_parse_anthropic_model() -> None: + """Anthropic models are normalized as LLM CatalogModel entries.""" + model = AnthropicListModelsResponseData.model_construct( + id="claude-sonnet", + created_at="2024-01-01T00:00:00Z", + display_name="Claude Sonnet", + max_input_tokens=200000, + max_tokens=8192, + type="model", + ) + + parsed = parse_anthropic_model(model) + + assert parsed.identifier == "claude-sonnet" + assert parsed.model_type == "llm" + assert parsed.provider_id == "anthropic" + assert parsed.metadata["display_name"] == "Claude Sonnet" + assert parsed.metadata["max_input_tokens"] == 200000 + assert parsed.metadata["max_tokens"] == 8192 + + +def test_parse_google_model() -> None: + """Google models are normalized as LLM CatalogModel entries.""" + model = GoogleListModelsResponseModel.model_construct( + name="models/gemini-pro", + display_name="Gemini Pro", + description="A Gemini model", + ) + + parsed = parse_google_model(model) + + assert parsed.identifier == "models/gemini-pro" + assert parsed.model_type == "llm" + assert parsed.provider_id == "google" + assert parsed.metadata["display_name"] == "Gemini Pro" + assert parsed.metadata["description"] == "A Gemini model" + + +def test_parse_model_list_response_openai() -> None: + """ListModelsResponse branch returns CatalogModel entries.""" + response = ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="p/m", + created=1, + owned_by="x", + custom_metadata={"model_type": "embedding", "provider_id": "p"}, + ) + ] + ) + + parsed = parse_model_list_response(response) + + assert len(parsed) == 1 + assert parsed[0].identifier == "p/m" + assert parsed[0].model_type == "embedding" + + +def test_parse_model_list_response_anthropic() -> None: + """AnthropicListModelsResponse branch returns CatalogModel entries.""" + response = AnthropicListModelsResponse.model_construct( + data=[ + AnthropicListModelsResponseData.model_construct( + id="claude", + created_at="2024-01-01T00:00:00Z", + display_name="Claude", + ) + ] + ) + + parsed = parse_model_list_response(response) + + assert len(parsed) == 1 + assert parsed[0].identifier == "claude" + assert parsed[0].provider_id == "anthropic" + + +def test_parse_model_list_response_google() -> None: + """GoogleListModelsResponse branch returns CatalogModel entries.""" + response = GoogleListModelsResponse.model_construct( + models=[ + GoogleListModelsResponseModel.model_construct( + name="models/gemini", + display_name="Gemini", + ) + ] + ) + + parsed = parse_model_list_response(response) + + assert len(parsed) == 1 + assert parsed[0].identifier == "models/gemini" + assert parsed[0].provider_id == "google" + + +def test_parse_model_list_response_unsupported_type() -> None: + """Unsupported response types yield an empty catalog list.""" + assert parse_model_list_response(object()) == [] # type: ignore[arg-type] diff --git a/tests/unit/utils/test_models_dumper.py b/tests/unit/utils/test_models_dumper.py index 0a6cf445f..07a903131 100644 --- a/tests/unit/utils/test_models_dumper.py +++ b/tests/unit/utils/test_models_dumper.py @@ -2143,8 +2143,7 @@ def test_dump_models(tmpdir: Path) -> None: "computer_call_output.output.image_url", "file_search_call.results", "message.input_image.image_url", - "message.output_text.logprobs", - "reasoning.encrypted_content" + "message.output_text.logprobs" ], "type": "string" }, @@ -2298,7 +2297,7 @@ def test_dump_models(tmpdir: Path) -> None: "type": "string" }, { - "$ref": "`#/components/schemas/`llama_stack_api__openai_responses__ApprovalFilter" + "$ref": "`#/components/schemas/`ogx_api__openai_responses__ApprovalFilter" } ], "default": "never", @@ -7332,7 +7331,7 @@ def test_dump_models(tmpdir: Path) -> None: { "detail": { "cause": "Connection error while trying to reach backend service.", - "response": "Unable to connect to Llama Stack" + "response": "Unable to connect to OGX" }, "label": "llama stack" }, @@ -9085,7 +9084,7 @@ def test_dump_models(tmpdir: Path) -> None: "title": "VectorStoresListResponse", "type": "object" }, - "llama_stack_api__openai_responses__ApprovalFilter": { + "ogx_api__openai_responses__ApprovalFilter": { "description": "Filter configuration for MCP tool approval requirements.\n\n:param always: (Optional) List of tool names that always require approval\n:param never: (Optional) List of tool names that never require approval", "properties": { "always": { @@ -9172,6 +9171,7 @@ def test_dump_models(tmpdir: Path) -> None: "BadRequestResponse", "ByokRag", "CORSConfiguration", + "CatalogShield", "CompactionConfiguration", "Configuration", "ConfigurationResponse", diff --git a/tests/unit/utils/test_pydantic_ai.py b/tests/unit/utils/test_pydantic_ai.py index 05809386f..bc477ce8e 100644 --- a/tests/unit/utils/test_pydantic_ai.py +++ b/tests/unit/utils/test_pydantic_ai.py @@ -2,21 +2,46 @@ # pylint: disable=protected-access +from collections.abc import Callable + import httpx -from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient -from llama_stack_client import AsyncLlamaStackClient +import pytest +from fastapi import HTTPException +from ogx.core.library_client import AsyncOGXAsLibraryClient +from ogx_client import AsyncOgxClient from pydantic_ai_skills import SkillsCapability from pytest_mock import MockerFixture +from configuration import AppConfig from models.common.responses.responses_api_params import ResponsesApiParams -from models.config import SkillsConfiguration +from models.config import ( + QuestionValidityConfig, + QuestionValidityShieldConfiguration, + RedactionConfig, + RedactionShieldConfiguration, + SkillsConfiguration, +) +from pydantic_ai_lightspeed.capabilities import QuestionValidity +from pydantic_ai_lightspeed.capabilities.redaction import PiiRedactionCapability from utils.pydantic_ai_helpers import ( _agent_capabilities, + _shield_capability, _skills_capability, build_agent, get_agent_capability_tools, ) +_QUESTION_VALIDITY_MODULE = ( + "pydantic_ai_lightspeed.capabilities.question_validity._capability" +) + + +@pytest.fixture(autouse=True) +def _mock_question_validity_model(mocker: MockerFixture) -> None: + """Avoid constructing a real client/model when building QuestionValidity.""" + mocker.patch(f"{_QUESTION_VALIDITY_MODULE}.AsyncOgxClientHolder") + mocker.patch(f"{_QUESTION_VALIDITY_MODULE}.OgxResponsesModel.from_ogx_client") + class TestSkillsCapability: """Tests for _skills_capability.""" @@ -39,6 +64,49 @@ def test_returns_capability_for_configured_paths( assert list(capability.toolset.skills) == ["test-skill"] +class TestShieldCapability: + """Tests for _shield_capability.""" + + def test_question_validity_shield_builds_question_validity_capability( + self, + ) -> None: + """Test that a question_validity shield builds a QuestionValidity capability.""" + shield = QuestionValidityShieldConfiguration( + name="topic-guard", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ) + + capability = _shield_capability(shield) + + assert isinstance(capability, QuestionValidity) + assert capability.config is shield.config + + def test_redaction_shield_builds_pii_redaction_capability(self) -> None: + """Test that a redaction shield builds a PiiRedactionCapability.""" + shield = RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig(rules=[]), + ) + + capability = _shield_capability(shield) + + assert isinstance(capability, PiiRedactionCapability) + assert capability.config is shield.config + + def test_unsupported_config_type_raises_value_error( + self, mocker: MockerFixture + ) -> None: + """Test that an unrecognized shield config type raises ValueError.""" + shield = mocker.Mock(name="bad-shield") + shield.name = "bad-shield" + shield.config = object() + + with pytest.raises(ValueError, match="Unsupported shield config type"): + _shield_capability(shield) + + class TestAgentCapabilities: """Tests for _agent_capabilities.""" @@ -46,6 +114,7 @@ def test_returns_none_when_no_capabilities_configured(self) -> None: """Test that missing configuration yields None for Agent construction.""" assert _agent_capabilities(None) is None assert _agent_capabilities(SkillsConfiguration(paths=[])) is None + assert _agent_capabilities(None, shields=[]) is None def test_returns_skills_capability_when_configured( self, mock_skills_configuration: SkillsConfiguration @@ -56,11 +125,55 @@ def test_returns_skills_capability_when_configured( assert len(capabilities) == 1 assert isinstance(capabilities[0], SkillsCapability) + def test_returns_shield_capabilities_when_configured(self) -> None: + """Test that configured shields are included in the capability list.""" + shields = [ + QuestionValidityShieldConfiguration( + name="topic-guard", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig(rules=[]), + ), + ] + + capabilities = _agent_capabilities(None, shields=shields) or [] + + assert len(capabilities) == 2 + assert isinstance(capabilities[0], QuestionValidity) + assert isinstance(capabilities[1], PiiRedactionCapability) + + def test_combines_shields_and_skills( + self, mock_skills_configuration: SkillsConfiguration + ) -> None: + """Test that shield and skill capabilities are both included together.""" + shields = [ + RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig(rules=[]), + ), + ] + + capabilities = ( + _agent_capabilities(mock_skills_configuration, shields=shields) or [] + ) + + capability_types = {type(capability) for capability in capabilities} + assert capability_types == {PiiRedactionCapability, SkillsCapability} + class TestBuildAgent: """Tests for the build_agent factory function.""" - def test_returns_agent_with_correct_model(self, mocker: MockerFixture) -> None: + def test_returns_agent_with_correct_model( + self, + mocker: MockerFixture, + make_agent_config: Callable[..., AppConfig], + ) -> None: """Test that build_agent returns an Agent with the specified model name.""" mock_client = mocker.Mock() mock_client.base_url = "http://localhost:8321" @@ -82,11 +195,15 @@ def test_returns_agent_with_correct_model(self, mocker: MockerFixture) -> None: mock_params.store = False mock_params.previous_response_id = None - agent = build_agent(mock_client, mock_params, None) + agent = build_agent(mock_client, mock_params, make_agent_config()) assert agent is not None - def test_agent_has_instructions(self, mocker: MockerFixture) -> None: + def test_agent_has_instructions( + self, + mocker: MockerFixture, + make_agent_config: Callable[..., AppConfig], + ) -> None: """Test that build_agent passes instructions to the Agent.""" mock_client = mocker.Mock() mock_client.base_url = "http://localhost:8321" @@ -105,13 +222,17 @@ def test_agent_has_instructions(self, mocker: MockerFixture) -> None: mock_params.store = False mock_params.previous_response_id = None - agent = build_agent(mock_client, mock_params, None) + agent = build_agent(mock_client, mock_params, make_agent_config()) assert "You are a helpful assistant." in agent._instructions - def test_agent_with_library_client(self, mocker: MockerFixture) -> None: + def test_agent_with_library_client( + self, + mocker: MockerFixture, + make_agent_config: Callable[..., AppConfig], + ) -> None: """Test that build_agent works with a library client.""" - mock_lib_client = mocker.Mock(spec=AsyncLlamaStackAsLibraryClient) + mock_lib_client = mocker.Mock(spec=AsyncOGXAsLibraryClient) mock_lib_client.provider_data = None mock_params = mocker.Mock() @@ -128,21 +249,22 @@ def test_agent_with_library_client(self, mocker: MockerFixture) -> None: mock_params.store = True mock_params.previous_response_id = None - agent = build_agent(mock_lib_client, mock_params, None) + agent = build_agent(mock_lib_client, mock_params, make_agent_config()) assert agent is not None def test_agent_includes_skills_capability_when_configured( self, - mock_client: AsyncLlamaStackClient, + mock_client: AsyncOgxClient, mock_params: ResponsesApiParams, mock_skills_configuration: SkillsConfiguration, + make_agent_config: Callable[..., AppConfig], ) -> None: - """Test that build_agent attaches SkillsCapability when skills are passed.""" + """Test that build_agent attaches SkillsCapability when skills are configured.""" agent = build_agent( mock_client, mock_params, - mock_skills_configuration, + make_agent_config(skills=mock_skills_configuration), ) capability_types = { @@ -152,28 +274,75 @@ def test_agent_includes_skills_capability_when_configured( def test_agent_has_no_skills_capability_when_not_configured( self, - mock_client: AsyncLlamaStackClient, + mock_client: AsyncOgxClient, mock_params: ResponsesApiParams, + make_agent_config: Callable[..., AppConfig], ) -> None: - """Test that build_agent omits SkillsCapability when skills are not passed.""" - agent = build_agent(mock_client, mock_params, None) + """Test that build_agent omits SkillsCapability when skills are not configured.""" + agent = build_agent(mock_client, mock_params, make_agent_config()) capability_types = { type(capability) for capability in agent._root_capability.capabilities } assert SkillsCapability not in capability_types + def test_agent_includes_shield_capabilities_when_configured( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + make_agent_config: Callable[..., AppConfig], + ) -> None: + """Test that build_agent attaches shield capabilities configured for the app.""" + shields = [ + QuestionValidityShieldConfiguration( + name="topic-guard", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig(rules=[]), + ), + ] + + agent = build_agent( + mock_client, mock_params, make_agent_config(shields=shields) + ) + + capability_types = { + type(capability) for capability in agent._root_capability.capabilities + } + assert QuestionValidity in capability_types + assert PiiRedactionCapability in capability_types + + def test_agent_has_no_shield_capabilities_when_not_configured( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + make_agent_config: Callable[..., AppConfig], + ) -> None: + """Test that build_agent omits shield capabilities when none are configured.""" + agent = build_agent(mock_client, mock_params, make_agent_config()) + + capability_types = { + type(capability) for capability in agent._root_capability.capabilities + } + assert QuestionValidity not in capability_types + assert PiiRedactionCapability not in capability_types + def test_agent_excludes_tool_capabilities_when_no_tools( self, - mock_client: AsyncLlamaStackClient, + mock_client: AsyncOgxClient, mock_params: ResponsesApiParams, mock_skills_configuration: SkillsConfiguration, + make_agent_config: Callable[..., AppConfig], ) -> None: """Test that build_agent omits tool-bearing capabilities when no_tools=True.""" agent = build_agent( mock_client, mock_params, - mock_skills_configuration, + make_agent_config(skills=mock_skills_configuration), no_tools=True, ) @@ -182,6 +351,76 @@ def test_agent_excludes_tool_capabilities_when_no_tools( } assert SkillsCapability not in capability_types + def test_agent_filters_shields_by_name( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + make_agent_config: Callable[..., AppConfig], + ) -> None: + """Test that the shields param filters configured shields by name. + + Mirrors ``QueryRequest.shield_ids``: only shields whose ``name`` is in + the requested list should be attached to the agent. + """ + shields = [ + QuestionValidityShieldConfiguration( + name="topic-guard", + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig(rules=[]), + ), + ] + config = make_agent_config(shields=shields) + + agent = build_agent(mock_client, mock_params, config, shields=["pii-guard"]) + + capability_types = { + type(capability) for capability in agent._root_capability.capabilities + } + assert PiiRedactionCapability in capability_types + assert QuestionValidity not in capability_types + + def test_agent_disables_all_shields_with_empty_list( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + make_agent_config: Callable[..., AppConfig], + ) -> None: + """Test that an empty shields list disables all configured shields.""" + shields = [ + RedactionShieldConfiguration( + name="pii-guard", + provider_id="redaction", + config=RedactionConfig(rules=[]), + ), + ] + config = make_agent_config(shields=shields) + + agent = build_agent(mock_client, mock_params, config, shields=[]) + + capability_types = { + type(capability) for capability in agent._root_capability.capabilities + } + assert PiiRedactionCapability not in capability_types + + def test_agent_raises_not_found_for_unknown_shield_name( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + make_agent_config: Callable[..., AppConfig], + ) -> None: + """Test that requesting an unconfigured shield name raises HTTPException.""" + config = make_agent_config(shields=[]) + + with pytest.raises(HTTPException) as exc_info: + build_agent(mock_client, mock_params, config, shields=["missing-shield"]) + + assert exc_info.value.status_code == 404 + class TestGetAgentCapabilityTools: """Tests for get_agent_capability_tools.""" @@ -197,22 +436,22 @@ def test_returns_skills_tools_when_configured( """Test that configured skills expose pydantic-ai skill tools.""" tools = get_agent_capability_tools(mock_skills_configuration) - assert [tool["identifier"] for tool in tools] == [ + assert [tool.identifier for tool in tools] == [ "list_skills", "load_skill", "read_skill_resource", "run_skill_script", ] assert all( - tool["provider_id"] == "agent-skills" - and tool["toolgroup_id"] == "builtin::agent-skills" - and tool["server_source"] == "builtin" - and tool["type"] == "tool" + tool.provider_id == "agent-skills" + and tool.toolgroup_id == "builtin::agent-skills" + and tool.server_source == "builtin" + and tool.type == "tool" for tool in tools ) - load_skill = next(tool for tool in tools if tool["identifier"] == "load_skill") - assert load_skill["parameters"] == [ + load_skill = next(tool for tool in tools if tool.identifier == "load_skill") + assert [parameter.model_dump() for parameter in load_skill.parameters] == [ { "name": "skill_name", "description": ( diff --git a/tests/unit/utils/test_query.py b/tests/unit/utils/test_query.py index 43ad847d0..d72cd8946 100644 --- a/tests/unit/utils/test_query.py +++ b/tests/unit/utils/test_query.py @@ -9,7 +9,8 @@ import psycopg2 import pytest from fastapi import HTTPException -from llama_stack_client.types import ModelListResponse +from ogx_client.types import ListModelsResponse, ModelListResponse +from ogx_client.types.model import Model from pydantic_ai.messages import ImageUrl from pytest_mock import MockerFixture from sqlalchemy.exc import SQLAlchemyError @@ -33,8 +34,6 @@ consume_query_tokens, extract_provider_and_model_from_model_id, handle_known_apistatus_errors, - is_input_shield, - is_output_shield, is_transcripts_enabled, persist_user_conversation_details, prepare_input, @@ -56,24 +55,25 @@ def mock_config_fixture() -> AppConfig: @pytest.fixture(name="mock_models") def mock_models_fixture() -> ModelListResponse: - """Create mock models list.""" - model1 = type( - "Model", - (), - { - "id": "provider1/model1", - "custom_metadata": {"model_type": "llm", "provider_id": "provider1"}, - }, - )() - model2 = type( - "Model", - (), - { - "id": "provider2/model2", - "custom_metadata": {"model_type": "llm", "provider_id": "provider2"}, - }, - )() - return [model1, model2] + """Create an OpenAI-style OGX models list response.""" + return ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={"model_type": "llm", "provider_id": "provider1"}, + ), + Model.model_construct( + id="provider2/model2", + created=0, + owned_by="test", + object="model", + custom_metadata={"model_type": "llm", "provider_id": "provider2"}, + ), + ] + ) class TestStoreConversationIntoCache: @@ -198,40 +198,6 @@ def test_responses_api_format_without_action(self) -> None: assert exc_info.value.status_code == 403 -class TestShieldFunctions: - """Tests for shield-related functions.""" - - def test_is_output_shield_output_prefix(self) -> None: - """Test is_output_shield returns True for output_ prefix.""" - shield = type("Shield", (), {"identifier": "output_test"})() - assert is_output_shield(shield) is True - - def test_is_output_shield_inout_prefix(self) -> None: - """Test is_output_shield returns True for inout_ prefix.""" - shield = type("Shield", (), {"identifier": "inout_test"})() - assert is_output_shield(shield) is True - - def test_is_output_shield_other(self) -> None: - """Test is_output_shield returns False for other prefixes.""" - shield = type("Shield", (), {"identifier": "input_test"})() - assert is_output_shield(shield) is False - - def test_is_input_shield_input_prefix(self) -> None: - """Test is_input_shield returns True for input prefix.""" - shield = type("Shield", (), {"identifier": "input_test"})() - assert is_input_shield(shield) is True - - def test_is_input_shield_inout_prefix(self) -> None: - """Test is_input_shield returns True for inout_ prefix.""" - shield = type("Shield", (), {"identifier": "inout_test"})() - assert is_input_shield(shield) is True - - def test_is_input_shield_output_prefix(self) -> None: - """Test is_input_shield returns False for output_ prefix.""" - shield = type("Shield", (), {"identifier": "output_test"})() - assert is_input_shield(shield) is False - - class TestPrepareInput: """Tests for prepare_input function.""" diff --git a/tests/unit/utils/test_responses.py b/tests/unit/utils/test_responses.py index d8530cec1..9e8c752e4 100644 --- a/tests/unit/utils/test_responses.py +++ b/tests/unit/utils/test_responses.py @@ -9,48 +9,50 @@ import pytest from fastapi import HTTPException -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( AllowedToolsFilter, OpenAIResponseInputToolChoiceAllowedTools, ) -from llama_stack_api.openai_responses import ApprovalFilter as LlamaStackApprovalFilter -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ApprovalFilter as OgxApprovalFilter +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoiceFileSearch as ToolChoiceFileSearch, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolChoiceMode as ToolChoiceMode, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolFileSearch as InputToolFileSearch, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolFunction as InputToolFunction, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseInputToolWebSearch as InputToolWebSearch, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMCPApprovalRequest as MCPApprovalRequest, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseMCPApprovalResponse as MCPApprovalResponse, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFileSearchToolCall as FileSearchCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageFunctionToolCall as FunctionCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPCall as MCPCall, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageMCPListTools as MCPListTools, ) -from llama_stack_api.openai_responses import ( +from ogx_api.openai_responses import ( OpenAIResponseOutputMessageWebSearchToolCall as WebSearchCall, ) -from llama_stack_client import APIConnectionError, APIStatusError, AsyncLlamaStackClient +from ogx_client import APIConnectionError, APIStatusError, AsyncOgxClient +from ogx_client.types import ListModelsResponse +from ogx_client.types.model import Model from pydantic import AnyUrl, BaseModel from pytest_mock import MockerFixture @@ -449,7 +451,7 @@ async def test_get_mcp_tools_require_approval_filter( tools = await get_mcp_tools(token=None) assert len(tools) == 1 - assert isinstance(tools[0].require_approval, LlamaStackApprovalFilter) + assert isinstance(tools[0].require_approval, OgxApprovalFilter) assert tools[0].require_approval.always == ["create_issue"] assert tools[0].require_approval.never == ["list_repos"] @@ -927,7 +929,7 @@ class TestGetTopicSummary: @pytest.mark.asyncio async def test_get_topic_summary_success(self, mocker: MockerFixture) -> None: """Test successful topic summary generation.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_output_item = make_output_item( item_type="message", role="assistant", content="Topic Summary" ) @@ -949,7 +951,7 @@ async def test_get_topic_summary_empty_response( self, mocker: MockerFixture ) -> None: """Test topic summary with empty response.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_response = mocker.Mock() mock_response.output = [] mock_client.responses.create = mocker.AsyncMock(return_value=mock_response) @@ -967,7 +969,7 @@ async def test_get_topic_summary_connection_error( self, mocker: MockerFixture ) -> None: """Test topic summary raises HTTPException on connection error.""" - mock_client = mocker.AsyncMock(spec=AsyncLlamaStackClient) + mock_client = mocker.AsyncMock(spec=AsyncOgxClient) mock_client.responses.create = mocker.AsyncMock( side_effect=APIConnectionError( message="Connection failed", request=mocker.Mock() @@ -986,7 +988,7 @@ async def test_get_topic_summary_connection_error( @pytest.mark.asyncio 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=AsyncLlamaStackClient) + 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 @@ -1851,10 +1853,22 @@ async def test_prepare_responses_params_with_conversation_id( ) -> None: """Test prepare_responses_params with existing conversation ID.""" mock_client = mocker.AsyncMock() - mock_model = mocker.Mock() - mock_model.id = "provider1/model1" - mock_model.custom_metadata = {"model_type": "llm", "provider_id": "provider1"} - mock_client.models.list = mocker.AsyncMock(return_value=[mock_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider1", + }, + ) + ] + ) + ) query_request = QueryRequest( query="test", conversation_id="123e4567-e89b-12d3-a456-426614174000" @@ -1883,10 +1897,22 @@ async def test_prepare_responses_params_create_conversation( ) -> None: """Test prepare_responses_params creates new conversation when ID not provided.""" mock_client = mocker.AsyncMock() - mock_model = mocker.Mock() - mock_model.id = "provider1/model1" - mock_model.custom_metadata = {"model_type": "llm", "provider_id": "provider1"} - mock_client.models.list = mocker.AsyncMock(return_value=[mock_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider1", + }, + ) + ] + ) + ) mock_conversation = mocker.Mock() mock_conversation.id = "new_conv_id" @@ -1936,10 +1962,22 @@ async def test_prepare_responses_params_connection_error_on_conversation( ) -> None: """Test prepare_responses_params raises HTTPException on connection error when creating conversation.""" mock_client = mocker.AsyncMock() - mock_model = mocker.Mock() - mock_model.id = "provider1/model1" - mock_model.custom_metadata = {"model_type": "llm", "provider_id": "provider1"} - mock_client.models.list = mocker.AsyncMock(return_value=[mock_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider1", + }, + ) + ] + ) + ) mock_client.conversations.create = mocker.AsyncMock( side_effect=APIConnectionError( message="Connection failed", request=mocker.Mock() @@ -1986,10 +2024,22 @@ async def test_prepare_responses_params_includes_mcp_provider_data_headers( ) -> None: """Test that extra_headers with x-llamastack-provider-data is set when MCP tools have headers.""" mock_client = mocker.AsyncMock() - mock_model = mocker.Mock() - mock_model.id = "provider1/model1" - mock_model.custom_metadata = {"model_type": "llm", "provider_id": "provider1"} - mock_client.models.list = mocker.AsyncMock(return_value=[mock_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider1", + }, + ) + ] + ) + ) mock_conversation = mocker.Mock() mock_conversation.id = "new_conv_id" @@ -2050,10 +2100,22 @@ async def test_prepare_responses_params_no_extra_headers_without_mcp_tools( ) -> None: """Test that extra_headers is None when no MCP tools have headers.""" mock_client = mocker.AsyncMock() - mock_model = mocker.Mock() - mock_model.id = "provider1/model1" - mock_model.custom_metadata = {"model_type": "llm", "provider_id": "provider1"} - mock_client.models.list = mocker.AsyncMock(return_value=[mock_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider1", + }, + ) + ] + ) + ) mock_conversation = mocker.Mock() mock_conversation.id = "new_conv_id" @@ -2083,10 +2145,22 @@ async def test_prepare_responses_params_api_status_error_on_conversation( ) -> None: """Test prepare_responses_params raises HTTPException on API status error when creating conversation.""" mock_client = mocker.AsyncMock() - mock_model = mocker.Mock() - mock_model.id = "provider1/model1" - mock_model.custom_metadata = {"model_type": "llm", "provider_id": "provider1"} - mock_client.models.list = mocker.AsyncMock(return_value=[mock_model]) + mock_client.models.list = mocker.AsyncMock( + return_value=ListModelsResponse.model_construct( + data=[ + Model.model_construct( + id="provider1/model1", + created=0, + owned_by="test", + object="model", + custom_metadata={ + "model_type": "llm", + "provider_id": "provider1", + }, + ) + ] + ) + ) mock_client.conversations.create = mocker.AsyncMock( side_effect=APIStatusError( message="API error", response=mocker.Mock(request=None), body=None diff --git a/tests/unit/utils/test_shields.py b/tests/unit/utils/test_shields.py index 745cdf0be..266a4d287 100644 --- a/tests/unit/utils/test_shields.py +++ b/tests/unit/utils/test_shields.py @@ -2,424 +2,29 @@ import pytest from fastapi import HTTPException, status -from llama_stack_client import APIConnectionError, APIStatusError +from pydantic_ai.exceptions import ModelAPIError, ModelHTTPError from pytest_mock import MockerFixture +from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed +from models.config import ( + QuestionValidityConfig, + QuestionValidityShieldConfiguration, + ShieldConfiguration, +) from utils.shields import ( - DEFAULT_VIOLATION_MESSAGE, - append_turn_to_conversation, - detect_shield_violations, - get_available_shields, get_shields_for_request, - run_shield_moderation, + run_shield_moderation_v2, validate_shield_ids_override, ) -class TestGetAvailableShields: - """Tests for get_available_shields function.""" - - @pytest.mark.asyncio - async def test_returns_shield_identifiers(self, mocker: MockerFixture) -> None: - """Test that get_available_shields returns list of shield identifiers.""" - mock_client = mocker.Mock() - shield1 = mocker.Mock() - shield1.identifier = "shield-1" - shield2 = mocker.Mock() - shield2.identifier = "shield-2" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield1, shield2]) - - result = await get_available_shields(mock_client) - - assert result == ["shield-1", "shield-2"] - mock_client.shields.list.assert_called_once() - - @pytest.mark.asyncio - async def test_returns_empty_list_when_no_shields( - self, mocker: MockerFixture - ) -> None: - """Test that get_available_shields returns empty list when no shields available.""" - mock_client = mocker.Mock() - mock_client.shields.list = mocker.AsyncMock(return_value=[]) - - result = await get_available_shields(mock_client) - - assert result == [] - - -class TestDetectShieldViolations: - """Tests for detect_shield_violations function.""" - - def test_detects_violation_when_refusal_present( - self, mocker: MockerFixture - ) -> None: - """Test that detect_shield_violations returns True when refusal is present.""" - mock_record_error = mocker.patch( - "utils.shields.recording.record_llm_validation_error" - ) - - output_item = mocker.Mock(type="message", refusal="Content blocked") - output_items = [output_item] - - result = detect_shield_violations(output_items) - - assert result is True - mock_record_error.assert_called_once() - - def test_returns_false_when_no_violation(self, mocker: MockerFixture) -> None: - """Test that detect_shield_violations returns False when no refusal.""" - mock_record_error = mocker.patch( - "utils.shields.recording.record_llm_validation_error" - ) - - output_item = mocker.Mock(type="message", refusal=None) - output_items = [output_item] - - result = detect_shield_violations(output_items) - - assert result is False - mock_record_error.assert_not_called() - - def test_returns_false_for_non_message_items(self, mocker: MockerFixture) -> None: - """Test that detect_shield_violations ignores non-message items.""" - mock_record_error = mocker.patch( - "utils.shields.recording.record_llm_validation_error" - ) - - output_item = mocker.Mock(type="tool_call", refusal="Content blocked") - output_items = [output_item] - - result = detect_shield_violations(output_items) - - assert result is False - mock_record_error.assert_not_called() - - def test_returns_false_for_empty_list(self, mocker: MockerFixture) -> None: - """Test that detect_shield_violations returns False for empty list.""" - mock_record_error = mocker.patch( - "utils.shields.recording.record_llm_validation_error" - ) - - result = detect_shield_violations([]) - - assert result is False - mock_record_error.assert_not_called() - - -class TestRunShieldModeration: - """Tests for run_shield_moderation function.""" - - @pytest.mark.asyncio - async def test_returns_not_blocked_when_no_shields( - self, mocker: MockerFixture - ) -> None: - """Test that run_shield_moderation returns not blocked when no shields.""" - mock_client = mocker.Mock() - mock_client.shields.list = mocker.AsyncMock(return_value=[]) - mock_client.models.list = mocker.AsyncMock(return_value=[]) - - result = await run_shield_moderation( - mock_client, "test input", "/test-endpoint" - ) - - assert result.decision == "passed" - - @pytest.mark.asyncio - async def test_returns_not_blocked_when_moderation_passes( - self, mocker: MockerFixture - ) -> None: - """Test that run_shield_moderation returns not blocked when content is safe.""" - mock_client = mocker.Mock() - - # Setup shield - shield = mocker.Mock() - shield.identifier = "test-shield" - shield.provider_resource_id = "moderation-model" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - # Setup model - model = mocker.Mock() - model.id = "moderation-model" - mock_client.models.list = mocker.AsyncMock(return_value=[model]) - - # Setup moderation result (not flagged) - moderation_result = mocker.Mock() - moderation_result.results = [mocker.Mock(flagged=False)] - mock_client.moderations.create = mocker.AsyncMock( - return_value=moderation_result - ) - - result = await run_shield_moderation( - mock_client, "safe input", "/test-endpoint" - ) - - assert result.decision == "passed" - mock_client.moderations.create.assert_called_once_with( - input="safe input", model="moderation-model" - ) - - @pytest.mark.asyncio - async def test_returns_blocked_when_content_flagged( - self, mocker: MockerFixture - ) -> None: - """Test that run_shield_moderation returns blocked when content is flagged.""" - mock_record_error = mocker.patch( - "utils.shields.recording.record_llm_validation_error" - ) - mock_client = mocker.Mock() - - # Setup shield - shield = mocker.Mock() - shield.identifier = "test-shield" - shield.provider_resource_id = "moderation-model" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - # Setup model - model = mocker.Mock() - model.id = "moderation-model" - mock_client.models.list = mocker.AsyncMock(return_value=[model]) - - # Setup moderation result (flagged) - flagged_result = mocker.Mock() - flagged_result.flagged = True - flagged_result.categories = ["violence"] - flagged_result.user_message = "Content blocked for violence" - moderation_result = mocker.Mock() - moderation_result.id = "mod_123" - moderation_result.results = [flagged_result] - mock_client.moderations.create = mocker.AsyncMock( - return_value=moderation_result - ) - - result = await run_shield_moderation( - mock_client, "violent content", "/test-endpoint" - ) - - assert result.decision == "blocked" - assert result.message == "Content blocked for violence" - mock_record_error.assert_called_once_with("/test-endpoint") - - @pytest.mark.asyncio - async def test_returns_blocked_with_default_message_when_no_user_message( - self, mocker: MockerFixture - ) -> None: - """Test that run_shield_moderation uses default message when user_message is None.""" - mock_record_error = mocker.patch( - "utils.shields.recording.record_llm_validation_error" - ) - mock_client = mocker.Mock() - - # Setup shield - shield = mocker.Mock() - shield.identifier = "test-shield" - shield.provider_resource_id = "moderation-model" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - # Setup model - model = mocker.Mock() - model.id = "moderation-model" - mock_client.models.list = mocker.AsyncMock(return_value=[model]) - - # Setup moderation result (flagged, no user_message) - flagged_result = mocker.Mock() - flagged_result.flagged = True - flagged_result.categories = ["spam"] - flagged_result.user_message = None - moderation_result = mocker.Mock() - moderation_result.id = "mod_456" - moderation_result.results = [flagged_result] - mock_client.moderations.create = mocker.AsyncMock( - return_value=moderation_result - ) - - result = await run_shield_moderation( - mock_client, "spam content", "/test-endpoint" - ) - - assert result.decision == "blocked" - assert result.message == DEFAULT_VIOLATION_MESSAGE - mock_record_error.assert_called_once() - - @pytest.mark.asyncio - async def test_skips_model_check_for_non_llama_guard_shields( - self, mocker: MockerFixture - ) -> None: - """Test that non-llama-guard shields skip model validation and proceed to moderation.""" - mock_client = mocker.Mock() - - # Setup custom shield (not llama-guard) with provider_resource_id not in models - shield = mocker.Mock() - shield.identifier = "custom-shield" - shield.provider_id = "lightspeed_question_validity" - shield.provider_resource_id = "not-a-model-id" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - # No matching models - should NOT raise for non-llama-guard - mock_client.models.list = mocker.AsyncMock(return_value=[]) - - # Setup moderation result (not flagged) - moderation_result = mocker.Mock() - moderation_result.results = [mocker.Mock(flagged=False)] - mock_client.moderations.create = mocker.AsyncMock( - return_value=moderation_result - ) - - result = await run_shield_moderation( - mock_client, "test input", "/test-endpoint" - ) - - assert result.decision == "passed" - mock_client.moderations.create.assert_called_once_with( - input="test input", model="not-a-model-id" - ) - - @pytest.mark.asyncio - async def test_raises_http_exception_when_shield_model_not_found( - self, mocker: MockerFixture - ) -> None: - """Test that run_shield_moderation raises HTTPException when shield model not in models.""" - mock_client = mocker.Mock() - - # Setup llama-guard shield with provider_resource_id not in models - shield = mocker.Mock() - shield.identifier = "test-shield" - shield.provider_id = "llama-guard" - shield.provider_resource_id = "missing-model" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - # Setup models (doesn't include the shield's model) - model = mocker.Mock() - model.id = "other-model" - mock_client.models.list = mocker.AsyncMock(return_value=[model]) - - with pytest.raises(HTTPException) as exc_info: - await run_shield_moderation(mock_client, "test input", "/test-endpoint") - - assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND - assert "missing-model" in exc_info.value.detail["cause"] # type: ignore[index] - - @pytest.mark.asyncio - async def test_raises_http_exception_when_shield_has_no_provider_resource_id( - self, mocker: MockerFixture - ) -> None: - """Test that run_shield_moderation raises HTTPException when no provider_resource_id.""" - mock_client = mocker.Mock() - - # Setup llama-guard shield without provider_resource_id - shield = mocker.Mock() - shield.identifier = "test-shield" - shield.provider_id = "llama-guard" - shield.provider_resource_id = None - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - mock_client.models.list = mocker.AsyncMock(return_value=[]) - - with pytest.raises(HTTPException) as exc_info: - await run_shield_moderation(mock_client, "test input", "/test-endpoint") - - assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND - - @pytest.mark.asyncio - async def test_shield_ids_empty_list_runs_no_shields_returns_passed( - self, mocker: MockerFixture - ) -> None: - """Test that shield_ids=[] runs no shields and returns passed.""" - mock_client = mocker.Mock() - shield = mocker.Mock() - shield.identifier = "shield-1" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - mock_client.models.list = mocker.AsyncMock(return_value=[]) - - result = await run_shield_moderation( - mock_client, "test input", "/test-endpoint", shield_ids=[] - ) - - assert result.decision == "passed" - - @pytest.mark.asyncio - async def test_shield_ids_raises_404_when_no_shields_found( - self, mocker: MockerFixture - ) -> None: - """Test shield_ids raises HTTPException 404 when requested shield not configured.""" - mock_client = mocker.Mock() - shield = mocker.Mock() - shield.identifier = "shield-1" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) - - with pytest.raises(HTTPException) as exc_info: - await run_shield_moderation( - mock_client, "test input", "/test-endpoint", shield_ids=["typo-shield"] - ) - - assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND - assert "Shield" in exc_info.value.detail["response"] # type: ignore[index] - assert "typo-shield" in exc_info.value.detail["cause"] # type: ignore[index] - - @pytest.mark.asyncio - async def test_shield_ids_filters_to_specific_shield( - self, mocker: MockerFixture - ) -> None: - """Test that shield_ids filters to only specified shields.""" - mock_client = mocker.Mock() - - shield1 = mocker.Mock() - shield1.identifier = "shield-1" - shield1.provider_resource_id = "model-1" - shield2 = mocker.Mock() - shield2.identifier = "shield-2" - shield2.provider_resource_id = "model-2" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield1, shield2]) - - model1 = mocker.Mock() - model1.id = "model-1" - mock_client.models.list = mocker.AsyncMock(return_value=[model1]) - - moderation_result = mocker.Mock() - moderation_result.results = [mocker.Mock(flagged=False)] - mock_client.moderations.create = mocker.AsyncMock( - return_value=moderation_result - ) - - result = await run_shield_moderation( - mock_client, "test input", "/test-endpoint", shield_ids=["shield-1"] - ) - - assert result.decision == "passed" - assert mock_client.moderations.create.call_count == 1 - mock_client.moderations.create.assert_called_with( - input="test input", model="model-1" - ) - - -class TestAppendTurnToConversation: # pylint: disable=too-few-public-methods - """Tests for append_turn_to_conversation function.""" - - @pytest.mark.asyncio - async def test_appends_user_and_assistant_messages( - self, mocker: MockerFixture - ) -> None: - """Test that append_turn_to_conversation creates conversation items correctly.""" - mock_client = mocker.Mock() - mock_client.conversations.items.create = mocker.AsyncMock(return_value=None) - - await append_turn_to_conversation( - mock_client, - conversation_id="conv-123", - user_message="Hello", - assistant_message="I cannot help with that", - ) - - mock_client.conversations.items.create.assert_called_once_with( - "conv-123", - items=[ - {"type": "message", "role": "user", "content": "Hello"}, - { - "type": "message", - "role": "assistant", - "content": "I cannot help with that", - }, - ], - ) +def _shield_config(name: str) -> QuestionValidityShieldConfiguration: + """Build a minimal question-validity shield configuration for tests.""" + return QuestionValidityShieldConfiguration( + name=name, + provider_id="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ) class TestValidateShieldIdsOverride: @@ -503,130 +108,206 @@ def test_raises_422_when_empty_list_shield_ids_and_override_disabled( assert exc_info.value.status_code == status.HTTP_422_UNPROCESSABLE_ENTITY -class TestGetShieldsForRequest: - """Tests for get_shields_for_request function.""" +class TestRunShieldModerationV2: + """Tests for run_shield_moderation_v2 function.""" @pytest.mark.asyncio - async def test_returns_all_shields_when_shield_ids_none( + async def test_returns_passed_when_no_shields(self) -> None: + """Return ShieldModerationPassed when shield list is empty.""" + result = await run_shield_moderation_v2("test input", []) + assert isinstance(result, ShieldModerationPassed) + + @pytest.mark.asyncio + async def test_returns_passed_when_all_shields_pass( self, mocker: MockerFixture ) -> None: - """Return all configured shields when shield_ids is None.""" - mock_client = mocker.Mock() - shield1 = mocker.Mock() - shield1.identifier = "shield-1" - shield2 = mocker.Mock() - shield2.identifier = "shield-2" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield1, shield2]) + """Return ShieldModerationPassed when every shield passes.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock(return_value=ShieldModerationPassed()) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) - result = await get_shields_for_request(mock_client, shield_ids=None) + shields: list[ShieldConfiguration] = [ + _shield_config("s1"), + _shield_config("s2"), + ] + result = await run_shield_moderation_v2("test input", shields) - assert len(result) == 2 - assert result[0].identifier == "shield-1" - assert result[1].identifier == "shield-2" - mock_client.shields.list.assert_called_once() + assert isinstance(result, ShieldModerationPassed) + assert mock_shield.run.call_count == 2 @pytest.mark.asyncio - async def test_returns_empty_list_when_no_shields_configured( - self, mocker: MockerFixture - ) -> None: - """Test that get_shields_for_request returns empty list when no shields configured.""" - mock_client = mocker.Mock() - mock_client.shields.list = mocker.AsyncMock(return_value=[]) + async def test_returns_blocked_on_first_block(self, mocker: MockerFixture) -> None: + """Return blocked result from first shield that blocks.""" + blocked = ShieldModerationBlocked(message="rejected", moderation_id="modr-123") + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock(return_value=blocked) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + shields: list[ShieldConfiguration] = [ + _shield_config("s1"), + _shield_config("s2"), + ] + result = await run_shield_moderation_v2("test input", shields) + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == "rejected" + mock_shield.run.assert_called_once() - result = await get_shields_for_request(mock_client, shield_ids=None) + @pytest.mark.asyncio + async def test_filters_by_selected_shield_ids(self, mocker: MockerFixture) -> None: + """Only run shields matching the selected IDs.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock(return_value=ShieldModerationPassed()) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + shields: list[ShieldConfiguration] = [ + _shield_config("s1"), + _shield_config("s2"), + _shield_config("s3"), + ] + result = await run_shield_moderation_v2( + "test input", shields, selected_shield_ids=["s2"] + ) - assert result == [] + assert isinstance(result, ShieldModerationPassed) + mock_shield.run.assert_called_once() @pytest.mark.asyncio - async def test_filters_to_requested_shields_when_all_exist( - self, mocker: MockerFixture - ) -> None: - """Test that get_shields_for_request returns only requested shields when all exist.""" - mock_client = mocker.Mock() - shield1 = mocker.Mock() - shield1.identifier = "shield-1" - shield2 = mocker.Mock() - shield2.identifier = "shield-2" - shield3 = mocker.Mock() - shield3.identifier = "shield-3" - mock_client.shields.list = mocker.AsyncMock( - return_value=[shield1, shield2, shield3] + async def test_shields_stops_on_first_block(self, mocker: MockerFixture) -> None: + """Stop at the first blocking shield.""" + blocked = ShieldModerationBlocked(message="rejected", moderation_id="modr-789") + mock_qv_shield = mocker.Mock() + mock_qv_shield.run = mocker.AsyncMock(return_value=blocked) + + mock_redact_shield = mocker.Mock() + mock_redact_shield.run = mocker.AsyncMock(return_value=ShieldModerationPassed()) + + mocker.patch( + "utils.shields.build_shield", + side_effect=[mock_qv_shield, mock_redact_shield], ) - result = await get_shields_for_request( - mock_client, shield_ids=["shield-1", "shield-3"] - ) + shields: list[ShieldConfiguration] = [ + _shield_config("s-1"), + _shield_config("s-2"), + ] + result = await run_shield_moderation_v2("test input", shields) - assert len(result) == 2 - assert result[0].identifier == "shield-1" - assert result[1].identifier == "shield-3" + assert isinstance(result, ShieldModerationBlocked) + mock_qv_shield.run.assert_called_once() + mock_redact_shield.run.assert_not_called() @pytest.mark.asyncio - async def test_raises_404_when_requested_shield_not_configured( - self, mocker: MockerFixture - ) -> None: - """Raise 404 when a requested shield is not configured.""" - mock_client = mocker.Mock() - shield = mocker.Mock() - shield.identifier = "shield-1" - mock_client.shields.list = mocker.AsyncMock(return_value=[shield]) + async def test_raise_503_on_model_api_error(self, mocker: MockerFixture) -> None: + """Raise HTTP 503 when a shield raises ModelAPIError.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock( + side_effect=ModelAPIError("test", "Incompatible mode") + ) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) with pytest.raises(HTTPException) as exc_info: - await get_shields_for_request( - mock_client, shield_ids=["shield-1", "missing-shield"] - ) + await run_shield_moderation_v2("test input", [_shield_config("s1")]) - assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND - assert "Shield" in exc_info.value.detail["response"] # type: ignore[index] - assert "missing-shield" in exc_info.value.detail["cause"] # type: ignore[index] + assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert "OGX" in str(exc_info.value.detail) @pytest.mark.asyncio - async def test_raises_404_when_multiple_requested_shields_not_configured( - self, mocker: MockerFixture - ) -> None: - """Raise 404 with all missing ids when multiple shields not configured.""" - mock_client = mocker.Mock() - mock_client.shields.list = mocker.AsyncMock(return_value=[]) + async def test_raise_429_when_exceeds_quota(self, mocker: MockerFixture) -> None: + """Raise HTTP 429 when a shield raises ModelHTTPError with status 429.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock( + side_effect=ModelHTTPError(429, "openai/gpt-4o-mini", "Quota exceeded") + ) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) with pytest.raises(HTTPException) as exc_info: - await get_shields_for_request( - mock_client, shield_ids=["missing-1", "missing-2"] - ) + await run_shield_moderation_v2("test input", [_shield_config("s1")]) - assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND - assert "Shields" in exc_info.value.detail["response"] # type: ignore[index] - cause = exc_info.value.detail["cause"] # type: ignore[index] - assert "missing-1" in cause - assert "missing-2" in cause + assert exc_info.value.status_code == status.HTTP_429_TOO_MANY_REQUESTS + assert "test-model" in str(exc_info.value.detail) + assert "The model quota has been exceeded" in str(exc_info.value.detail) @pytest.mark.asyncio - async def test_raises_503_on_connection_error(self, mocker: MockerFixture) -> None: - """Raise 503 on APIConnectionError.""" - mock_client = mocker.Mock() - mock_client.shields.list = mocker.AsyncMock( - side_effect=APIConnectionError( - message="Connection failed", request=mocker.Mock() + async def test_raise_413_when_exceeds_context_length( + self, mocker: MockerFixture + ) -> None: + """Raise HTTP 413 when a shield raises ModelHTTPError due to context length exceeded.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock( + side_effect=ModelHTTPError( + 413, "openai/gpt-4o-mini", "Context length exceeded" ) ) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) with pytest.raises(HTTPException) as exc_info: - await get_shields_for_request(mock_client, shield_ids=None) + await run_shield_moderation_v2("test input", [_shield_config("s1")]) - assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert exc_info.value.status_code == status.HTTP_413_CONTENT_TOO_LARGE + assert "test-model" in str(exc_info.value.detail) + assert "Prompt is too long" in str(exc_info.value.detail) - @pytest.mark.asyncio - async def test_raises_500_on_api_status_error(self, mocker: MockerFixture) -> None: - """Raise 500 on APIStatusError.""" - mock_client = mocker.Mock() - mock_client.shields.list = mocker.AsyncMock( - side_effect=APIStatusError( - message="Server error", - response=mocker.Mock(request=None), - body=None, - ) + +class TestGetShieldsForRequest: + """Tests for get_shields_for_request function.""" + + def test_returns_all_shields_when_shield_ids_none(self) -> None: + """Return all configured shields when shield_ids is None.""" + shields = [ + _shield_config("shield-1"), + _shield_config("shield-2"), + ] + + result = get_shields_for_request(shields, shield_ids=None) + + assert result == shields + + def test_returns_empty_list_when_shield_ids_empty(self) -> None: + """Return no shields when an empty shield_ids list is provided.""" + shields = [ + _shield_config("shield-1"), + _shield_config("shield-2"), + ] + + result = get_shields_for_request(shields, shield_ids=[]) + + assert result == [] + + def test_filters_to_requested_shields_when_all_exist(self) -> None: + """Return only shields whose names appear in shield_ids.""" + shield1 = _shield_config("shield-1") + shield2 = _shield_config("shield-2") + shield3 = _shield_config("shield-3") + + result = get_shields_for_request( + [shield1, shield2, shield3], shield_ids=["shield-1", "shield-3"] ) + assert result == [shield1, shield3] + + def test_raises_404_when_requested_shield_not_configured(self) -> None: + """Raise 404 when a requested shield name is not configured.""" with pytest.raises(HTTPException) as exc_info: - await get_shields_for_request(mock_client, shield_ids=None) + get_shields_for_request( + [_shield_config("shield-1")], + shield_ids=["shield-1", "missing-shield"], + ) + + assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND + detail = exc_info.value.detail + assert isinstance(detail, dict) + assert "Shield" in detail["response"] + assert "missing-shield" in detail["cause"] - assert exc_info.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR + def test_raises_404_when_multiple_requested_shields_not_configured(self) -> None: + """Raise 404 listing all missing shield names.""" + with pytest.raises(HTTPException) as exc_info: + get_shields_for_request([], shield_ids=["missing-1", "missing-2"]) + + assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND + detail = exc_info.value.detail + assert isinstance(detail, dict) + assert "Shields" in detail["response"] + assert "missing-1" in detail["cause"] + assert "missing-2" in detail["cause"] diff --git a/tests/unit/utils/test_types.py b/tests/unit/utils/test_types.py index eb20f13ae..57e514d11 100644 --- a/tests/unit/utils/test_types.py +++ b/tests/unit/utils/test_types.py @@ -1,7 +1,7 @@ """Unit tests for functions and types defined in utils/types.py.""" import pytest -from llama_stack_api import URL, ImageContentItem, TextContentItem, _URLOrData +from ogx_api import URL, ImageContentItem, TextContentItem, _URLOrData from pydantic import AnyUrl, ValidationError from models.common.responses.responses_api_params import ResponsesApiParams diff --git a/uv.lock b/uv.lock index d2923707d..2d1cbf84f 100644 --- a/uv.lock +++ b/uv.lock @@ -560,15 +560,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/52/93/342cc62a70ab727e093ed98e02a725d85b746345f05d2b5e5034649f4ec8/chevron-0.14.0-py3-none-any.whl", hash = "sha256:fbf996a709f8da2e745ef763f482ce2d311aa817d287593a5b990d6d6e4f0443", size = 11595, upload-time = "2021-01-02T22:47:57.847Z" }, ] -[[package]] -name = "circuitbreaker" -version = "2.1.3" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/df/ac/de7a92c4ed39cba31fe5ad9203b76a25ca67c530797f6bb420fff5f65ccb/circuitbreaker-2.1.3.tar.gz", hash = "sha256:1a4baee510f7bea3c91b194dcce7c07805fe96c4423ed5594b75af438531d084", size = 10787, upload-time = "2025-03-31T08:12:08.963Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/ae/34/15f08edd4628f65217de1fc3c1a27c82e46fe357d60c217fc9881e12ebcc/circuitbreaker-2.1.3-py3-none-any.whl", hash = "sha256:87ba6a3ed03fdc7032bc175561c2b04d52ade9d5faf94ca2b035fbdc5e6b1dd1", size = 7737, upload-time = "2025-03-31T08:12:07.802Z" }, -] - [[package]] name = "click" version = "8.4.2" @@ -629,41 +620,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ec/82/32e3bd191d498e64f6f911ad55d14006a0861e54869d2d32452326399e65/coverage-7.15.2-py3-none-any.whl", hash = "sha256:eb6bcae8d1a9d305351ecb108232441d11c5cfe9de840a04388ba5d2db8d735c", size = 213375, upload-time = "2026-07-15T18:56:17.305Z" }, ] -[[package]] -name = "crc32c" -version = "2.8" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/e3/66/7e97aa77af7cf6afbff26e3651b564fe41932599bc2d3dce0b2f73d4829a/crc32c-2.8.tar.gz", hash = "sha256:578728964e59c47c356aeeedee6220e021e124b9d3e8631d95d9a5e5f06e261c", size = 48179, upload-time = "2025-10-17T06:20:13.61Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/b6/36/fd18ef23c42926b79c7003e16cb0f79043b5b179c633521343d3b499e996/crc32c-2.8-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:572ffb1b78cce3d88e8d4143e154d31044a44be42cb3f6fbbf77f1e7a941c5ab", size = 66379, upload-time = "2025-10-17T06:19:10.115Z" }, - { url = "https://files.pythonhosted.org/packages/7f/b8/c584958e53f7798dd358f5bdb1bbfc97483134f053ee399d3eeb26cca075/crc32c-2.8-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:cf827b3758ee0c4aacd21ceca0e2da83681f10295c38a10bfeb105f7d98f7a68", size = 63042, upload-time = "2025-10-17T06:19:10.946Z" }, - { url = "https://files.pythonhosted.org/packages/62/e6/6f2af0ec64a668a46c861e5bc778ea3ee42171fedfc5440f791f470fd783/crc32c-2.8-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:106fbd79013e06fa92bc3b51031694fcc1249811ed4364ef1554ee3dd2c7f5a2", size = 61528, upload-time = "2025-10-17T06:19:11.768Z" }, - { url = "https://files.pythonhosted.org/packages/17/8b/4a04bd80a024f1a23978f19ae99407783e06549e361ab56e9c08bba3c1d3/crc32c-2.8-cp312-cp312-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:6dde035f91ffbfe23163e68605ee5a4bb8ceebd71ed54bb1fb1d0526cdd125a2", size = 80028, upload-time = "2025-10-17T06:19:12.554Z" }, - { url = "https://files.pythonhosted.org/packages/21/8f/01c7afdc76ac2007d0e6a98e7300b4470b170480f8188475b597d1f4b4c6/crc32c-2.8-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e41ebe7c2f0fdcd9f3a3fd206989a36b460b4d3f24816d53e5be6c7dba72c5e1", size = 81531, upload-time = "2025-10-17T06:19:13.406Z" }, - { url = "https://files.pythonhosted.org/packages/32/2b/8f78c5a8cc66486be5f51b6f038fc347c3ba748d3ea68be17a014283c331/crc32c-2.8-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:ecf66cf90266d9c15cea597d5cc86c01917cd1a238dc3c51420c7886fa750d7e", size = 80608, upload-time = "2025-10-17T06:19:14.223Z" }, - { url = "https://files.pythonhosted.org/packages/db/86/fad1a94cdeeeb6b6e2323c87f970186e74bfd6fbfbc247bf5c88ad0873d5/crc32c-2.8-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:59eee5f3a69ad0793d5fa9cdc9b9d743b0cd50edf7fccc0a3988a821fef0208c", size = 79886, upload-time = "2025-10-17T06:19:15.345Z" }, - { url = "https://files.pythonhosted.org/packages/d5/db/1a7cb6757a1e32376fa2dfce00c815ea4ee614a94f9bff8228e37420c183/crc32c-2.8-cp312-cp312-win32.whl", hash = "sha256:a73d03ce3604aa5d7a2698e9057a0eef69f529c46497b27ee1c38158e90ceb76", size = 64896, upload-time = "2025-10-17T06:19:16.457Z" }, - { url = "https://files.pythonhosted.org/packages/bf/8e/2024de34399b2e401a37dcb54b224b56c747b0dc46de4966886827b4d370/crc32c-2.8-cp312-cp312-win_amd64.whl", hash = "sha256:56b3b7d015247962cf58186e06d18c3d75a1a63d709d3233509e1c50a2d36aa2", size = 66645, upload-time = "2025-10-17T06:19:17.235Z" }, - { url = "https://files.pythonhosted.org/packages/e8/d8/3ae227890b3be40955a7144106ef4dd97d6123a82c2a5310cdab58ca49d8/crc32c-2.8-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:36f1e03ee9e9c6938e67d3bcb60e36f260170aa5f37da1185e04ef37b56af395", size = 66380, upload-time = "2025-10-17T06:19:18.009Z" }, - { url = "https://files.pythonhosted.org/packages/bd/8b/178d3f987cd0e049b484615512d3f91f3d2caeeb8ff336bb5896ae317438/crc32c-2.8-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:b2f3226b94b85a8dd9b3533601d7a63e9e3e8edf03a8a169830ee8303a199aeb", size = 63048, upload-time = "2025-10-17T06:19:18.853Z" }, - { url = "https://files.pythonhosted.org/packages/f2/a1/48145ae2545ebc0169d3283ebe882da580ea4606bfb67cf4ca922ac3cfc3/crc32c-2.8-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:6e08628bc72d5b6bc8e0730e8f142194b610e780a98c58cb6698e665cb885a5b", size = 61530, upload-time = "2025-10-17T06:19:19.974Z" }, - { url = "https://files.pythonhosted.org/packages/06/4b/cf05ed9d934cc30e5ae22f97c8272face420a476090e736615d9a6b53de0/crc32c-2.8-cp313-cp313-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:086f64793c5ec856d1ab31a026d52ad2b895ac83d7a38fce557d74eb857f0a82", size = 80001, upload-time = "2025-10-17T06:19:20.784Z" }, - { url = "https://files.pythonhosted.org/packages/15/ab/4b04801739faf36345f6ba1920be5b1c70282fec52f8280afd3613fb13e2/crc32c-2.8-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:bcf72ee7e0135b3d941c34bb2c26c3fc6bc207106b49fd89aaafaeae223ae209", size = 81543, upload-time = "2025-10-17T06:19:21.557Z" }, - { url = "https://files.pythonhosted.org/packages/a9/1b/6e38dde5bfd2ea69b7f2ab6ec229fcd972a53d39e2db4efe75c0ac0382ce/crc32c-2.8-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:8a717dd9c3fd777d9bc6603717eae172887d402c4ab589d124ebd0184a83f89e", size = 80644, upload-time = "2025-10-17T06:19:22.325Z" }, - { url = "https://files.pythonhosted.org/packages/ce/45/012176ffee90059ae8ec7131019c71724ea472aa63e72c0c8edbd1fad1d7/crc32c-2.8-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:0450bb845b3c3c7b9bdc0b4e95620ec9a40824abdc8c86d6285c919a90743c1a", size = 79919, upload-time = "2025-10-17T06:19:23.101Z" }, - { url = "https://files.pythonhosted.org/packages/f0/2b/f557629842f9dec2b3461cb3a0d854bb586ec45b814cea58b082c32f0dde/crc32c-2.8-cp313-cp313-win32.whl", hash = "sha256:765d220bfcbcffa6598ac11eb1e10af0ee4802b49fe126aa6bf79f8ddb9931d1", size = 64896, upload-time = "2025-10-17T06:19:23.88Z" }, - { url = "https://files.pythonhosted.org/packages/d0/db/fd0f698c15d1e21d47c64181a98290665a08fcbb3940cd559e9c15bda57e/crc32c-2.8-cp313-cp313-win_amd64.whl", hash = "sha256:171ff0260d112c62abcce29332986950a57bddee514e0a2418bfde493ea06bb3", size = 66646, upload-time = "2025-10-17T06:19:24.702Z" }, - { url = "https://files.pythonhosted.org/packages/db/b9/8e5d7054fe8e7eecab10fd0c8e7ffb01439417bdb6de1d66a81c38fc4a20/crc32c-2.8-cp313-cp313t-macosx_10_13_universal2.whl", hash = "sha256:b977a32a3708d6f51703c8557008f190aaa434d7347431efb0e86fcbe78c2a50", size = 66203, upload-time = "2025-10-17T06:19:25.872Z" }, - { url = "https://files.pythonhosted.org/packages/55/5f/cc926c70057a63cc0c98a3c8a896eb15fc7e74d3034eadd53c94917c6cc3/crc32c-2.8-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:7399b01db4adaf41da2fb36fe2408e75a8d82a179a9564ed7619412e427b26d6", size = 62956, upload-time = "2025-10-17T06:19:26.652Z" }, - { url = "https://files.pythonhosted.org/packages/a1/8a/0660c44a2dd2cb6ccbb529eb363b9280f5c766f1017bc8355ed8d695bd94/crc32c-2.8-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:4379f73f9cdad31958a673d11a332ec725ca71572401ca865867229f5f15e853", size = 61442, upload-time = "2025-10-17T06:19:27.74Z" }, - { url = "https://files.pythonhosted.org/packages/f5/5a/6108d2dfc0fe33522ce83ba07aed4b22014911b387afa228808a278e27cd/crc32c-2.8-cp313-cp313t-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:2e68264555fab19bab08331550dab58573e351a63ed79c869d455edd3b0aa417", size = 79109, upload-time = "2025-10-17T06:19:28.535Z" }, - { url = "https://files.pythonhosted.org/packages/84/1e/c054f9e390090c197abf3d2936f4f9effaf0c6ee14569ae03d6ddf86958a/crc32c-2.8-cp313-cp313t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:b48f2486727b8d0e7ccbae4a34cb0300498433d2a9d6b49cb13cb57c2e3f19cb", size = 80987, upload-time = "2025-10-17T06:19:29.305Z" }, - { url = "https://files.pythonhosted.org/packages/c8/ad/1650e5c3341e4a485f800ea83116d72965030c5d48ccc168fcc685756e4d/crc32c-2.8-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:ecf123348934a086df8c8fde7f9f2d716d523ca0707c5a1367b8bb00d8134823", size = 79994, upload-time = "2025-10-17T06:19:30.109Z" }, - { url = "https://files.pythonhosted.org/packages/d7/3b/f2ed924b177729cbb2ab30ca2902abff653c31d48c95e7b66717a9ca9fcc/crc32c-2.8-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:e636ac60f76de538f7a2c0d0f3abf43104ee83a8f5e516f6345dc283ed1a4df7", size = 79046, upload-time = "2025-10-17T06:19:30.894Z" }, - { url = "https://files.pythonhosted.org/packages/4b/80/413b05ee6ace613208b31b3670c3135ee1cf451f0e72a9c839b4946acc04/crc32c-2.8-cp313-cp313t-win32.whl", hash = "sha256:8dd4a19505e0253892e1b2f1425cc3bd47f79ae5a04cb8800315d00aad7197f2", size = 64837, upload-time = "2025-10-17T06:19:32.03Z" }, - { url = "https://files.pythonhosted.org/packages/3b/1b/85eddb6ac5b38496c4e35c20298aae627970c88c3c624a22ab33e84f16c7/crc32c-2.8-cp313-cp313t-win_amd64.whl", hash = "sha256:4bb18e4bd98fb266596523ffc6be9c5b2387b2fa4e505ec56ca36336f49cb639", size = 66574, upload-time = "2025-10-17T06:19:33.143Z" }, -] - [[package]] name = "cryptography" version = "49.0.0" @@ -880,7 +836,7 @@ wheels = [ [[package]] name = "fastapi" -version = "0.140.13" +version = "0.141.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "annotated-doc" }, @@ -889,9 +845,9 @@ dependencies = [ { name = "typing-extensions" }, { name = "typing-inspection" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/2f/cb/7a4d2c2eb5a5d8a91763c05b7383d72917862e32f780daa0e27ffbb34cc6/fastapi-0.140.13.tar.gz", hash = "sha256:500172a08cf1459901f90b05c37d93060dada3b573fec8f0862445db52ba6b4b", size = 424843, upload-time = "2026-07-28T15:37:00.805Z" } +sdist = { url = "https://files.pythonhosted.org/packages/8a/02/91e3416a8fdd715abb903a952a6bec7cdd8d14eed55d415fc8595524c319/fastapi-0.141.1.tar.gz", hash = "sha256:e8822fc40db1e1858054d7a949a888695bc9bdce70139178e33bd2871a453ca1", size = 425799, upload-time = "2026-07-29T17:18:05.568Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/84/4e/f9e8c762ef5e05c40482131e3d5e8b36bca13fa127578261f1d6b35a25d4/fastapi-0.140.13-py3-none-any.whl", hash = "sha256:8b017110e1e9f30a95e8bdb8f71fbe2f0fe3af5717109e5b14f9e069df54f6d4", size = 131222, upload-time = "2026-07-28T15:37:02.124Z" }, + { url = "https://files.pythonhosted.org/packages/cb/03/10388a42375ee7e4ac9b94eb2c5c569c8b5795e377e701c9ac3ad63de890/fastapi-0.141.1-py3-none-any.whl", hash = "sha256:bfb91aa2d334c61cb35ba9a116fc123b3d3df31640b801cf57a7a78ec3f603b3", size = 131954, upload-time = "2026-07-29T17:18:04.364Z" }, ] [[package]] @@ -954,11 +910,11 @@ wheels = [ [[package]] name = "filelock" -version = "3.32.0" +version = "3.32.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/c0/80/8232b582c4b318b817cf1274ba74976b07b34d35ef439b3eb948f98645a1/filelock-3.32.0.tar.gz", hash = "sha256:7be2ad23a14607ccc71808e68fe30848aeace7058ace17852f68e2a68e310402", size = 213757, upload-time = "2026-07-21T13:17:42.898Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f6/57/3ba6e6cb097f85b855b00163d169f35365f44277df044dcf96d55b8f62a3/filelock-3.32.2.tar.gz", hash = "sha256:c33351e1f49cae33414acbc6d56784e6ecee82514ec90795da1161fc4836b5b8", size = 217172, upload-time = "2026-07-29T22:46:04.895Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/06/79/b4c714bef36bc4ec2beeae1e0c124f0223888cd8c6feb1cdc56038116920/filelock-3.32.0-py3-none-any.whl", hash = "sha256:d396bea984af47333ef05e50eae7eff88c84256de6112aea0ec48a233c064fe3", size = 97732, upload-time = "2026-07-21T13:17:41.55Z" }, + { url = "https://files.pythonhosted.org/packages/c1/e8/72f8cef9fdfeffe06213fe8508039396ee48daa0e3259457ed766173bfd6/filelock-3.32.2-py3-none-any.whl", hash = "sha256:87dd94cf281e586d135fa51132b8e3d9a598b316e90377a288663c9321036c82", size = 98830, upload-time = "2026-07-29T22:46:03.52Z" }, ] [[package]] @@ -1046,15 +1002,15 @@ http = [ [[package]] name = "genai-prices" -version = "0.0.72" +version = "0.0.73" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "httpx2" }, { name = "pydantic" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/6f/c8/2549fa8ceaaf0bd61114cf392250a107be657233723bc4e65cb91dba934f/genai_prices-0.0.72.tar.gz", hash = "sha256:a7e481d0ea85922fcf48df6864f2491fc81a4f03dcf0ecbd9aa8c2c6d9fcdbe8", size = 82753, upload-time = "2026-07-22T21:11:03.997Z" } +sdist = { url = "https://files.pythonhosted.org/packages/84/0b/4b430d9c6ff0c76e42b91fc4cedff84dfbefa1a35b885fad45a22d29303f/genai_prices-0.0.73.tar.gz", hash = "sha256:ddad5b23dadd7aac8a58ab5af7506addb062114037e96427324703e50bc35077", size = 88406, upload-time = "2026-07-29T12:48:58.215Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/22/2e/1e2c31666f7e4e2d0327a80f1c920fd8764ccd27672c4bc389d79219ece6/genai_prices-0.0.72-py3-none-any.whl", hash = "sha256:21281d8df34d9bfbb736e0763a60ceda02226938cd47379ae437a44fcf8ba23d", size = 85238, upload-time = "2026-07-22T21:11:02.911Z" }, + { url = "https://files.pythonhosted.org/packages/a8/85/527a729ecd58b7169430e47febdab2e3f91193d0e3f5a56d721c04525709/genai_prices-0.0.73-py3-none-any.whl", hash = "sha256:64dfbbf55df91db6b74ff2b24ac752b8154952e0aed4e62868d6cd42710bf742", size = 92507, upload-time = "2026-07-29T12:48:57.194Z" }, ] [[package]] @@ -1209,7 +1165,7 @@ wheels = [ [[package]] name = "google-genai" -version = "2.14.0" +version = "2.15.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio" }, @@ -1223,9 +1179,9 @@ dependencies = [ { name = "typing-extensions" }, { name = "websockets" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/dc/df/4f820054c99f29f2fe3de4a8a7c9534dd795302e4a07483a0cb07c3a29b6/google_genai-2.14.0.tar.gz", hash = "sha256:a9d1f4f362d76280f1be1340fcb3c86e63dbca128f6a4ae09d86ab47ff7148e8", size = 641055, upload-time = "2026-07-22T21:35:44.717Z" } +sdist = { url = "https://files.pythonhosted.org/packages/69/53/b2c9b0a74b817a393d388a2303ec4da8bda27ea744b23914480d1b024d84/google_genai-2.15.0.tar.gz", hash = "sha256:ef71bdb79ce9931bca1cf0a393c8cfb606e1075b6100fcdde02b7b467db8235d", size = 640674, upload-time = "2026-07-29T17:43:21.853Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ba/86/5ac5fb53e44cca4a6607fb917eb331fa237c65a103b9ec2e8e8acc8a42db/google_genai-2.14.0-py3-none-any.whl", hash = "sha256:ae7172cdd35695189b516b33a878e4132e5daa2dbc03a5b44cddfa8a82fad664", size = 1030738, upload-time = "2026-07-22T21:35:42.785Z" }, + { url = "https://files.pythonhosted.org/packages/cf/86/2ef6d992955307bf525f305b44cb3e2066eb01636549607082a0407f7047/google_genai-2.15.0-py3-none-any.whl", hash = "sha256:f322a94c3c1ddb1b1cc536086708f5f8e13101347d062f63dcd7947b598c1d16", size = 1030459, upload-time = "2026-07-29T17:43:20.286Z" }, ] [[package]] @@ -1768,9 +1724,9 @@ dependencies = [ { name = "jsonpath-ng" }, { name = "kubernetes" }, { name = "litellm" }, - { name = "llama-stack" }, - { name = "llama-stack-api" }, - { name = "llama-stack-client" }, + { name = "ogx" }, + { name = "ogx-api" }, + { name = "ogx-client" }, { name = "openai" }, { name = "opentelemetry-distro" }, { name = "opentelemetry-exporter-otlp" }, @@ -1880,9 +1836,9 @@ requires-dist = [ { name = "jsonpath-ng", specifier = ">=1.6.1" }, { name = "kubernetes", specifier = ">=30.1.0" }, { name = "litellm", specifier = ">=1.83.7" }, - { name = "llama-stack", specifier = "==0.6.0" }, - { name = "llama-stack-api", specifier = "==0.6.0" }, - { name = "llama-stack-client", specifier = "==0.6.0" }, + { name = "ogx", specifier = "==1.0.2" }, + { name = "ogx-api", specifier = "==1.0.2" }, + { name = "ogx-client", specifier = "==1.0.2" }, { name = "openai", specifier = ">=1.99.9" }, { name = "opentelemetry-distro", specifier = ">=0.49b0" }, { name = "opentelemetry-exporter-otlp", specifier = ">=1.34.1" }, @@ -2001,93 +1957,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/3f/2e/b0506b0308b36b41fda0a195e07a2202541cd9c71750d75653e6149bfc34/litellm-1.94.0-cp313-cp313-win_amd64.whl", hash = "sha256:dd05c0dee67325f5a580e2adc43fbe21c0ab8b2f2641991b33d096f10fe8c665", size = 20268579, upload-time = "2026-07-28T20:21:41.401Z" }, ] -[[package]] -name = "llama-stack" -version = "0.6.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "aiohttp" }, - { name = "aiosqlite" }, - { name = "asyncpg" }, - { name = "fastapi" }, - { name = "fire" }, - { name = "h11" }, - { name = "httpx" }, - { name = "jinja2" }, - { name = "jsonschema" }, - { name = "llama-stack-api" }, - { name = "mcp" }, - { name = "numpy" }, - { name = "oci" }, - { name = "openai" }, - { name = "opentelemetry-distro" }, - { name = "opentelemetry-exporter-otlp-proto-http" }, - { name = "opentelemetry-sdk" }, - { name = "oracledb" }, - { name = "prompt-toolkit" }, - { name = "psycopg2-binary" }, - { name = "pydantic" }, - { name = "pyjwt", extra = ["crypto"] }, - { name = "python-dotenv" }, - { name = "python-multipart" }, - { name = "pyyaml" }, - { name = "rich" }, - { name = "sqlalchemy", extra = ["asyncio"] }, - { name = "starlette" }, - { name = "termcolor" }, - { name = "tiktoken" }, - { name = "tornado" }, - { name = "urllib3" }, - { name = "uvicorn" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/f4/53/5bc3ae19e9a42475b42682456898ced5a0b48e43f918920fc790b665f223/llama_stack-0.6.0.tar.gz", hash = "sha256:d92711791633f5505a4473ffba3f3e26acb700716fddab5aec419d99e614c802", size = 13631563, upload-time = "2026-03-11T15:06:13.071Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/9d/17/b4427e1db7409f698c95b6f8b2b9e662bbcf1f819beb1af180bab55ddfb5/llama_stack-0.6.0-py3-none-any.whl", hash = "sha256:b804830664dc91e54c7225a7a081cb1874c48fc18573569c19fac4a9397e8076", size = 770027, upload-time = "2026-03-11T15:06:10.649Z" }, -] - -[[package]] -name = "llama-stack-api" -version = "0.6.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "fastapi" }, - { name = "jsonschema" }, - { name = "openai" }, - { name = "opentelemetry-exporter-otlp-proto-http" }, - { name = "opentelemetry-sdk" }, - { name = "pydantic" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/f3/4f/0c6fbc861fb9f6074f877f9c39ea40b02ad1fc81c9a455b020b32dcc471f/llama_stack_api-0.6.0.tar.gz", hash = "sha256:f0f3a1a6239a5d3b8c7ef02cefdf817c96c6461dcd8a82c1689ac67ec3107270", size = 136402, upload-time = "2026-03-11T15:05:30.843Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/4d/44/3c7a8dc82ddcc45a375681051450979837f28b27250ee057cabcfb8421f3/llama_stack_api-0.6.0-py3-none-any.whl", hash = "sha256:b99a03aba3659736b6b540c9e5e674b1daac2bf5eeb2a68795113d62b8250672", size = 161069, upload-time = "2026-03-11T15:05:29.072Z" }, -] - -[[package]] -name = "llama-stack-client" -version = "0.6.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "anyio" }, - { name = "click" }, - { name = "distro" }, - { name = "fire" }, - { name = "httpx" }, - { name = "pandas" }, - { name = "prompt-toolkit" }, - { name = "pyaml" }, - { name = "pydantic" }, - { name = "requests" }, - { name = "rich" }, - { name = "sniffio" }, - { name = "termcolor" }, - { name = "tqdm" }, - { name = "typing-extensions" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/b7/e9/62dc71e7d6003d9b56a1e632445065f55687c891e62eff1636e10b5dd629/llama_stack_client-0.6.0.tar.gz", hash = "sha256:3290aac36dcafbd1bc0baaf995522e2037f57056672b5a1516af112a4210f3ea", size = 368695, upload-time = "2026-03-11T15:04:19.267Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/83/a3/33d3e066a320a993b6f9cca9c8efe8da7deb2045df61235d327d0a05b25f/llama_stack_client-0.6.0-py3-none-any.whl", hash = "sha256:7e514a6ffd92f237aceb062dadc4db44e24a3cd9c4ea35e25173d1e0739beb8e", size = 392001, upload-time = "2026-03-11T15:04:17.772Z" }, -] - [[package]] name = "logfire" version = "4.39.0" @@ -2496,23 +2365,80 @@ wheels = [ ] [[package]] -name = "oci" -version = "2.183.0" +name = "ogx" +version = "1.0.2" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "certifi" }, - { name = "circuitbreaker" }, - { name = "crc32c" }, - { name = "cryptography" }, - { name = "pyjwt" }, - { name = "pyopenssl" }, - { name = "python-dateutil" }, - { name = "pytz" }, - { name = "urllib3" }, + { name = "aiosqlite" }, + { name = "asyncpg" }, + { name = "fastapi" }, + { name = "httpx" }, + { name = "jinja2" }, + { name = "jsonschema" }, + { name = "mcp" }, + { name = "ogx-api" }, + { name = "openai" }, + { name = "opentelemetry-distro" }, + { name = "opentelemetry-exporter-otlp-proto-http" }, + { name = "opentelemetry-sdk" }, + { name = "pydantic" }, + { name = "pyjwt", extra = ["crypto"] }, + { name = "python-dotenv" }, + { name = "pyyaml" }, + { name = "rich" }, + { name = "sqlalchemy", extra = ["asyncio"] }, + { name = "structlog" }, + { name = "termcolor" }, + { name = "tiktoken" }, + { name = "uvicorn" }, + { name = "zstandard" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1e/2a/77bd6cbf1c69b2f368fe3d6462d84369b0cba15e37ce713cdc08d459b95a/oci-2.183.0.tar.gz", hash = "sha256:ff572ef5f2030a788796bb509d257e6a41c6510ef9b4b6a75a079efd06e533ce", size = 17759723, upload-time = "2026-07-28T06:02:29.76Z" } +sdist = { url = "https://files.pythonhosted.org/packages/44/a0/c8f0c17c7297f8f989d8246ebd30312b08a0e707c83a6bd433e9954dbbb4/ogx-1.0.2.tar.gz", hash = "sha256:b4978ed93f89d7b4f91c70610e338a53265cb3c4c6a00dd5e6b572ba38ed8ccd", size = 16911718, upload-time = "2026-05-13T22:53:53.338Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a9/de/8574b3e527996a099d196e87794a4652d91a0c3185fcc7fdbb5649b75a8a/oci-2.183.0-py3-none-any.whl", hash = "sha256:bd789c98a94d7c5ea08c20d11dcf68c9cd1ad479b134727d80a930b84387070b", size = 36133501, upload-time = "2026-07-28T06:02:18.239Z" }, + { url = "https://files.pythonhosted.org/packages/4d/ca/8681d636307f13a7af201a8544bb903f14de2359080612db1c748ee3a5f5/ogx-1.0.2-py3-none-any.whl", hash = "sha256:6904b34c0c8300dab84e581bf7fbb5b612eac1759c5c389f7a5f31db9a6ca522", size = 696448, upload-time = "2026-05-13T22:53:50.461Z" }, +] + +[[package]] +name = "ogx-api" +version = "1.0.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "fastapi" }, + { name = "jsonschema" }, + { name = "openai" }, + { name = "opentelemetry-exporter-otlp-proto-http" }, + { name = "opentelemetry-sdk" }, + { name = "pydantic" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/c0/4e/9e8919d50328f182ea5066751c80daaaab745da70fb3b1b8b13c6a72f3cb/ogx_api-1.0.2.tar.gz", hash = "sha256:33a5ef09761eea649415c0b581395f0638e98f2a45ea03461594fbdbfed439e4", size = 145638, upload-time = "2026-05-13T22:53:01.679Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/de/99/e25be35713869b730d4a1d47a3d3b2ae5705fb044279f95c1c83b773cdef/ogx_api-1.0.2-py3-none-any.whl", hash = "sha256:e88817c97e00b1e6ddd604d87dfb0c1c32f6c4ab4a9539fed9063f7339ef672a", size = 145201, upload-time = "2026-05-13T22:52:59.773Z" }, +] + +[[package]] +name = "ogx-client" +version = "1.0.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "click" }, + { name = "distro" }, + { name = "fire" }, + { name = "httpx" }, + { name = "pandas" }, + { name = "prompt-toolkit" }, + { name = "pyaml" }, + { name = "pydantic" }, + { name = "requests" }, + { name = "rich" }, + { name = "sniffio" }, + { name = "termcolor" }, + { name = "tqdm" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/91/74/3edc7a8f396137741d9b29a728bcf45afd4a8266b40f1080ec0118c52065/ogx_client-1.0.2.tar.gz", hash = "sha256:4f5d0a5285fdbfcd3c260c291d70d325cf44b743a67c51e65d8a0c5f5cd76177", size = 435217, upload-time = "2026-05-13T22:52:04.778Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b0/fc/4bf0216fd6e13943837f30164937d6db78dd8743424984ce22c763491c22/ogx_client-1.0.2-py3-none-any.whl", hash = "sha256:167ad9f36ec632cd579bb6451dc8137d90d755d78cd7073d8fc9dcc9bf622fa5", size = 309278, upload-time = "2026-05-13T22:52:03.149Z" }, ] [[package]] @@ -2746,34 +2672,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/23/3f/ab8d29df207ce5f470a07fa96ebb48af4e95b7fab7e7635311b9a32f2fab/opentelemetry_util_http-0.65b0-py3-none-any.whl", hash = "sha256:7553b606f963097cb190536dc30556cce85090692e471a422fff30ca29b04348", size = 8245, upload-time = "2026-07-16T15:25:46.482Z" }, ] -[[package]] -name = "oracledb" -version = "4.0.2" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "cryptography" }, - { name = "typing-extensions" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/87/ae/4576e5df7b8eadec51bb7d981a2bcc0d8387d7d6998a51d146c0886a523e/oracledb-4.0.2.tar.gz", hash = "sha256:0a380ab72853487ea2764c5df772f35026b4219868fc3eba68e193c9aea230ac", size = 881658, upload-time = "2026-07-14T17:21:28.876Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/a5/89/2568c2d32afb3c0a66cef2e74dac7769f69e117a14402301b106970475dd/oracledb-4.0.2-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:cbe3a75b463a334da9b3d76026062bcfb24ec563b7df291f700c9f7eefcbeefe", size = 4366002, upload-time = "2026-07-14T17:21:58.951Z" }, - { url = "https://files.pythonhosted.org/packages/9c/c8/1825d240aa68b255eb78867c62964d2a69c6ae6197ab3543fe220539bf69/oracledb-4.0.2-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9d11723018b6aeae035f4e1feae4150b7cca1161023c2d68d957d7310fe0dd56", size = 2286174, upload-time = "2026-07-14T17:22:00.649Z" }, - { url = "https://files.pythonhosted.org/packages/10/31/5b9d6fa28942ff38e2d76dcef3184e7569d4b0006d0425f92a3d4a764dc4/oracledb-4.0.2-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:579f2c568433523a990cde5bea73c980d144754dc54d3ab2cd37efd670dc31d6", size = 2490198, upload-time = "2026-07-14T17:22:02.322Z" }, - { url = "https://files.pythonhosted.org/packages/28/4e/77b21ec50c786270a7471abd6f48ec3c2e629ce304ec792eeeb6f03ba4d0/oracledb-4.0.2-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:b82d3d83bdc6246e53f8ad8b598d2508e9f5d983e352ce2e8a71a020200ba2d6", size = 2330058, upload-time = "2026-07-14T17:23:04.213Z" }, - { url = "https://files.pythonhosted.org/packages/24/83/834e07805b8b3aaab7e87b818ef44ab3cb5394a73149a580ecf73cd92d4d/oracledb-4.0.2-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:b4f2b51837248d4acfb323a64c48d89b62467b0f73d1211c99db9b80df0cdfbc", size = 2510815, upload-time = "2026-07-14T17:23:05.838Z" }, - { url = "https://files.pythonhosted.org/packages/49/10/f05a3348a50ecab07b86317424751a4f63891e31c228a6d95f92b9538afd/oracledb-4.0.2-cp312-cp312-win32.whl", hash = "sha256:b5d203426f0d191842b4cfd87c223163ccc57f7e7875ec7ed073f163de2229af", size = 1488879, upload-time = "2026-07-14T17:23:07.53Z" }, - { url = "https://files.pythonhosted.org/packages/85/e2/99fe3fa29466533df10fdfed90faee133aa1e7147b9edfc13b839e06f738/oracledb-4.0.2-cp312-cp312-win_amd64.whl", hash = "sha256:5fe6e07ed29f84a6656e0580b4e662109d3bb90fc13d8f98ae3a56a67c75b80e", size = 1865771, upload-time = "2026-07-14T17:23:09.374Z" }, - { url = "https://files.pythonhosted.org/packages/87/cb/009980df826442900419a7bc317e35dc85a031bc2d3806256d469de7d670/oracledb-4.0.2-cp312-cp312-win_arm64.whl", hash = "sha256:a1f46b01a089e0e0dd44c4ac4c5c79903760976346c74a698955c1fb76d68bf5", size = 1519969, upload-time = "2026-07-14T17:23:11.531Z" }, - { url = "https://files.pythonhosted.org/packages/2b/cb/e9ad24c2fa20ff977ca52b5a68a2010a054f724d420681be75e7e3d92b57/oracledb-4.0.2-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:af4121112f0b68e8d61ce46a086c6aec8256fa2ca32d2e6467195a6c575c207b", size = 4352202, upload-time = "2026-07-14T17:23:13.202Z" }, - { url = "https://files.pythonhosted.org/packages/f6/ad/583d85906b5c4b344600be2bf30077f3a6c2bf04ee01b4fee9856d1da269/oracledb-4.0.2-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a814307ca8a6bf0649d8bf0a650a72f23030d9af052a2d4a481b94b6ab50d5a4", size = 2274847, upload-time = "2026-07-14T17:23:14.967Z" }, - { url = "https://files.pythonhosted.org/packages/cf/53/badcc7e29ba9e9f2158f2e18214197d63c53d99c4720d436d8fc0b6b175f/oracledb-4.0.2-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1b89c02670bafc9b1fcc0d15175f097bdc1968bc36ee4d66aa63aee09bfbbe30", size = 2474482, upload-time = "2026-07-14T17:23:16.605Z" }, - { url = "https://files.pythonhosted.org/packages/b9/76/0bc64bdfc7ca794f4576c757da89617922be402f9fcf434bf5de7e44235a/oracledb-4.0.2-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:db0d76199b6dffd721bda48dad4eeea6603942ce2e246f3b9301d85af6bda0ee", size = 2327237, upload-time = "2026-07-14T17:23:18.602Z" }, - { url = "https://files.pythonhosted.org/packages/15/f5/2fc4e30b24a60e40fa6419dbadd8f869878a3ac49d28b7d16bca59dd2519/oracledb-4.0.2-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:f5ed684ab419603d53981e0f32b51da673e21a42d7dce1bf3f215a9301c7c108", size = 2504800, upload-time = "2026-07-14T17:23:20.025Z" }, - { url = "https://files.pythonhosted.org/packages/2e/03/0a15a43a88addfb7b4f8d8343f0c47df68173510a27e2196b6b7347efb06/oracledb-4.0.2-cp313-cp313-win32.whl", hash = "sha256:a45a13e33509db5bd3622a05cb2e4a6cb41d56d1bfd2a32caeacd26f4ce8b988", size = 1493052, upload-time = "2026-07-14T17:23:21.603Z" }, - { url = "https://files.pythonhosted.org/packages/bb/4a/9895cb5a1fffa68f2d9fb1cfb43de88628c049d313734de7ecef17099460/oracledb-4.0.2-cp313-cp313-win_amd64.whl", hash = "sha256:6444be4991f33754cd98f624cecef3eb98db3e87f8ac14e9b65bd3591c9dd252", size = 1864126, upload-time = "2026-07-14T17:23:22.985Z" }, - { url = "https://files.pythonhosted.org/packages/0b/14/5c7c44a8f441783d461b9657994e03f8798e28bebe67395951c53c86c270/oracledb-4.0.2-cp313-cp313-win_arm64.whl", hash = "sha256:087fdf5bc36b03dc3c55f7abc944a163ba8edc9c802bc8baf91698a52041ee13", size = 1519906, upload-time = "2026-07-14T17:23:24.338Z" }, -] - [[package]] name = "packaging" version = "26.2" @@ -3453,19 +3351,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ab/da/acb2e7d4dbd2dfb792d38c0d850481f29ad7049b356d23f56c687d35203b/pylint-4.0.6-py3-none-any.whl", hash = "sha256:d11a0e1fdb7b1cd46ec5d6fc78fee8b95f28695b2d6140e5809925f61e32ea54", size = 538389, upload-time = "2026-06-14T14:43:24.873Z" }, ] -[[package]] -name = "pyopenssl" -version = "26.3.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "cryptography" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/74/b7/da07bae88f5a9506b4def6f2f4903cf4c3b8831e560dba8fa18ca08f758f/pyopenssl-26.3.0.tar.gz", hash = "sha256:589de7fae1c9ea670d18422ed00fc04da787bbde8e1454aea872aa57b49ad341", size = 182024, upload-time = "2026-06-12T20:28:07.458Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/54/18/1dd71c9b43192ab83f1d531ad6002dc81108ac36c475f79fb7a295abe2f4/pyopenssl-26.3.0-py3-none-any.whl", hash = "sha256:46367f8f66b92271e6d218da9c87607e1ef5a0bc5c8dea5bb3db82f395c385a3", size = 56008, upload-time = "2026-06-12T20:28:05.999Z" }, -] - [[package]] name = "pypdf" version = "6.14.2" @@ -3601,14 +3486,14 @@ wheels = [ [[package]] name = "pythainlp" -version = "5.3.4" +version = "5.3.5" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "tzdata", marker = "sys_platform == 'win32'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/56/47/2364543c559eacdc63633a3f29d7467a36a33089f7a8b226a247dd12de7e/pythainlp-5.3.4.tar.gz", hash = "sha256:e66fd76fb5931834fd4e32ed54337ec62350d7654f187850e4dd4f915e9f624f", size = 19306201, upload-time = "2026-04-02T18:42:46.044Z" } +sdist = { url = "https://files.pythonhosted.org/packages/8b/5a/f893095176843b998b93abac1df5976ce2e738366fb84784dfd7ee56c283/pythainlp-5.3.5.tar.gz", hash = "sha256:3be53b97e44fdfc55669705b31a2fe96546146d0c1d90f18999c9b04c4e50c83", size = 19306143, upload-time = "2026-07-29T18:47:29.05Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0c/bd/bedc1e416eb0fa2cb5eaa990207ae67dfaf02d4986f22d5ec372c0b97ae3/pythainlp-5.3.4-py3-none-any.whl", hash = "sha256:76744e51e27c895630bafd74f53a1f0aa8782cef2f7f02eebd6427fe8ce8d84d", size = 19849568, upload-time = "2026-04-02T18:42:42.384Z" }, + { url = "https://files.pythonhosted.org/packages/0c/21/805d31f57a6b1d93b12a3007d04621789aaec2f11931098fc0e56f1fff3f/pythainlp-5.3.5-py3-none-any.whl", hash = "sha256:147a7a77c5c6d5b387b827ed3b00ee23c3665d242e9d021565cae4c3bca7b2c2", size = 19849547, upload-time = "2026-07-29T18:47:25.983Z" }, ] [[package]] @@ -3660,15 +3545,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/c6/78/397db326746f0a342855b81216ae1f0a32965deccfd7c830a2dbc66d2483/pytokens-0.4.1-py3-none-any.whl", hash = "sha256:26cef14744a8385f35d0e095dc8b3a7583f6c953c2e3d269c7f82484bf5ad2de", size = 13729, upload-time = "2026-01-30T01:03:45.029Z" }, ] -[[package]] -name = "pytz" -version = "2026.3.post1" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/fb/48/fb042503b6ca6cd271261dc559fd6432f7d8c713153e9ec5c591af4dfc1c/pytz-2026.3.post1.tar.gz", hash = "sha256:2211d3fcf9a797d3405cac96ac7f61d80e6a644f72a3309607282fe8a2010c5d", size = 319745, upload-time = "2026-07-25T15:12:07.385Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/0f/7b/39c34ca613b0b198cb866466651b26b045e2009864c5183c979a3b83f383/pytz-2026.3.post1-py2.py3-none-any.whl", hash = "sha256:dd95840dd199baea12d9cc096a1d452caa6596a1c1e4b5f3dbd1541855d5e815", size = 508283, upload-time = "2026-07-25T15:12:05.782Z" }, -] - [[package]] name = "pywin32" version = "312" @@ -4182,6 +4058,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/62/8d/008761f6e1000600e5303db30d05724bdcf3d2d186cbb59fac79b52e39ed/stevedore-5.9.0-py3-none-any.whl", hash = "sha256:e520945d4c257700eddc1eb1d79df04b2ea578eef185e0e3fa5b442fc848d3f7", size = 54463, upload-time = "2026-07-02T11:38:07.43Z" }, ] +[[package]] +name = "structlog" +version = "26.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/5e/89/b4a0bcfdf4f71a3dea31379f095929613d7e4528a0996bca6aa964cd0dca/structlog-26.1.0.tar.gz", hash = "sha256:f63a716cbd1b1291cf7661de7794b455acfa4c43c5bcf1630e6ad5ddc1adb3b7", size = 1459881, upload-time = "2026-06-06T07:33:39.348Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a9/18/489c97b834dfff9cf2fc2507cede4bcd4b11e67f84bc462acd1992496f86/structlog-26.1.0-py3-none-any.whl", hash = "sha256:e081a26d6c373e6d201eca24eede26d8ffab07f88f477822e679183428d3d91e", size = 73764, upload-time = "2026-06-06T07:33:38.046Z" }, +] + [[package]] name = "sympy" version = "1.14.0" @@ -4348,23 +4233,6 @@ wheels = [ { url = "https://download-r2.pytorch.org/whl/cpu/torch-2.11.0%2Bcpu-cp313-cp313t-win_amd64.whl", hash = "sha256:62ec1f1694c185f601eab74eb7fc0e8e10c64c06ae82f13c3592774c231c4877", upload-time = "2026-04-28T00:07:47Z" }, ] -[[package]] -name = "tornado" -version = "6.5.7" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/64/24/95ec527ad67b76d59299e5465b3935d05e4294b7e0290a3924b7487df30b/tornado-6.5.7.tar.gz", hash = "sha256:66c513a76cda70d53907bc27cf1447557699c2e95aa48ba27a442ff61c3ddfc2", size = 519252, upload-time = "2026-06-08T17:34:51.232Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/02/dc/c7043cab6fed8ae159fc1923ce829ada35c4dbd797d408a43858ffaf9639/tornado-6.5.7-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:148b2eb15c2c765a50796172c1e499649b35f30d2e3c3d3e15913cfa56bfb163", size = 448543, upload-time = "2026-06-08T17:34:38.052Z" }, - { url = "https://files.pythonhosted.org/packages/92/4f/090b1431e5a43df696feceffc268c5383cc079ecb5f08ce58f917109aafe/tornado-6.5.7-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:9da38de27f1da3b78a966f0dae12b5a1ea9afe72ca805d84ff06508272ddf100", size = 446707, upload-time = "2026-06-08T17:34:39.594Z" }, - { url = "https://files.pythonhosted.org/packages/37/d8/ef374952fd5da67d4463122c2b8e5a96536ec10b4b339254c6dcde81d01c/tornado-6.5.7-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:8d759e71906ee783f8867b93bf26a265743da4c1e2f4a018464c1ba019862972", size = 449774, upload-time = "2026-06-08T17:34:41.204Z" }, - { url = "https://files.pythonhosted.org/packages/35/37/d434c73f4c6e014b745b9b37085f34f40c022f007efff3d7fe65991899f3/tornado-6.5.7-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8a46347a18f23fb92b396beebe0fb78f61dda0cc302445202c16203d8a18848b", size = 450745, upload-time = "2026-06-08T17:34:42.531Z" }, - { url = "https://files.pythonhosted.org/packages/b6/2b/56b9aff361d7f1ab728a805ec7d7ea835f8807afa9f5cc690ea0e630efb9/tornado-6.5.7-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:7778b30bef919231265e91c69963ce0f49a1e9c07ac900bbe75b19ce2575ba92", size = 450578, upload-time = "2026-06-08T17:34:43.787Z" }, - { url = "https://files.pythonhosted.org/packages/02/30/a7444fb23aa76860a14198fab96ac79f1866b0a6e19e26c4381b0938e50f/tornado-6.5.7-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:e726f0c75da7726eec023aa62751ff8878bd2737e34fbdd33b1ae5897d2200f5", size = 449985, upload-time = "2026-06-08T17:34:45.326Z" }, - { url = "https://files.pythonhosted.org/packages/5c/42/5f0e56c01e8d9d36f4e23f367b85ae6cae0c1ecddd5e6977d8388ad27488/tornado-6.5.7-cp39-abi3-win32.whl", hash = "sha256:f8de3bf12d3efdd0cbe7c8887868198f8a91415e3f29fcf258d9b8eb7b1d9ae4", size = 451047, upload-time = "2026-06-08T17:34:46.784Z" }, - { url = "https://files.pythonhosted.org/packages/c9/a4/b393076ffb21b469eec5b328a0534cf03a3b90bfc6b1f09507cdd075d938/tornado-6.5.7-cp39-abi3-win_amd64.whl", hash = "sha256:de942f843533a039ef9fa3d9c88c7cd8a7c94553fb5ad0154270989b3d99a2c4", size = 451485, upload-time = "2026-06-08T17:34:48.248Z" }, - { url = "https://files.pythonhosted.org/packages/71/2e/7b1c769803121b809112cf9a00681c472eae1d80e32d7ec0e0bd61d0d0e1/tornado-6.5.7-cp39-abi3-win_arm64.whl", hash = "sha256:ff934fce95643af5f11efdae618eaa73d469dc588641e5c8d19295a0c65c4796", size = 450506, upload-time = "2026-06-08T17:34:49.702Z" }, -] - [[package]] name = "tqdm" version = "4.70.0" @@ -4591,15 +4459,15 @@ wheels = [ [[package]] name = "uvicorn" -version = "0.51.0" +version = "0.52.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "click" }, { name = "h11" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/a2/65/b7c6c443ccc58678c91e1e973bbe2a878591538655d6e1d47f24ba1c51f3/uvicorn-0.51.0.tar.gz", hash = "sha256:f6f4b69b657c312f516dd2d268ab9ae6f254b11e4bac504f37b2ab58b24dd0b0", size = 94412, upload-time = "2026-07-08T10:59:05.962Z" } +sdist = { url = "https://files.pythonhosted.org/packages/05/c8/2d307868453a4bca6e64fa3581d122ae0748a0869c53f159339def179c7c/uvicorn-0.52.0.tar.gz", hash = "sha256:ca8876ad6c1983f394157c168b39d52f6dd56dabf5602fa0982751cffc2293ae", size = 97504, upload-time = "2026-07-29T08:45:34.065Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/45/ec/dbb7e5a6b91f86bfb9eb7d2988a2730907b6a729875b949c7f022e8b88fa/uvicorn-0.51.0-py3-none-any.whl", hash = "sha256:5d38af6cd620f2ae3849fb44fd4879e0890aa1febe8d47eb355fb45d93fe6a5b", size = 73219, upload-time = "2026-07-08T10:59:04.44Z" }, + { url = "https://files.pythonhosted.org/packages/39/e6/b5c0630ace9757232aec07112be8146b812787db52141ff9d50674aa7634/uvicorn-0.52.0-py3-none-any.whl", hash = "sha256:3d887809810b89ed33501bcf0a9aba469b06ecd608158efce04bd6b48d8c9b08", size = 79058, upload-time = "2026-07-29T08:45:32.492Z" }, ] [[package]] @@ -4852,3 +4720,45 @@ sdist = { url = "https://files.pythonhosted.org/packages/b9/d8/eab98a517c14134c0 wheels = [ { url = "https://files.pythonhosted.org/packages/3a/13/547360d81e6d88d58492968ffda9f9542854f11310ee556fef14260cc886/zipp-4.1.0-py3-none-any.whl", hash = "sha256:25ad4e16390cd314347dd8f1de67a2ac538ae658ed4ab9db16029c07c188e97f", size = 10238, upload-time = "2026-05-18T20:08:57.045Z" }, ] + +[[package]] +name = "zstandard" +version = "0.25.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/fd/aa/3e0508d5a5dd96529cdc5a97011299056e14c6505b678fd58938792794b1/zstandard-0.25.0.tar.gz", hash = "sha256:7713e1179d162cf5c7906da876ec2ccb9c3a9dcbdffef0cc7f70c3667a205f0b", size = 711513, upload-time = "2025-09-14T22:15:54.002Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/82/fc/f26eb6ef91ae723a03e16eddb198abcfce2bc5a42e224d44cc8b6765e57e/zstandard-0.25.0-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:7b3c3a3ab9daa3eed242d6ecceead93aebbb8f5f84318d82cee643e019c4b73b", size = 795738, upload-time = "2025-09-14T22:16:56.237Z" }, + { url = "https://files.pythonhosted.org/packages/aa/1c/d920d64b22f8dd028a8b90e2d756e431a5d86194caa78e3819c7bf53b4b3/zstandard-0.25.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:913cbd31a400febff93b564a23e17c3ed2d56c064006f54efec210d586171c00", size = 640436, upload-time = "2025-09-14T22:16:57.774Z" }, + { url = "https://files.pythonhosted.org/packages/53/6c/288c3f0bd9fcfe9ca41e2c2fbfd17b2097f6af57b62a81161941f09afa76/zstandard-0.25.0-cp312-cp312-manylinux2010_i686.manylinux2014_i686.manylinux_2_12_i686.manylinux_2_17_i686.whl", hash = "sha256:011d388c76b11a0c165374ce660ce2c8efa8e5d87f34996aa80f9c0816698b64", size = 5343019, upload-time = "2025-09-14T22:16:59.302Z" }, + { url = "https://files.pythonhosted.org/packages/1e/15/efef5a2f204a64bdb5571e6161d49f7ef0fffdbca953a615efbec045f60f/zstandard-0.25.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:6dffecc361d079bb48d7caef5d673c88c8988d3d33fb74ab95b7ee6da42652ea", size = 5063012, upload-time = "2025-09-14T22:17:01.156Z" }, + { url = "https://files.pythonhosted.org/packages/b7/37/a6ce629ffdb43959e92e87ebdaeebb5ac81c944b6a75c9c47e300f85abdf/zstandard-0.25.0-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:7149623bba7fdf7e7f24312953bcf73cae103db8cae49f8154dd1eadc8a29ecb", size = 5394148, upload-time = "2025-09-14T22:17:03.091Z" }, + { url = "https://files.pythonhosted.org/packages/e3/79/2bf870b3abeb5c070fe2d670a5a8d1057a8270f125ef7676d29ea900f496/zstandard-0.25.0-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:6a573a35693e03cf1d67799fd01b50ff578515a8aeadd4595d2a7fa9f3ec002a", size = 5451652, upload-time = "2025-09-14T22:17:04.979Z" }, + { url = "https://files.pythonhosted.org/packages/53/60/7be26e610767316c028a2cbedb9a3beabdbe33e2182c373f71a1c0b88f36/zstandard-0.25.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:5a56ba0db2d244117ed744dfa8f6f5b366e14148e00de44723413b2f3938a902", size = 5546993, upload-time = "2025-09-14T22:17:06.781Z" }, + { url = "https://files.pythonhosted.org/packages/85/c7/3483ad9ff0662623f3648479b0380d2de5510abf00990468c286c6b04017/zstandard-0.25.0-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:10ef2a79ab8e2974e2075fb984e5b9806c64134810fac21576f0668e7ea19f8f", size = 5046806, upload-time = "2025-09-14T22:17:08.415Z" }, + { url = "https://files.pythonhosted.org/packages/08/b3/206883dd25b8d1591a1caa44b54c2aad84badccf2f1de9e2d60a446f9a25/zstandard-0.25.0-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:aaf21ba8fb76d102b696781bddaa0954b782536446083ae3fdaa6f16b25a1c4b", size = 5576659, upload-time = "2025-09-14T22:17:10.164Z" }, + { url = "https://files.pythonhosted.org/packages/9d/31/76c0779101453e6c117b0ff22565865c54f48f8bd807df2b00c2c404b8e0/zstandard-0.25.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:1869da9571d5e94a85a5e8d57e4e8807b175c9e4a6294e3b66fa4efb074d90f6", size = 4953933, upload-time = "2025-09-14T22:17:11.857Z" }, + { url = "https://files.pythonhosted.org/packages/18/e1/97680c664a1bf9a247a280a053d98e251424af51f1b196c6d52f117c9720/zstandard-0.25.0-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:809c5bcb2c67cd0ed81e9229d227d4ca28f82d0f778fc5fea624a9def3963f91", size = 5268008, upload-time = "2025-09-14T22:17:13.627Z" }, + { url = "https://files.pythonhosted.org/packages/1e/73/316e4010de585ac798e154e88fd81bb16afc5c5cb1a72eeb16dd37e8024a/zstandard-0.25.0-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:f27662e4f7dbf9f9c12391cb37b4c4c3cb90ffbd3b1fb9284dadbbb8935fa708", size = 5433517, upload-time = "2025-09-14T22:17:16.103Z" }, + { url = "https://files.pythonhosted.org/packages/5b/60/dd0f8cfa8129c5a0ce3ea6b7f70be5b33d2618013a161e1ff26c2b39787c/zstandard-0.25.0-cp312-cp312-musllinux_1_2_s390x.whl", hash = "sha256:99c0c846e6e61718715a3c9437ccc625de26593fea60189567f0118dc9db7512", size = 5814292, upload-time = "2025-09-14T22:17:17.827Z" }, + { url = "https://files.pythonhosted.org/packages/fc/5f/75aafd4b9d11b5407b641b8e41a57864097663699f23e9ad4dbb91dc6bfe/zstandard-0.25.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:474d2596a2dbc241a556e965fb76002c1ce655445e4e3bf38e5477d413165ffa", size = 5360237, upload-time = "2025-09-14T22:17:19.954Z" }, + { url = "https://files.pythonhosted.org/packages/ff/8d/0309daffea4fcac7981021dbf21cdb2e3427a9e76bafbcdbdf5392ff99a4/zstandard-0.25.0-cp312-cp312-win32.whl", hash = "sha256:23ebc8f17a03133b4426bcc04aabd68f8236eb78c3760f12783385171b0fd8bd", size = 436922, upload-time = "2025-09-14T22:17:24.398Z" }, + { url = "https://files.pythonhosted.org/packages/79/3b/fa54d9015f945330510cb5d0b0501e8253c127cca7ebe8ba46a965df18c5/zstandard-0.25.0-cp312-cp312-win_amd64.whl", hash = "sha256:ffef5a74088f1e09947aecf91011136665152e0b4b359c42be3373897fb39b01", size = 506276, upload-time = "2025-09-14T22:17:21.429Z" }, + { url = "https://files.pythonhosted.org/packages/ea/6b/8b51697e5319b1f9ac71087b0af9a40d8a6288ff8025c36486e0c12abcc4/zstandard-0.25.0-cp312-cp312-win_arm64.whl", hash = "sha256:181eb40e0b6a29b3cd2849f825e0fa34397f649170673d385f3598ae17cca2e9", size = 462679, upload-time = "2025-09-14T22:17:23.147Z" }, + { url = "https://files.pythonhosted.org/packages/35/0b/8df9c4ad06af91d39e94fa96cc010a24ac4ef1378d3efab9223cc8593d40/zstandard-0.25.0-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:ec996f12524f88e151c339688c3897194821d7f03081ab35d31d1e12ec975e94", size = 795735, upload-time = "2025-09-14T22:17:26.042Z" }, + { url = "https://files.pythonhosted.org/packages/3f/06/9ae96a3e5dcfd119377ba33d4c42a7d89da1efabd5cb3e366b156c45ff4d/zstandard-0.25.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:a1a4ae2dec3993a32247995bdfe367fc3266da832d82f8438c8570f989753de1", size = 640440, upload-time = "2025-09-14T22:17:27.366Z" }, + { url = "https://files.pythonhosted.org/packages/d9/14/933d27204c2bd404229c69f445862454dcc101cd69ef8c6068f15aaec12c/zstandard-0.25.0-cp313-cp313-manylinux2010_i686.manylinux2014_i686.manylinux_2_12_i686.manylinux_2_17_i686.whl", hash = "sha256:e96594a5537722fdfb79951672a2a63aec5ebfb823e7560586f7484819f2a08f", size = 5343070, upload-time = "2025-09-14T22:17:28.896Z" }, + { url = "https://files.pythonhosted.org/packages/6d/db/ddb11011826ed7db9d0e485d13df79b58586bfdec56e5c84a928a9a78c1c/zstandard-0.25.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:bfc4e20784722098822e3eee42b8e576b379ed72cca4a7cb856ae733e62192ea", size = 5063001, upload-time = "2025-09-14T22:17:31.044Z" }, + { url = "https://files.pythonhosted.org/packages/db/00/87466ea3f99599d02a5238498b87bf84a6348290c19571051839ca943777/zstandard-0.25.0-cp313-cp313-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:457ed498fc58cdc12fc48f7950e02740d4f7ae9493dd4ab2168a47c93c31298e", size = 5394120, upload-time = "2025-09-14T22:17:32.711Z" }, + { url = "https://files.pythonhosted.org/packages/2b/95/fc5531d9c618a679a20ff6c29e2b3ef1d1f4ad66c5e161ae6ff847d102a9/zstandard-0.25.0-cp313-cp313-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:fd7a5004eb1980d3cefe26b2685bcb0b17989901a70a1040d1ac86f1d898c551", size = 5451230, upload-time = "2025-09-14T22:17:34.41Z" }, + { url = "https://files.pythonhosted.org/packages/63/4b/e3678b4e776db00f9f7b2fe58e547e8928ef32727d7a1ff01dea010f3f13/zstandard-0.25.0-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:8e735494da3db08694d26480f1493ad2cf86e99bdd53e8e9771b2752a5c0246a", size = 5547173, upload-time = "2025-09-14T22:17:36.084Z" }, + { url = "https://files.pythonhosted.org/packages/4e/d5/ba05ed95c6b8ec30bd468dfeab20589f2cf709b5c940483e31d991f2ca58/zstandard-0.25.0-cp313-cp313-musllinux_1_1_aarch64.whl", hash = "sha256:3a39c94ad7866160a4a46d772e43311a743c316942037671beb264e395bdd611", size = 5046736, upload-time = "2025-09-14T22:17:37.891Z" }, + { url = "https://files.pythonhosted.org/packages/50/d5/870aa06b3a76c73eced65c044b92286a3c4e00554005ff51962deef28e28/zstandard-0.25.0-cp313-cp313-musllinux_1_1_x86_64.whl", hash = "sha256:172de1f06947577d3a3005416977cce6168f2261284c02080e7ad0185faeced3", size = 5576368, upload-time = "2025-09-14T22:17:40.206Z" }, + { url = "https://files.pythonhosted.org/packages/5d/35/398dc2ffc89d304d59bc12f0fdd931b4ce455bddf7038a0a67733a25f550/zstandard-0.25.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:3c83b0188c852a47cd13ef3bf9209fb0a77fa5374958b8c53aaa699398c6bd7b", size = 4954022, upload-time = "2025-09-14T22:17:41.879Z" }, + { url = "https://files.pythonhosted.org/packages/9a/5c/36ba1e5507d56d2213202ec2b05e8541734af5f2ce378c5d1ceaf4d88dc4/zstandard-0.25.0-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:1673b7199bbe763365b81a4f3252b8e80f44c9e323fc42940dc8843bfeaf9851", size = 5267889, upload-time = "2025-09-14T22:17:43.577Z" }, + { url = "https://files.pythonhosted.org/packages/70/e8/2ec6b6fb7358b2ec0113ae202647ca7c0e9d15b61c005ae5225ad0995df5/zstandard-0.25.0-cp313-cp313-musllinux_1_2_ppc64le.whl", hash = "sha256:0be7622c37c183406f3dbf0cba104118eb16a4ea7359eeb5752f0794882fc250", size = 5433952, upload-time = "2025-09-14T22:17:45.271Z" }, + { url = "https://files.pythonhosted.org/packages/7b/01/b5f4d4dbc59ef193e870495c6f1275f5b2928e01ff5a81fecb22a06e22fb/zstandard-0.25.0-cp313-cp313-musllinux_1_2_s390x.whl", hash = "sha256:5f5e4c2a23ca271c218ac025bd7d635597048b366d6f31f420aaeb715239fc98", size = 5814054, upload-time = "2025-09-14T22:17:47.08Z" }, + { url = "https://files.pythonhosted.org/packages/b2/e5/fbd822d5c6f427cf158316d012c5a12f233473c2f9c5fe5ab1ae5d21f3d8/zstandard-0.25.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:4f187a0bb61b35119d1926aee039524d1f93aaf38a9916b8c4b78ac8514a0aaf", size = 5360113, upload-time = "2025-09-14T22:17:48.893Z" }, + { url = "https://files.pythonhosted.org/packages/8e/e0/69a553d2047f9a2c7347caa225bb3a63b6d7704ad74610cb7823baa08ed7/zstandard-0.25.0-cp313-cp313-win32.whl", hash = "sha256:7030defa83eef3e51ff26f0b7bfb229f0204b66fe18e04359ce3474ac33cbc09", size = 436936, upload-time = "2025-09-14T22:17:52.658Z" }, + { url = "https://files.pythonhosted.org/packages/d9/82/b9c06c870f3bd8767c201f1edbdf9e8dc34be5b0fbc5682c4f80fe948475/zstandard-0.25.0-cp313-cp313-win_amd64.whl", hash = "sha256:1f830a0dac88719af0ae43b8b2d6aef487d437036468ef3c2ea59c51f9d55fd5", size = 506232, upload-time = "2025-09-14T22:17:50.402Z" }, + { url = "https://files.pythonhosted.org/packages/d4/57/60c3c01243bb81d381c9916e2a6d9e149ab8627c0c7d7abb2d73384b3c0c/zstandard-0.25.0-cp313-cp313-win_arm64.whl", hash = "sha256:85304a43f4d513f5464ceb938aa02c1e78c2943b29f44a750b48b25ac999a049", size = 462671, upload-time = "2025-09-14T22:17:51.533Z" }, +]