Skip to content

Commit 3171422

Browse files
committed
Restructure pydantic_ai utils into LlamaStack model and provider classes
Move llama_stack_provider_from_client and _model_settings_from_responses_params out of utils/pydantic_ai.py into their owning classes as static factory methods: - LlamaStackProvider.from_llama_stack_client() in _provider.py - LlamaStackResponsesModel.from_llama_stack_client() in _model.py - _model_settings_from_responses_params and _LLS_RESPONSES_EXTRA_FIELDS into _model.py Update callers (QuestionValidity, build_agent) to use the new factory methods and relocate corresponding tests to test_model.py and test_provider.py.
1 parent 513df21 commit 3171422

8 files changed

Lines changed: 1031 additions & 333 deletions

File tree

src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py

Lines changed: 7 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,6 @@
2727
QuestionValidityConfig,
2828
)
2929
from pydantic_ai_lightspeed.llamastack import LlamaStackResponsesModel
30-
from utils.pydantic_ai import llama_stack_provider_from_client
3130

3231
logger = get_logger(__name__)
3332

@@ -55,21 +54,6 @@ def _extract_message_str_from_user_content(user_content: Sequence[UserContent])
5554
return "\n".join(str_arr)
5655

5756

58-
def _create_model_from_llama_stack_client(model_id: str) -> LlamaStackResponsesModel:
59-
"""Create a LlamaStackResponsesModel from the shared Llama Stack client.
60-
61-
Parameters:
62-
model_id: The model identifier to use for the responses model.
63-
64-
Returns:
65-
A configured LlamaStackResponsesModel instance.
66-
"""
67-
client = AsyncLlamaStackClientHolder().get_client()
68-
provider = llama_stack_provider_from_client(client)
69-
settings = OpenAIResponsesModelSettings(openai_store=False)
70-
return LlamaStackResponsesModel(model_id, provider=provider, settings=settings)
71-
72-
7357
@dataclass
7458
class QuestionValidity(AbstractCapability[None]):
7559
"""Block or modify user input based on a guardrail check.
@@ -91,7 +75,13 @@ class QuestionValidity(AbstractCapability[None]):
9175

9276
def __post_init__(self) -> None:
9377
"""Initialize the model instance from the configured model ID."""
94-
self._model = _create_model_from_llama_stack_client(self.config.model_id)
78+
llama_stack_client = AsyncLlamaStackClientHolder().get_client()
79+
80+
self._model = LlamaStackResponsesModel.from_llama_stack_client(
81+
self.config.model_id,
82+
llama_stack_client,
83+
model_settings=OpenAIResponsesModelSettings(openai_store=False),
84+
)
9585

9686
def _build_prompt(self, message: str | Sequence[UserContent] | None) -> str:
9787
"""Build the classification prompt from the user message.

src/pydantic_ai_lightspeed/llamastack/_model.py

Lines changed: 97 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,8 +19,10 @@
1919
from collections import defaultdict
2020
from collections.abc import AsyncIterator
2121
from contextlib import asynccontextmanager
22-
from typing import Any, cast
22+
from typing import Any, Final, cast
2323

24+
from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient
25+
from llama_stack_client import AsyncLlamaStackClient
2426
from openai import AsyncStream
2527
from openai.types import responses
2628
from pydantic_ai import UnexpectedModelBehavior
@@ -38,12 +40,56 @@
3840
OpenAIResponsesStreamedResponse,
3941
_map_api_errors,
4042
)
43+
from pydantic_ai.profiles import ModelProfileSpec
4144
from pydantic_ai.settings import ModelSettings
4245

4346
from log import get_logger
47+
from models.common.responses.responses_api_params import ResponsesApiParams
48+
from pydantic_ai_lightspeed.llamastack._provider import LlamaStackProvider
4449

4550
logger = get_logger(__name__)
4651

52+
_LLS_RESPONSES_EXTRA_FIELDS: Final[frozenset[str]] = frozenset(
53+
{
54+
"conversation",
55+
"max_infer_iters",
56+
"tools",
57+
"tool_choice",
58+
"include",
59+
"text",
60+
"reasoning",
61+
"prompt",
62+
"metadata",
63+
"max_tool_calls",
64+
"safety_identifier",
65+
}
66+
)
67+
68+
69+
def _model_settings_from_responses_params(
70+
responses_params: ResponsesApiParams,
71+
) -> OpenAIResponsesModelSettings:
72+
"""Map ``ResponsesApiParams`` into Pydantic AI OpenAI Responses model settings."""
73+
payload = responses_params.model_dump(exclude_none=True)
74+
extra_body = {k: v for k, v in payload.items() if k in _LLS_RESPONSES_EXTRA_FIELDS}
75+
settings_dict: dict[str, Any] = {}
76+
if extra_body:
77+
settings_dict["extra_body"] = extra_body
78+
if responses_params.max_output_tokens is not None:
79+
settings_dict["max_tokens"] = responses_params.max_output_tokens
80+
if responses_params.temperature is not None:
81+
settings_dict["temperature"] = responses_params.temperature
82+
if responses_params.parallel_tool_calls is not None:
83+
settings_dict["parallel_tool_calls"] = responses_params.parallel_tool_calls
84+
if responses_params.extra_headers:
85+
settings_dict["extra_headers"] = dict(responses_params.extra_headers)
86+
settings_dict["openai_store"] = responses_params.store
87+
if responses_params.previous_response_id is not None:
88+
settings_dict["openai_previous_response_id"] = (
89+
responses_params.previous_response_id
90+
)
91+
return cast(OpenAIResponsesModelSettings, settings_dict)
92+
4793

4894
class _FilteredResponseStream:
4995
"""Wraps an OpenAI AsyncStream to reorder spurious events from Llama Stack.
@@ -307,3 +353,53 @@ async def request_stream( # pylint: disable=unused-argument
307353
else None
308354
),
309355
)
356+
357+
@staticmethod
358+
def from_llama_stack_client(
359+
model_name: str,
360+
client: AsyncLlamaStackClient | AsyncLlamaStackAsLibraryClient,
361+
*,
362+
responses_params: ResponsesApiParams | None = None,
363+
model_settings: ModelSettings | None = None,
364+
profile: ModelProfileSpec | None = None,
365+
) -> LlamaStackResponsesModel:
366+
"""Create a ``LlamaStackResponsesModel`` from a Llama Stack client.
367+
368+
Mirrors ``OpenAIResponsesModel.__init__`` parameters, but accepts a
369+
Llama Stack client instead of a provider. Exactly one of
370+
``responses_params`` or ``model_settings`` may be provided.
371+
372+
Args:
373+
model_name: The model name/ID to use.
374+
client: Llama Stack client to build the provider from.
375+
responses_params: Optional ``ResponsesApiParams``, converted to
376+
``OpenAIResponsesModelSettings`` internally. Mutually
377+
exclusive with ``model_settings``.
378+
model_settings: Optional raw ``ModelSettings`` passed through
379+
directly. Mutually exclusive with ``responses_params``.
380+
profile: Optional model profile specification.
381+
382+
Raises:
383+
ValueError: If both ``responses_params`` and ``model_settings``
384+
are provided.
385+
386+
Returns:
387+
Configured ``LlamaStackResponsesModel`` instance.
388+
"""
389+
provider = LlamaStackProvider.from_llama_stack_client(client)
390+
391+
if responses_params is not None and model_settings is not None:
392+
raise ValueError(
393+
"You can only pass either ResponsesApiParams or ModelSetting not both."
394+
)
395+
396+
_settings: OpenAIResponsesModelSettings | ModelSettings | None = None
397+
398+
if responses_params is not None:
399+
_settings = _model_settings_from_responses_params(responses_params)
400+
elif model_settings is not None:
401+
_settings = model_settings
402+
403+
return LlamaStackResponsesModel(
404+
model_name, provider=provider, profile=profile, settings=_settings
405+
)

src/pydantic_ai_lightspeed/llamastack/_provider.py

Lines changed: 32 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@
55
from typing import TYPE_CHECKING, Optional
66

77
import httpx
8+
from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient
9+
from llama_stack_client import AsyncLlamaStackClient
810
from openai import AsyncOpenAI
911
from pydantic_ai import ModelProfile
1012
from pydantic_ai.models import create_async_http_client
@@ -14,7 +16,9 @@
1416
from pydantic_ai_lightspeed.llamastack._transport import LlamaStackLibraryTransport
1517

1618
if TYPE_CHECKING:
17-
from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient
19+
from llama_stack.core.library_client import (
20+
AsyncLlamaStackAsLibraryClient, # pylint: disable=reimported
21+
)
1822

1923
DEFAULT_BASE_URL = "http://localhost:8321/v1"
2024

@@ -48,6 +52,33 @@ def model_profile(model_name: str) -> Optional[ModelProfile]:
4852
"""Return the model profile for the named model, if available."""
4953
return openai_model_profile(model_name)
5054

55+
@staticmethod
56+
def from_llama_stack_client(
57+
client: AsyncLlamaStackClient | AsyncLlamaStackAsLibraryClient,
58+
) -> LlamaStackProvider:
59+
"""Create a ``LlamaStackProvider`` from a Llama Stack client.
60+
61+
For an ``AsyncLlamaStackAsLibraryClient``, delegates to library mode.
62+
For an ``AsyncLlamaStackClient``, extracts the base URL, API key, and
63+
underlying HTTP client to create a server-mode provider.
64+
65+
Args:
66+
client: A Llama Stack client (server or library variant).
67+
68+
Returns:
69+
Configured ``LlamaStackProvider`` instance.
70+
"""
71+
if isinstance(client, AsyncLlamaStackAsLibraryClient):
72+
return LlamaStackProvider(library_client=client)
73+
api_key = client.api_key or "not-needed"
74+
base = str(client.base_url).rstrip("/")
75+
base_url = base if base.endswith("/v1") else f"{base}/v1"
76+
return LlamaStackProvider(
77+
base_url=base_url,
78+
api_key=api_key,
79+
http_client=client._client, # pylint: disable=protected-access
80+
)
81+
5182
def __init__(
5283
self,
5384
*,

src/utils/pydantic_ai.py

Lines changed: 4 additions & 67 deletions
Original file line numberDiff line numberDiff line change
@@ -2,80 +2,20 @@
22

33
from __future__ import annotations
44

5-
from typing import Any, Final, Optional, cast
5+
from typing import Optional
66

77
from llama_stack.core.library_client import AsyncLlamaStackAsLibraryClient
88
from llama_stack_client import AsyncLlamaStackClient
99
from pydantic_ai.agent import Agent
1010
from pydantic_ai.capabilities import AbstractCapability, AgentCapability
11-
from pydantic_ai.models.openai import OpenAIResponsesModelSettings
1211
from pydantic_ai_skills import SkillsCapability
1312

1413
from models.common.responses.responses_api_params import ResponsesApiParams
1514
from models.config import SkillsConfiguration
1615
from pydantic_ai_lightspeed.llamastack import (
17-
LlamaStackProvider,
1816
LlamaStackResponsesModel,
1917
)
2018

21-
_LLS_RESPONSES_EXTRA_FIELDS: Final[frozenset[str]] = frozenset(
22-
{
23-
"conversation",
24-
"max_infer_iters",
25-
"tool_choice",
26-
"include",
27-
"text",
28-
"reasoning",
29-
"prompt",
30-
"metadata",
31-
"max_tool_calls",
32-
"safety_identifier",
33-
}
34-
)
35-
36-
37-
def llama_stack_provider_from_client(
38-
client: AsyncLlamaStackClient | AsyncLlamaStackAsLibraryClient,
39-
) -> LlamaStackProvider:
40-
"""Construct a Pydantic AI Llama Stack provider backed by the same client as ``/query``."""
41-
if isinstance(client, AsyncLlamaStackAsLibraryClient):
42-
return LlamaStackProvider(library_client=client)
43-
api_key = client.api_key or "not-needed"
44-
base = str(client.base_url).rstrip("/")
45-
base_url = base if base.endswith("/v1") else f"{base}/v1"
46-
return LlamaStackProvider(
47-
base_url=base_url,
48-
api_key=api_key,
49-
http_client=client._client, # pylint: disable=protected-access
50-
)
51-
52-
53-
def _model_settings_from_responses_params(
54-
responses_params: ResponsesApiParams,
55-
) -> OpenAIResponsesModelSettings:
56-
"""Map ``ResponsesApiParams`` into Pydantic AI OpenAI Responses model settings."""
57-
payload = responses_params.model_dump(exclude_none=True)
58-
extra_body = {k: v for k, v in payload.items() if k in _LLS_RESPONSES_EXTRA_FIELDS}
59-
settings_dict: dict[str, Any] = {}
60-
if extra_body:
61-
settings_dict["extra_body"] = extra_body
62-
if responses_params.max_output_tokens is not None:
63-
settings_dict["max_tokens"] = responses_params.max_output_tokens
64-
if responses_params.temperature is not None:
65-
settings_dict["temperature"] = responses_params.temperature
66-
if responses_params.parallel_tool_calls is not None:
67-
settings_dict["parallel_tool_calls"] = responses_params.parallel_tool_calls
68-
if responses_params.extra_headers:
69-
settings_dict["extra_headers"] = dict(responses_params.extra_headers)
70-
settings_dict["openai_store"] = responses_params.store
71-
if responses_params.tools is not None:
72-
settings_dict["openai_native_tools"] = responses_params.tools
73-
if responses_params.previous_response_id is not None:
74-
settings_dict["openai_previous_response_id"] = (
75-
responses_params.previous_response_id
76-
)
77-
return cast(OpenAIResponsesModelSettings, settings_dict)
78-
7919

8020
def _skills_capability(
8121
skills_config: Optional[SkillsConfiguration],
@@ -147,15 +87,12 @@ def build_agent(
14787
``Agent`` configured for ``await agent.run(...)`` (or streaming) against the same
14888
stack configuration as ``client.responses.create(**responses_params.model_dump())``.
14989
"""
150-
provider = llama_stack_provider_from_client(client)
151-
settings = _model_settings_from_responses_params(responses_params)
15290
capabilities = _agent_capabilities(skills, no_tools=no_tools)
15391

154-
model = LlamaStackResponsesModel(
155-
responses_params.model,
156-
provider=provider,
157-
settings=settings,
92+
model = LlamaStackResponsesModel.from_llama_stack_client(
93+
responses_params.model, client, responses_params=responses_params
15894
)
95+
15996
return Agent(
16097
model,
16198
instructions=responses_params.instructions,

0 commit comments

Comments
 (0)