codekingpro/portable-devtools
114k
1"""Types for setting agent response formats."""2 3from __future__ import annotations4 5import json6import uuid7from dataclasses import dataclass, is_dataclass8from types import UnionType9from typing import (10 TYPE_CHECKING,11 Any,12 Generic,13 Literal,14 TypeVar,15 Union,16 get_args,17 get_origin,18)19 20from langchain_core.tools import BaseTool, StructuredTool21from pydantic import BaseModel, TypeAdapter22from typing_extensions import Self, is_typeddict23 24if TYPE_CHECKING:25 from collections.abc import Callable, Iterable26 27 from langchain_core.messages import AIMessage28 29# Supported schema types: Pydantic models, dataclasses, TypedDict, JSON schema dicts30SchemaT = TypeVar("SchemaT")31 32SchemaKind = Literal["pydantic", "dataclass", "typeddict", "json_schema"]33 34 35class StructuredOutputError(Exception):36 """Base class for structured output errors."""37 38 ai_message: AIMessage39 40 41class MultipleStructuredOutputsError(StructuredOutputError):42 """Raised when model returns multiple structured output tool calls when only one is expected."""43 44 def __init__(self, tool_names: list[str], ai_message: AIMessage) -> None:45 """Initialize `MultipleStructuredOutputsError`.46 47 Args:48 tool_names: The names of the tools called for structured output.49 ai_message: The AI message that contained the invalid multiple tool calls.50 """51 self.tool_names = tool_names52 self.ai_message = ai_message53 54 super().__init__(55 "Model incorrectly returned multiple structured responses "56 f"({', '.join(tool_names)}) when only one is expected."57 )58 59 60class StructuredOutputValidationError(StructuredOutputError):61 """Raised when structured output tool call arguments fail to parse according to the schema."""62 63 def __init__(self, tool_name: str, source: Exception, ai_message: AIMessage) -> None:64 """Initialize `StructuredOutputValidationError`.65 66 Args:67 tool_name: The name of the tool that failed.68 source: The exception that occurred.69 ai_message: The AI message that contained the invalid structured output.70 """71 self.tool_name = tool_name72 self.source = source73 self.ai_message = ai_message74 super().__init__(f"Failed to parse structured output for tool '{tool_name}': {source}.")75 76 77def _parse_with_schema(78 schema: type[SchemaT] | dict[str, Any], schema_kind: SchemaKind, data: dict[str, Any]79) -> Any:80 """Parse data using for any supported schema type.81 82 Args:83 schema: The schema type (Pydantic model, `dataclass`, or `TypedDict`)84 schema_kind: One of `'pydantic'`, `'dataclass'`, `'typeddict'`, or85 `'json_schema'`86 data: The data to parse87 88 Returns:89 The parsed instance according to the schema type90 91 Raises:92 ValueError: If parsing fails93 """94 if schema_kind == "json_schema":95 return data96 try:97 adapter: TypeAdapter[SchemaT] = TypeAdapter(schema)98 return adapter.validate_python(data)99 except Exception as e:100 schema_name = getattr(schema, "__name__", str(schema))101 msg = f"Failed to parse data to {schema_name}: {e}"102 raise ValueError(msg) from e103 104 105@dataclass(init=False)106class _SchemaSpec(Generic[SchemaT]):107 """Describes a structured output schema."""108 109 schema: type[SchemaT] | dict[str, Any]110 """The schema for the response, can be a Pydantic model, `dataclass`, `TypedDict`,111 or JSON schema dict.112 """113 114 name: str115 """Name of the schema, used for tool calling.116 117 If not provided, the name will be the class name for models/dataclasses/TypedDicts,118 or the `title` field for JSON schemas.119 120 Falls back to a generated name if unavailable.121 """122 123 description: str124 """Custom description of the schema.125 126 If not provided, will use the model's docstring.127 """128 129 schema_kind: SchemaKind130 """The kind of schema."""131 132 json_schema: dict[str, Any]133 """JSON schema associated with the schema."""134 135 strict: bool | None = None136 """Whether to enforce strict validation of the schema."""137 138 def __init__(139 self,140 schema: type[SchemaT] | dict[str, Any],141 *,142 name: str | None = None,143 description: str | None = None,144 strict: bool | None = None,145 ) -> None:146 """Initialize `SchemaSpec` with schema and optional parameters.147 148 Args:149 schema: Schema to describe.150 name: Optional name for the schema.151 description: Optional description for the schema.152 strict: Whether to enforce strict validation of the schema.153 154 Raises:155 ValueError: If the schema type is unsupported.156 """157 self.schema = schema158 159 if name:160 self.name = name161 elif isinstance(schema, dict):162 self.name = str(schema.get("title", f"response_format_{str(uuid.uuid4())[:4]}"))163 else:164 self.name = str(getattr(schema, "__name__", f"response_format_{str(uuid.uuid4())[:4]}"))165 166 self.description = description or (167 schema.get("description", "")168 if isinstance(schema, dict)169 else getattr(schema, "__doc__", None) or ""170 )171 172 self.strict = strict173 174 if isinstance(schema, dict):175 self.schema_kind = "json_schema"176 self.json_schema = schema177 elif isinstance(schema, type) and issubclass(schema, BaseModel):178 self.schema_kind = "pydantic"179 self.json_schema = schema.model_json_schema()180 elif is_dataclass(schema):181 self.schema_kind = "dataclass"182 self.json_schema = TypeAdapter(schema).json_schema()183 elif is_typeddict(schema):184 self.schema_kind = "typeddict"185 self.json_schema = TypeAdapter(schema).json_schema()186 else:187 msg = (188 f"Unsupported schema type: {type(schema)}. "189 f"Supported types: Pydantic models, dataclasses, TypedDicts, and JSON schema dicts."190 )191 raise ValueError(msg)192 193 194@dataclass(init=False)195class ToolStrategy(Generic[SchemaT]):196 """Use a tool calling strategy for model responses."""197 198 schema: type[SchemaT] | UnionType | dict[str, Any]199 """Schema for the tool calls."""200 201 schema_specs: list[_SchemaSpec[Any]]202 """Schema specs for the tool calls."""203 204 tool_message_content: str | None205 """The content of the tool message to be returned when the model calls206 an artificial structured output tool.207 """208 209 handle_errors: (210 bool | str | type[Exception] | tuple[type[Exception], ...] | Callable[[Exception], str]211 )212 """Error handling strategy for structured output via `ToolStrategy`.213 214 - `True`: Catch all errors with default error template215 - `str`: Catch all errors with this custom message216 - `type[Exception]`: Only catch this exception type with default message217 - `tuple[type[Exception], ...]`: Only catch these exception types with default218 message219 - `Callable[[Exception], str]`: Custom function that returns error message220 - `False`: No retry, let exceptions propagate221 """222 223 def __init__(224 self,225 schema: type[SchemaT] | UnionType | dict[str, Any],226 *,227 tool_message_content: str | None = None,228 handle_errors: bool229 | str230 | type[Exception]231 | tuple[type[Exception], ...]232 | Callable[[Exception], str] = True,233 ) -> None:234 """Initialize `ToolStrategy`.235 236 Initialize `ToolStrategy` with schemas, tool message content, and error handling237 strategy.238 """239 self.schema = schema240 self.tool_message_content = tool_message_content241 self.handle_errors = handle_errors242 243 def _iter_variants(schema: Any) -> Iterable[Any]:244 """Yield leaf variants from Union and JSON Schema oneOf."""245 if get_origin(schema) in {UnionType, Union}:246 for arg in get_args(schema):247 yield from _iter_variants(arg)248 return249 250 if isinstance(schema, dict) and "oneOf" in schema:251 for sub in schema.get("oneOf", []):252 yield from _iter_variants(sub)253 return254 255 yield schema256 257 self.schema_specs = [_SchemaSpec(s) for s in _iter_variants(schema)]258 259 260@dataclass(init=False)261class ProviderStrategy(Generic[SchemaT]):262 """Use the model provider's native structured output method."""263 264 schema: type[SchemaT] | dict[str, Any]265 """Schema for native mode."""266 267 schema_spec: _SchemaSpec[SchemaT]268 """Schema spec for native mode."""269 270 def __init__(271 self,272 schema: type[SchemaT] | dict[str, Any],273 *,274 strict: bool | None = None,275 ) -> None:276 """Initialize `ProviderStrategy` with schema.277 278 Args:279 schema: Schema to enforce via the provider's native structured output.280 strict: Whether to request strict provider-side schema enforcement.281 """282 self.schema = schema283 self.schema_spec = _SchemaSpec(schema, strict=strict)284 285 def to_model_kwargs(self) -> dict[str, Any]:286 """Convert to kwargs to bind to a model to force structured output.287 288 Returns:289 The kwargs to bind to a model.290 """291 # OpenAI:292 # - see https://platform.openai.com/docs/guides/structured-outputs293 json_schema: dict[str, Any] = {294 "name": self.schema_spec.name,295 "schema": self.schema_spec.json_schema,296 }297 if self.schema_spec.strict:298 json_schema["strict"] = True299 300 response_format: dict[str, Any] = {301 "type": "json_schema",302 "json_schema": json_schema,303 }304 return {"response_format": response_format}305 306 307@dataclass308class OutputToolBinding(Generic[SchemaT]):309 """Information for tracking structured output tool metadata.310 311 This contains all necessary information to handle structured responses generated via312 tool calls, including the original schema, its type classification, and the313 corresponding tool implementation used by the tools strategy.314 """315 316 schema: type[SchemaT] | dict[str, Any]317 """The original schema provided for structured output (Pydantic model, dataclass,318 TypedDict, or JSON schema dict).319 """320 321 schema_kind: SchemaKind322 """Classification of the schema type for proper response construction."""323 324 tool: BaseTool325 """LangChain tool instance created from the schema for model binding."""326 327 @classmethod328 def from_schema_spec(cls, schema_spec: _SchemaSpec[SchemaT]) -> Self:329 """Create an `OutputToolBinding` instance from a `SchemaSpec`.330 331 Args:332 schema_spec: The `SchemaSpec` to convert333 334 Returns:335 An `OutputToolBinding` instance with the appropriate tool created336 """337 return cls(338 schema=schema_spec.schema,339 schema_kind=schema_spec.schema_kind,340 tool=StructuredTool(341 args_schema=schema_spec.json_schema,342 name=schema_spec.name,343 description=schema_spec.description,344 ),345 )346 347 def parse(self, tool_args: dict[str, Any]) -> SchemaT:348 """Parse tool arguments according to the schema.349 350 Args:351 tool_args: The arguments from the tool call352 353 Returns:354 The parsed response according to the schema type355 356 Raises:357 ValueError: If parsing fails358 """359 return _parse_with_schema(self.schema, self.schema_kind, tool_args)360 361 362@dataclass363class ProviderStrategyBinding(Generic[SchemaT]):364 """Information for tracking native structured output metadata.365 366 This contains all necessary information to handle structured responses generated via367 native provider output, including the original schema, its type classification, and368 parsing logic for provider-enforced JSON.369 """370 371 schema: type[SchemaT] | dict[str, Any]372 """The original schema provided for structured output (Pydantic model, `dataclass`,373 `TypedDict`, or JSON schema dict).374 """375 376 schema_kind: SchemaKind377 """Classification of the schema type for proper response construction."""378 379 @classmethod380 def from_schema_spec(cls, schema_spec: _SchemaSpec[SchemaT]) -> Self:381 """Create a `ProviderStrategyBinding` instance from a `SchemaSpec`.382 383 Args:384 schema_spec: The `SchemaSpec` to convert385 386 Returns:387 A `ProviderStrategyBinding` instance for parsing native structured output388 """389 return cls(390 schema=schema_spec.schema,391 schema_kind=schema_spec.schema_kind,392 )393 394 def parse(self, response: AIMessage) -> SchemaT:395 """Parse `AIMessage` content according to the schema.396 397 Args:398 response: The `AIMessage` containing the structured output399 400 Returns:401 The parsed response according to the schema402 403 Raises:404 ValueError: If text extraction, JSON parsing or schema validation fails405 """406 # Extract text content from AIMessage and parse as JSON407 raw_text = self._extract_text_content_from_message(response)408 409 try:410 data = json.loads(raw_text)411 except Exception as e:412 schema_name = getattr(self.schema, "__name__", "response_format")413 msg = (414 f"Native structured output expected valid JSON for {schema_name}, "415 f"but parsing failed: {e}."416 )417 raise ValueError(msg) from e418 419 # Parse according to schema420 return _parse_with_schema(self.schema, self.schema_kind, data)421 422 @staticmethod423 def _extract_text_content_from_message(message: AIMessage) -> str:424 """Extract text content from an `AIMessage`.425 426 Args:427 message: The AI message to extract text from428 429 Returns:430 The extracted text content431 """432 content = message.content433 if isinstance(content, str):434 return content435 parts: list[str] = []436 for c in content:437 if isinstance(c, dict):438 if c.get("type") == "text" and "text" in c:439 parts.append(str(c["text"]))440 elif "content" in c and isinstance(c["content"], str):441 parts.append(c["content"])442 else:443 parts.append(str(c))444 return "".join(parts)445 446 447class AutoStrategy(Generic[SchemaT]):448 """Automatically select the best strategy for structured output."""449 450 schema: type[SchemaT] | dict[str, Any]451 """Schema for automatic mode."""452 453 def __init__(454 self,455 schema: type[SchemaT] | dict[str, Any],456 ) -> None:457 """Initialize `AutoStrategy` with schema."""458 self.schema = schema459 460 461ResponseFormat = ToolStrategy[SchemaT] | ProviderStrategy[SchemaT] | AutoStrategy[SchemaT]462"""Union type for all supported response format strategies."""463 