Skip to content

Commit 52f2d12

Browse files
authored
Merge pull request #2025 from Jazzcort/restructure-pydantic-ai-utils-function
LCORE-:Restructure pydantic_ai utils into LlamaStack model and provider classes
2 parents 7d520c9 + 278a033 commit 52f2d12

8 files changed

Lines changed: 1017 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, Optional, cast
22+
from typing import Any, Final, Optional, 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 ( # pylint: disable=reimported
20+
AsyncLlamaStackAsLibraryClient,
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
@@ -3,19 +3,17 @@
33
from __future__ import annotations
44

55
import re
6-
from typing import Any, Final, Optional, cast
6+
from typing import Any, Final, Optional
77

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

1514
from models.common.responses.responses_api_params import ResponsesApiParams
1615
from models.config import SkillsConfiguration
1716
from pydantic_ai_lightspeed.llamastack import (
18-
LlamaStackProvider,
1917
LlamaStackResponsesModel,
2018
)
2119

@@ -24,64 +22,6 @@
2422
_BUILTIN_CAPABILITY_SERVER_SOURCE: Final[str] = "builtin"
2523
_CAPABILITY_TOOL_TYPE: Final[str] = "tool"
2624

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

8626
def _skills_capability(
8727
skills_config: Optional[SkillsConfiguration],
@@ -239,15 +179,12 @@ def build_agent(
239179
``Agent`` configured for ``await agent.run(...)`` (or streaming) against the same
240180
stack configuration as ``client.responses.create(**responses_params.model_dump())``.
241181
"""
242-
provider = llama_stack_provider_from_client(client)
243-
settings = _model_settings_from_responses_params(responses_params)
244182
capabilities = _agent_capabilities(skills, no_tools=no_tools)
245183

246-
model = LlamaStackResponsesModel(
247-
responses_params.model,
248-
provider=provider,
249-
settings=settings,
184+
model = LlamaStackResponsesModel.from_llama_stack_client(
185+
responses_params.model, client, responses_params=responses_params
250186
)
187+
251188
return Agent(
252189
model,
253190
instructions=responses_params.instructions,

0 commit comments

Comments
 (0)