fix
This commit is contained in:
@@ -9,6 +9,7 @@ from pydantic import AliasChoices, BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
ConfigT = TypeVar("ConfigT")
|
||||
DataclassConfigT = TypeVar("DataclassConfigT")
|
||||
PydanticConfigT = TypeVar("PydanticConfigT", bound=BaseModel)
|
||||
|
||||
|
||||
@@ -155,8 +156,9 @@ def coerce_dataclass_config(
|
||||
f"got {type(value).__name__}"
|
||||
)
|
||||
|
||||
valid_fields = {config_field.name for config_field in fields(cast(Any, config_type))}
|
||||
unknown_fields = sorted(str(name) for name in set(value) - valid_fields)
|
||||
unknown_fields = sorted(
|
||||
str(name) for name in value if _dataclass_field_name_for_input(config_type, name) is None
|
||||
)
|
||||
if unknown_fields:
|
||||
raise TypeError(f"Unknown {parameter_name} settings: {', '.join(unknown_fields)}")
|
||||
|
||||
@@ -167,16 +169,22 @@ def coerce_dataclass_config(
|
||||
except (NameError, TypeError):
|
||||
resolved_annotations = {}
|
||||
|
||||
for config_field in fields(cast(Any, config_type)):
|
||||
field_value = value.get(config_field.name)
|
||||
annotation = resolved_annotations.get(config_field.name, config_field.type)
|
||||
config_fields = {
|
||||
config_field.name: config_field for config_field in fields(cast(Any, config_type))
|
||||
}
|
||||
for input_name, field_value in value.items():
|
||||
field_name = _dataclass_field_name_for_input(config_type, input_name)
|
||||
if field_name is None:
|
||||
continue
|
||||
config_field = config_fields[field_name]
|
||||
annotation = resolved_annotations.get(field_name, config_field.type)
|
||||
normalized_field_value = _normalize_declared_dataclass_value(
|
||||
field_value,
|
||||
annotation,
|
||||
parameter_name=f"{parameter_name}.{config_field.name}",
|
||||
parameter_name=f"{parameter_name}.{field_name}",
|
||||
)
|
||||
if normalized_field_value is not field_value:
|
||||
normalized[config_field.name] = normalized_field_value
|
||||
normalized[input_name] = normalized_field_value
|
||||
|
||||
return config_type(**normalized)
|
||||
|
||||
@@ -300,6 +308,77 @@ def _pydantic_fields_for_type(config_type: type[Any]) -> Mapping[str, FieldInfo]
|
||||
return {}
|
||||
|
||||
|
||||
def _dataclass_field_name_for_input(
|
||||
config_type: type[Any],
|
||||
input_name: object,
|
||||
) -> str | None:
|
||||
if resolved_field := _pydantic_field_for_input(config_type, input_name):
|
||||
return resolved_field[0]
|
||||
if not isinstance(input_name, str):
|
||||
return None
|
||||
|
||||
pydantic_fields = _pydantic_fields_for_type(config_type)
|
||||
if input_name in pydantic_fields:
|
||||
return None
|
||||
|
||||
dataclass_field_names = {config_field.name for config_field in fields(cast(Any, config_type))}
|
||||
return input_name if input_name in dataclass_field_names else None
|
||||
|
||||
|
||||
def _dataclass_input_values(
|
||||
value: Mapping[Any, Any],
|
||||
config_type: type[Any],
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
field_name: field_value
|
||||
for input_name, field_value in value.items()
|
||||
if (field_name := _dataclass_field_name_for_input(config_type, input_name)) is not None
|
||||
}
|
||||
|
||||
|
||||
def _dataclass_input_name_for_field(
|
||||
config_type: type[Any],
|
||||
field_name: str,
|
||||
) -> str:
|
||||
field_info = _pydantic_fields_for_type(config_type).get(field_name)
|
||||
if field_info is None:
|
||||
return field_name
|
||||
validation_alias = field_info.validation_alias
|
||||
if isinstance(validation_alias, str):
|
||||
return validation_alias
|
||||
if isinstance(validation_alias, AliasChoices):
|
||||
if string_alias := next(
|
||||
(choice for choice in validation_alias.choices if isinstance(choice, str)),
|
||||
None,
|
||||
):
|
||||
return string_alias
|
||||
return field_name
|
||||
|
||||
|
||||
def _replace_dataclass(
|
||||
value: DataclassConfigT,
|
||||
changes: Mapping[str, Any],
|
||||
) -> DataclassConfigT:
|
||||
config_fields = fields(cast(Any, value))
|
||||
valid_fields = {config_field.name for config_field in config_fields}
|
||||
unknown_fields = sorted(str(name) for name in set(changes) - valid_fields)
|
||||
if unknown_fields:
|
||||
raise TypeError(f"Unknown dataclass replacement fields: {', '.join(unknown_fields)}")
|
||||
|
||||
positional_values: list[Any] = []
|
||||
keyword_values: dict[str, Any] = {}
|
||||
for config_field in config_fields:
|
||||
if not config_field.init:
|
||||
continue
|
||||
field_value = changes.get(config_field.name, getattr(value, config_field.name))
|
||||
if config_field.kw_only:
|
||||
input_name = _dataclass_input_name_for_field(type(value), config_field.name)
|
||||
keyword_values[input_name] = field_value
|
||||
else:
|
||||
positional_values.append(field_value)
|
||||
return type(value)(*positional_values, **keyword_values)
|
||||
|
||||
|
||||
def _reject_ignored_pydantic_fields(
|
||||
value: object,
|
||||
normalized: object,
|
||||
@@ -337,6 +416,7 @@ def _reject_ignored_pydantic_fields(
|
||||
return
|
||||
|
||||
if isinstance(value, Mapping) and isinstance(normalized, Mapping):
|
||||
dropped_fields: list[str] = []
|
||||
for key, child_value in value.items():
|
||||
if key in normalized:
|
||||
normalized_child = normalized[key]
|
||||
@@ -350,6 +430,7 @@ def _reject_ignored_pydantic_fields(
|
||||
None,
|
||||
)
|
||||
if matching_key is None:
|
||||
dropped_fields.append(str(key))
|
||||
continue
|
||||
normalized_child = normalized[matching_key]
|
||||
_reject_ignored_pydantic_fields(
|
||||
@@ -357,6 +438,10 @@ def _reject_ignored_pydantic_fields(
|
||||
normalized_child,
|
||||
parameter_name=f"{parameter_name}[{key}]",
|
||||
)
|
||||
if dropped_fields:
|
||||
raise TypeError(
|
||||
f"Unknown {parameter_name} settings: {', '.join(sorted(dropped_fields))}"
|
||||
)
|
||||
return
|
||||
|
||||
if (
|
||||
|
||||
@@ -3,12 +3,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
from dataclasses import fields, replace
|
||||
from dataclasses import fields
|
||||
from typing import Any
|
||||
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from .._config_coercion import coerce_dataclass_config
|
||||
from .._config_coercion import (
|
||||
_dataclass_input_values,
|
||||
_replace_dataclass,
|
||||
coerce_dataclass_config,
|
||||
)
|
||||
|
||||
|
||||
def resolve_session_limit(
|
||||
@@ -39,21 +43,31 @@ class SessionSettings:
|
||||
override on top of this instance."""
|
||||
if override is None:
|
||||
return self
|
||||
override = coerce_session_settings(override)
|
||||
override_fields = (
|
||||
set(_dataclass_input_values(override, type(self)))
|
||||
if isinstance(override, dict)
|
||||
else None
|
||||
)
|
||||
override = coerce_session_settings(override, settings_type=type(self))
|
||||
|
||||
changes = {
|
||||
field.name: getattr(override, field.name)
|
||||
for field in fields(self)
|
||||
if getattr(override, field.name) is not None
|
||||
if (override_fields is None or field.name in override_fields)
|
||||
and getattr(override, field.name) is not None
|
||||
}
|
||||
|
||||
return replace(self, **changes)
|
||||
return _replace_dataclass(self, changes)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""Convert settings to a dictionary."""
|
||||
return dataclasses.asdict(self)
|
||||
|
||||
|
||||
def coerce_session_settings(value: SessionSettings | dict[str, Any]) -> SessionSettings:
|
||||
def coerce_session_settings(
|
||||
value: SessionSettings | dict[str, Any],
|
||||
*,
|
||||
settings_type: type[SessionSettings] = SessionSettings,
|
||||
) -> SessionSettings:
|
||||
"""Normalize session settings while preserving existing typed instances."""
|
||||
return coerce_dataclass_config(value, SessionSettings, parameter_name="session")
|
||||
return coerce_dataclass_config(value, settings_type, parameter_name="session")
|
||||
|
||||
+146
-45
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from dataclasses import fields, is_dataclass, replace
|
||||
from dataclasses import fields, is_dataclass
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Literal, TypeAlias, cast, get_type_hints
|
||||
|
||||
from openai import Omit as _Omit
|
||||
@@ -15,10 +15,13 @@ from pydantic.fields import FieldInfo
|
||||
from pydantic_core import ArgsKwargs, core_schema
|
||||
|
||||
from ._config_coercion import (
|
||||
_dataclass_field_name_for_input,
|
||||
_dataclass_input_values,
|
||||
_get_single_dataclass_type,
|
||||
_get_single_pydantic_model_type,
|
||||
_pydantic_field_for_input,
|
||||
_pydantic_fields_for_type,
|
||||
_replace_dataclass,
|
||||
)
|
||||
from .retry import (
|
||||
ModelRetryBackoffInput,
|
||||
@@ -210,28 +213,54 @@ class ModelSettings:
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def _validate_dictionary_values(cls, value: Any) -> Any:
|
||||
positional_values: dict[str, Any] = {}
|
||||
if isinstance(value, ArgsKwargs):
|
||||
values = {
|
||||
positional_values = {
|
||||
model_field.name: argument
|
||||
for model_field, argument in zip(fields(cls), value.args, strict=False)
|
||||
}
|
||||
values = dict(positional_values)
|
||||
values.update(value.kwargs or {})
|
||||
unknown_fields = sorted(
|
||||
str(name)
|
||||
for name in (value.kwargs or {})
|
||||
if _dataclass_field_name_for_input(cls, name) is None
|
||||
)
|
||||
canonical_values = {
|
||||
**positional_values,
|
||||
**_dataclass_input_values(value.kwargs or {}, cls),
|
||||
}
|
||||
elif isinstance(value, dict):
|
||||
values = value
|
||||
unknown_fields = sorted(
|
||||
str(name) for name in values if _dataclass_field_name_for_input(cls, name) is None
|
||||
)
|
||||
canonical_values = _dataclass_input_values(values, cls)
|
||||
else:
|
||||
return value
|
||||
|
||||
valid_fields = {model_field.name for model_field in fields(cls)}
|
||||
unknown_fields = sorted(str(name) for name in set(values) - valid_fields)
|
||||
if unknown_fields:
|
||||
raise TypeError(f"Unknown model settings: {', '.join(unknown_fields)}")
|
||||
|
||||
context_management = values.get("context_management")
|
||||
context_management = canonical_values.get("context_management")
|
||||
if isinstance(context_management, Iterable) and not isinstance(
|
||||
context_management, Sequence | Mapping
|
||||
):
|
||||
normalized_context_management = list(context_management)
|
||||
values = {**values, "context_management": normalized_context_management}
|
||||
context_management_input_name = (
|
||||
"context_management"
|
||||
if isinstance(value, ArgsKwargs) and "context_management" in positional_values
|
||||
else next(
|
||||
input_name
|
||||
for input_name in values
|
||||
if _dataclass_field_name_for_input(cls, input_name) == "context_management"
|
||||
)
|
||||
)
|
||||
values = {
|
||||
**values,
|
||||
context_management_input_name: normalized_context_management,
|
||||
}
|
||||
canonical_values["context_management"] = normalized_context_management
|
||||
if isinstance(value, ArgsKwargs):
|
||||
context_management_index = next(
|
||||
index
|
||||
@@ -247,7 +276,7 @@ class ModelSettings:
|
||||
value.args,
|
||||
{
|
||||
**(value.kwargs or {}),
|
||||
"context_management": normalized_context_management,
|
||||
context_management_input_name: normalized_context_management,
|
||||
},
|
||||
)
|
||||
else:
|
||||
@@ -255,7 +284,7 @@ class ModelSettings:
|
||||
|
||||
retry_settings_type, backoff_settings_type = _retry_settings_types(cls)
|
||||
_validate_structured_model_settings(
|
||||
values,
|
||||
canonical_values,
|
||||
reasoning_settings_type=_reasoning_settings_type(cls),
|
||||
tool_choice_settings_type=_tool_choice_settings_type(cls),
|
||||
retry_settings_type=retry_settings_type,
|
||||
@@ -302,6 +331,33 @@ class ModelSettings:
|
||||
if override is None:
|
||||
return _coerce_model_settings(self, parameter_name="ModelSettings")
|
||||
|
||||
override_fields: set[str] | None = None
|
||||
retry_override_fields: set[str] | None = None
|
||||
backoff_override_fields: set[str] | None = None
|
||||
if isinstance(override, dict):
|
||||
canonical_override = _dataclass_input_values(override, type(self))
|
||||
override_fields = set(canonical_override)
|
||||
retry_settings_type, backoff_settings_type = _retry_settings_types(
|
||||
type(self),
|
||||
inherited_model_settings=self,
|
||||
)
|
||||
if isinstance(retry_override := canonical_override.get("retry"), Mapping):
|
||||
canonical_retry_override = _dataclass_input_values(
|
||||
retry_override,
|
||||
retry_settings_type,
|
||||
)
|
||||
retry_override_fields = set(canonical_retry_override)
|
||||
if isinstance(
|
||||
backoff_override := canonical_retry_override.get("backoff"),
|
||||
Mapping,
|
||||
):
|
||||
backoff_override_fields = set(
|
||||
_dataclass_input_values(
|
||||
backoff_override,
|
||||
backoff_settings_type,
|
||||
)
|
||||
)
|
||||
|
||||
override = _coerce_model_settings(
|
||||
override,
|
||||
parameter_name="ModelSettings override",
|
||||
@@ -311,11 +367,14 @@ class ModelSettings:
|
||||
changes = {
|
||||
field.name: getattr(override, field.name)
|
||||
for field in fields(self)
|
||||
if getattr(override, field.name, None) is not None
|
||||
if (override_fields is None or field.name in override_fields)
|
||||
and getattr(override, field.name, None) is not None
|
||||
}
|
||||
|
||||
# Handle extra_args merging specially - merge dictionaries instead of replacing.
|
||||
if self.extra_args is not None or override.extra_args is not None:
|
||||
if (override_fields is None or "extra_args" in override_fields) and (
|
||||
self.extra_args is not None or override.extra_args is not None
|
||||
):
|
||||
merged_args = {}
|
||||
if self.extra_args:
|
||||
merged_args.update(self.extra_args)
|
||||
@@ -323,10 +382,17 @@ class ModelSettings:
|
||||
merged_args.update(override.extra_args)
|
||||
changes["extra_args"] = merged_args if merged_args else None
|
||||
|
||||
if self.retry is not None or override.retry is not None:
|
||||
changes["retry"] = _merge_retry_settings(self.retry, override.retry)
|
||||
if (override_fields is None or "retry" in override_fields) and (
|
||||
self.retry is not None or override.retry is not None
|
||||
):
|
||||
changes["retry"] = _merge_retry_settings(
|
||||
self.retry,
|
||||
override.retry,
|
||||
override_fields=retry_override_fields,
|
||||
backoff_override_fields=backoff_override_fields,
|
||||
)
|
||||
|
||||
return replace(self, **changes)
|
||||
return _replace_dataclass(self, changes)
|
||||
|
||||
def to_json_dict(self) -> dict[str, Any]:
|
||||
return cast(dict[str, Any], TypeAdapter(ModelSettings).dump_python(self, mode="json"))
|
||||
@@ -396,11 +462,13 @@ def _coerce_model_settings(
|
||||
or normalized_tool_choice is not value.tool_choice
|
||||
or normalized_retry is not value.retry
|
||||
):
|
||||
return replace(
|
||||
return _replace_dataclass(
|
||||
value,
|
||||
reasoning=normalized_reasoning,
|
||||
tool_choice=normalized_tool_choice,
|
||||
retry=normalized_retry,
|
||||
{
|
||||
"reasoning": normalized_reasoning,
|
||||
"tool_choice": normalized_tool_choice,
|
||||
"retry": normalized_retry,
|
||||
},
|
||||
)
|
||||
return value
|
||||
|
||||
@@ -410,11 +478,15 @@ def _coerce_model_settings(
|
||||
f"got {type(value).__name__}"
|
||||
)
|
||||
|
||||
valid_fields = {model_field.name for model_field in fields(model_settings_type)}
|
||||
unknown_fields = sorted(str(name) for name in set(value) - valid_fields)
|
||||
unknown_fields = sorted(
|
||||
str(name)
|
||||
for name in value
|
||||
if _dataclass_field_name_for_input(model_settings_type, name) is None
|
||||
)
|
||||
if unknown_fields:
|
||||
raise TypeError(f"Unknown model settings: {', '.join(unknown_fields)}")
|
||||
|
||||
canonical_value = _dataclass_input_values(value, model_settings_type)
|
||||
retry_settings_type, backoff_settings_type = _retry_settings_types(
|
||||
model_settings_type,
|
||||
inherited_model_settings=inherited_model_settings,
|
||||
@@ -428,7 +500,7 @@ def _coerce_model_settings(
|
||||
inherited_model_settings=inherited_model_settings,
|
||||
)
|
||||
_validate_structured_model_settings(
|
||||
value,
|
||||
canonical_value,
|
||||
reasoning_settings_type=reasoning_settings_type,
|
||||
tool_choice_settings_type=tool_choice_settings_type,
|
||||
retry_settings_type=retry_settings_type,
|
||||
@@ -436,23 +508,29 @@ def _coerce_model_settings(
|
||||
)
|
||||
|
||||
normalized_value = dict(value)
|
||||
if isinstance(reasoning := value.get("reasoning"), Mapping):
|
||||
normalized_value["reasoning"] = TypeAdapter(reasoning_settings_type).validate_python(
|
||||
dict(reasoning)
|
||||
)
|
||||
if isinstance(tool_choice := value.get("tool_choice"), Mapping):
|
||||
normalized_value["tool_choice"] = TypeAdapter(tool_choice_settings_type).validate_python(
|
||||
dict(tool_choice)
|
||||
)
|
||||
if isinstance(retry := value.get("retry"), Mapping):
|
||||
normalized_retry_payload = dict(retry)
|
||||
if isinstance(backoff := normalized_retry_payload.get("backoff"), Mapping):
|
||||
normalized_retry_payload["backoff"] = TypeAdapter(
|
||||
backoff_settings_type
|
||||
).validate_python(dict(backoff))
|
||||
normalized_value["retry"] = TypeAdapter(retry_settings_type).validate_python(
|
||||
normalized_retry_payload
|
||||
)
|
||||
for input_name, input_value in value.items():
|
||||
field_name = _dataclass_field_name_for_input(model_settings_type, input_name)
|
||||
if field_name == "reasoning" and isinstance(input_value, Mapping):
|
||||
normalized_value[input_name] = TypeAdapter(reasoning_settings_type).validate_python(
|
||||
dict(input_value)
|
||||
)
|
||||
elif field_name == "tool_choice" and isinstance(input_value, Mapping):
|
||||
normalized_value[input_name] = TypeAdapter(tool_choice_settings_type).validate_python(
|
||||
dict(input_value)
|
||||
)
|
||||
elif field_name == "retry" and isinstance(input_value, Mapping):
|
||||
normalized_retry_payload = dict(input_value)
|
||||
for retry_input_name, retry_input_value in input_value.items():
|
||||
if _dataclass_field_name_for_input(
|
||||
retry_settings_type,
|
||||
retry_input_name,
|
||||
) == "backoff" and isinstance(retry_input_value, Mapping):
|
||||
normalized_retry_payload[retry_input_name] = TypeAdapter(
|
||||
backoff_settings_type
|
||||
).validate_python(dict(retry_input_value))
|
||||
normalized_value[input_name] = TypeAdapter(retry_settings_type).validate_python(
|
||||
normalized_retry_payload
|
||||
)
|
||||
|
||||
return TypeAdapter(model_settings_type).validate_python(normalized_value)
|
||||
|
||||
@@ -650,13 +728,14 @@ def _validate_structured_model_settings(
|
||||
)
|
||||
|
||||
if isinstance(retry := value.get("retry"), Mapping):
|
||||
canonical_retry = _dataclass_input_values(retry, retry_settings_type)
|
||||
reject_unknown_pydantic_fields(
|
||||
retry,
|
||||
retry_settings_type,
|
||||
"retry",
|
||||
allow_model_extra=retry_settings_type is not ModelRetrySettings,
|
||||
)
|
||||
if isinstance(backoff := retry.get("backoff"), Mapping):
|
||||
if isinstance(backoff := canonical_retry.get("backoff"), Mapping):
|
||||
reject_unknown_pydantic_fields(
|
||||
backoff,
|
||||
backoff_settings_type,
|
||||
@@ -685,6 +764,9 @@ def _validate_structured_model_settings(
|
||||
def _merge_retry_settings(
|
||||
inherited: ModelRetrySettings | None,
|
||||
override: ModelRetrySettings | None,
|
||||
*,
|
||||
override_fields: set[str] | None = None,
|
||||
backoff_override_fields: set[str] | None = None,
|
||||
) -> ModelRetrySettings | None:
|
||||
inherited = _normalize_retry_settings(inherited)
|
||||
override = _normalize_retry_settings(override)
|
||||
@@ -693,18 +775,36 @@ def _merge_retry_settings(
|
||||
if override is None:
|
||||
return inherited
|
||||
|
||||
merged_backoff = _merge_backoff_settings(inherited.backoff, override.backoff)
|
||||
merged_backoff = (
|
||||
_merge_backoff_settings(
|
||||
inherited.backoff,
|
||||
override.backoff,
|
||||
override_fields=backoff_override_fields,
|
||||
)
|
||||
if override_fields is None or "backoff" in override_fields
|
||||
else _normalize_backoff_settings(_coerce_backoff_settings(inherited.backoff))
|
||||
)
|
||||
retry_changes = {
|
||||
field.name: getattr(override, field.name)
|
||||
for field in fields(inherited)
|
||||
if field.name != "backoff" and getattr(override, field.name, None) is not None
|
||||
if field.name != "backoff"
|
||||
and (override_fields is None or field.name in override_fields)
|
||||
and getattr(override, field.name, None) is not None
|
||||
}
|
||||
return replace(inherited, **retry_changes, backoff=merged_backoff)
|
||||
return _replace_dataclass(
|
||||
inherited,
|
||||
{
|
||||
**retry_changes,
|
||||
"backoff": merged_backoff,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _merge_backoff_settings(
|
||||
inherited: ModelRetryBackoffInput | None,
|
||||
override: ModelRetryBackoffInput | None,
|
||||
*,
|
||||
override_fields: set[str] | None = None,
|
||||
) -> ModelRetryBackoffSettings | None:
|
||||
inherited = _normalize_backoff_settings(_coerce_backoff_settings(inherited))
|
||||
override = _normalize_backoff_settings(_coerce_backoff_settings(override))
|
||||
@@ -716,9 +816,10 @@ def _merge_backoff_settings(
|
||||
changes = {
|
||||
field.name: getattr(override, field.name)
|
||||
for field in fields(inherited)
|
||||
if getattr(override, field.name, None) is not None
|
||||
if (override_fields is None or field.name in override_fields)
|
||||
and getattr(override, field.name, None) is not None
|
||||
}
|
||||
return replace(inherited, **changes)
|
||||
return _replace_dataclass(inherited, changes)
|
||||
|
||||
|
||||
def _normalize_retry_settings(
|
||||
@@ -735,7 +836,7 @@ def _normalize_retry_settings(
|
||||
backoff = _normalize_backoff_settings(_coerce_backoff_settings(settings.backoff))
|
||||
if backoff is not settings.backoff:
|
||||
changes["backoff"] = backoff
|
||||
return replace(settings, **changes) if changes else settings
|
||||
return _replace_dataclass(settings, changes) if changes else settings
|
||||
|
||||
|
||||
def _normalize_backoff_settings(
|
||||
@@ -749,4 +850,4 @@ def _normalize_backoff_settings(
|
||||
for model_field in fields(settings)
|
||||
if isinstance(value := getattr(settings, model_field.name), FieldInfo)
|
||||
}
|
||||
return replace(settings, **changes) if changes else settings
|
||||
return _replace_dataclass(settings, changes) if changes else settings
|
||||
|
||||
@@ -237,6 +237,7 @@ class SandboxRunConfig:
|
||||
self.snapshot,
|
||||
type(snapshot),
|
||||
parameter_name="sandbox.snapshot",
|
||||
allow_model_extra=True,
|
||||
)
|
||||
if isinstance(self.options, dict):
|
||||
from .sandbox.session.sandbox_client import BaseSandboxClientOptions
|
||||
|
||||
@@ -6,6 +6,8 @@ import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from pydantic import Field
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from agents import Agent, RunConfig, Runner, SessionSettings, SQLiteSession, TResponseInputItem
|
||||
from tests.fake_model import FakeModel
|
||||
@@ -782,6 +784,26 @@ async def test_session_settings_resolve():
|
||||
assert final_none.limit == 100
|
||||
|
||||
|
||||
def test_session_settings_subclass_resolves_dictionary_with_declared_type() -> None:
|
||||
@dataclass
|
||||
class ProviderSessionSettings(SessionSettings):
|
||||
provider_limit: int | None = Field(default=None, alias="providerLimit")
|
||||
provider_mode: str | None = "default"
|
||||
|
||||
settings = ProviderSessionSettings(
|
||||
limit=100,
|
||||
providerLimit=3,
|
||||
provider_mode="custom",
|
||||
)
|
||||
|
||||
resolved = settings.resolve({"limit": 10, "providerLimit": 4})
|
||||
|
||||
assert isinstance(resolved, ProviderSessionSettings)
|
||||
assert resolved.limit == 10
|
||||
assert resolved.provider_limit == 4
|
||||
assert resolved.provider_mode == "custom"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_with_session_settings_override():
|
||||
"""Test that RunConfig can override session's default settings."""
|
||||
|
||||
@@ -435,6 +435,74 @@ def test_model_settings_subclass_honors_declared_nested_aliases(use_resolve: boo
|
||||
assert settings.retry.backoff.provider_field == "backoff-value"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("use_resolve", [False, True], ids=["constructor", "resolve"])
|
||||
def test_model_settings_subclass_honors_declared_top_level_aliases(
|
||||
use_resolve: bool,
|
||||
) -> None:
|
||||
@pydantic_dataclass
|
||||
class ProviderModelSettings(ModelSettings):
|
||||
provider_field: str | None = Field(default=None, alias="providerField")
|
||||
|
||||
payload = {"providerField": "provider-value"}
|
||||
settings = (
|
||||
ProviderModelSettings().resolve(payload)
|
||||
if use_resolve
|
||||
else ProviderModelSettings(**cast(Any, payload))
|
||||
)
|
||||
|
||||
assert isinstance(settings, ProviderModelSettings)
|
||||
assert settings.provider_field == "provider-value"
|
||||
|
||||
|
||||
def test_dictionary_resolve_preserves_omitted_subclass_defaults() -> None:
|
||||
@pydantic_dataclass
|
||||
class ProviderBackoffSettings(ModelRetryBackoffSettings):
|
||||
provider_mode: str | None = "backoff-default"
|
||||
|
||||
@pydantic_dataclass
|
||||
class ProviderRetrySettings(ModelRetrySettings):
|
||||
backoff: ProviderBackoffSettings | None = None
|
||||
provider_mode: str | None = "retry-default"
|
||||
|
||||
@pydantic_dataclass
|
||||
class ProviderModelSettings(ModelSettings):
|
||||
retry: ProviderRetrySettings | None = None
|
||||
provider_mode: str | None = "model-default"
|
||||
|
||||
settings = ProviderModelSettings(
|
||||
provider_mode="model-custom",
|
||||
retry=ProviderRetrySettings(
|
||||
max_retries=1,
|
||||
provider_mode="retry-custom",
|
||||
backoff=ProviderBackoffSettings(
|
||||
initial_delay=0.1,
|
||||
provider_mode="backoff-custom",
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
resolved = settings.resolve(
|
||||
{
|
||||
"temperature": 0.2,
|
||||
"retry": {
|
||||
"max_retries": 2,
|
||||
"backoff": {"max_delay": 3.0},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert isinstance(resolved, ProviderModelSettings)
|
||||
assert resolved.temperature == 0.2
|
||||
assert resolved.provider_mode == "model-custom"
|
||||
assert isinstance(resolved.retry, ProviderRetrySettings)
|
||||
assert resolved.retry.max_retries == 2
|
||||
assert resolved.retry.provider_mode == "retry-custom"
|
||||
assert isinstance(resolved.retry.backoff, ProviderBackoffSettings)
|
||||
assert resolved.retry.backoff.initial_delay == 0.1
|
||||
assert resolved.retry.backoff.max_delay == 3.0
|
||||
assert resolved.retry.backoff.provider_mode == "backoff-custom"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("use_resolve", [False, True], ids=["constructor", "resolve"])
|
||||
def test_model_settings_subclass_uses_declared_tool_choice_type(use_resolve: bool) -> None:
|
||||
@dataclass
|
||||
|
||||
@@ -6,6 +6,7 @@ from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from agents import (
|
||||
Agent,
|
||||
@@ -198,6 +199,39 @@ class _CustomRawUnionPydanticOptionsSandboxClient(
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _CustomTypedDictPayload(TypedDict):
|
||||
name: str
|
||||
|
||||
|
||||
class _CustomTypedDictPydanticSandboxClientOptions(BaseModel):
|
||||
payload: _CustomTypedDictPayload
|
||||
|
||||
|
||||
class _CustomTypedDictPydanticOptionsSandboxClient(
|
||||
BaseSandboxClient[_CustomTypedDictPydanticSandboxClientOptions]
|
||||
):
|
||||
backend_id = "custom-typed-dict-pydantic"
|
||||
options_type = _CustomTypedDictPydanticSandboxClientOptions
|
||||
|
||||
async def create(
|
||||
self,
|
||||
*,
|
||||
snapshot: SnapshotSpec | SnapshotBase | None = None,
|
||||
manifest: Manifest | None = None,
|
||||
options: _CustomTypedDictPydanticSandboxClientOptions,
|
||||
) -> SandboxSession:
|
||||
raise NotImplementedError
|
||||
|
||||
async def delete(self, session: SandboxSession) -> SandboxSession:
|
||||
raise NotImplementedError
|
||||
|
||||
async def resume(self, state: SandboxSessionState) -> SandboxSession:
|
||||
raise NotImplementedError
|
||||
|
||||
def deserialize_session_state(self, payload: dict[str, object]) -> SandboxSessionState:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _CustomSnapshotSpec(SnapshotSpec):
|
||||
type: Literal["custom-run-config-spec"] = "custom-run-config-spec"
|
||||
id: str
|
||||
@@ -207,6 +241,15 @@ class _CustomSnapshotSpec(SnapshotSpec):
|
||||
return NoopSnapshot(id=f"{snapshot_id}-{self.id}-{self.custom_value}")
|
||||
|
||||
|
||||
class _CustomExtraSnapshotSpec(SnapshotSpec):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
type: Literal["custom-extra-run-config-spec"] = "custom-extra-run-config-spec"
|
||||
|
||||
def build(self, snapshot_id: str) -> SnapshotBase:
|
||||
return NoopSnapshot(id=snapshot_id)
|
||||
|
||||
|
||||
class _DynamicTypeSnapshotSpec(SnapshotSpec):
|
||||
type: str
|
||||
|
||||
@@ -442,6 +485,20 @@ def test_sandbox_config_normalizes_registered_custom_snapshot_spec_dictionary()
|
||||
assert resolved.id == "snapshot-template-value"
|
||||
|
||||
|
||||
def test_sandbox_config_preserves_registered_custom_snapshot_spec_extras() -> None:
|
||||
spec = _CustomExtraSnapshotSpec.model_validate(
|
||||
{
|
||||
"type": "custom-extra-run-config-spec",
|
||||
"providerFlag": True,
|
||||
}
|
||||
)
|
||||
|
||||
config = SandboxRunConfig(snapshot=spec.model_dump())
|
||||
|
||||
assert isinstance(config.snapshot, _CustomExtraSnapshotSpec)
|
||||
assert config.snapshot.model_extra == {"providerFlag": True}
|
||||
|
||||
|
||||
def test_sandbox_config_rejects_unknown_custom_snapshot_spec_fields() -> None:
|
||||
with pytest.raises(TypeError, match="Unknown sandbox.snapshot settings: custom_valeu"):
|
||||
SandboxRunConfig(
|
||||
@@ -593,6 +650,22 @@ def test_sandbox_config_preserves_raw_mapping_sequence_elements() -> None:
|
||||
assert config.options.raw_rules[0] is raw_rule
|
||||
|
||||
|
||||
def test_sandbox_config_rejects_typed_dict_fields_dropped_by_pydantic() -> None:
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match=r"Unknown sandbox\.options\.payload settings: naem",
|
||||
):
|
||||
SandboxRunConfig(
|
||||
client=_CustomTypedDictPydanticOptionsSandboxClient(),
|
||||
options={
|
||||
"payload": {
|
||||
"name": "read",
|
||||
"naem": "typo",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_sandbox_config_preserves_explicitly_raw_nested_custom_client_options() -> None:
|
||||
raw_timeouts = {"exec_timeout_s": 1.5}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user