Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
structured_output.py463 linesDownload Raw Back to agents
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 
codekingpro/portable-devtools · Team Ai