This commit is contained in:
Kazuhiro Sera
2026-07-23 21:09:45 +09:00
parent 125f9f88bb
commit 94130c1ac9
7 changed files with 423 additions and 59 deletions
+92 -7
View File
@@ -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 (
+21 -7
View File
@@ -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
View File
@@ -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
+1
View File
@@ -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
+22
View File
@@ -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
+73
View File
@@ -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}