Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
116 changes: 75 additions & 41 deletions src/cohere/core/pydantic_utilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import json
import logging
from collections import defaultdict
from dataclasses import asdict
from dataclasses import asdict, dataclass
from typing import (
TYPE_CHECKING,
Any,
Expand Down Expand Up @@ -184,6 +184,59 @@ def _get_type_adapter(type_: Type[Any]) -> Any:
return adapter


@dataclass(frozen=True)
class _FieldAliasInfo:
name_to_alias: Dict[str, str]
differing_names: Set[str]
ambiguous_keys: Set[str]
relevant_keys: Set[str]


_field_alias_info_cache: Dict[int, _FieldAliasInfo] = {}


def _get_field_alias_info(cls: Any) -> _FieldAliasInfo:
"""
Computes, once per class, the field-name-to-alias map, the subset of names whose alias
differs from the name, and the set of keys that are ambiguous (a key that is simultaneously
a field name and the alias of a different field). Cached by class id, mirroring
`_type_adapter_cache`, since Pydantic fields are static per class.
"""
key = id(cls)
info = _field_alias_info_cache.get(key)
if info is not None:
return info

name_to_alias: Dict[str, str] = {}
alias_to_name: Dict[str, str] = {}

if IS_PYDANTIC_V2:
for name, field_info in getattr(cls, "model_fields", {}).items(): # type: ignore[attr-defined]
alias = getattr(field_info, "alias", None) or name
name_to_alias[name] = alias
if alias != name:
alias_to_name[alias] = name
else:
for name, field in getattr(cls, "__fields__", {}).items():
alias = getattr(field, "alias", None) or name
name_to_alias[name] = alias
if alias != name:
alias_to_name[alias] = name

differing_names = {name for name, alias in name_to_alias.items() if alias != name}
ambiguous_keys = set(alias_to_name.keys()).intersection(name_to_alias.keys())
relevant_keys = differing_names | ambiguous_keys

info = _FieldAliasInfo(
name_to_alias=name_to_alias,
differing_names=differing_names,
ambiguous_keys=ambiguous_keys,
relevant_keys=relevant_keys,
)
_field_alias_info_cache[key] = info
return info


def parse_obj_as(type_: Type[T], object_: Any) -> T:
# convert_and_respect_annotation_metadata is required for TypedDict aliasing.
#
Expand All @@ -193,20 +246,7 @@ def parse_obj_as(type_: Type[T], object_: Any) -> T:
# - If the model encodes aliasing only via FieldMetadata annotations, then we MUST pre-dealias because Pydantic
# will not recognize those aliases during validation.
if inspect.isclass(type_) and issubclass(type_, pydantic.BaseModel):
has_pydantic_aliases = False
if IS_PYDANTIC_V2:
for field_name, field_info in getattr(type_, "model_fields", {}).items(): # type: ignore[attr-defined]
alias = getattr(field_info, "alias", None)
if alias is not None and alias != field_name:
has_pydantic_aliases = True
break
else:
for field in getattr(type_, "__fields__", {}).values():
alias = getattr(field, "alias", None)
name = getattr(field, "name", None)
if alias is not None and name is not None and alias != name:
has_pydantic_aliases = True
break
has_pydantic_aliases = bool(_get_field_alias_info(type_).differing_names)

dealiased_object = (
object_
Expand Down Expand Up @@ -246,19 +286,14 @@ def _coerce_field_names_to_aliases(cls, data: Any) -> Any:
if not isinstance(data, Mapping):
return data

fields = getattr(cls, "model_fields", {}) # type: ignore[attr-defined]
name_to_alias: Dict[str, str] = {}
alias_to_name: Dict[str, str] = {}
alias_info = _get_field_alias_info(cls)
if alias_info.relevant_keys.isdisjoint(data.keys()):
return data

for name, field_info in fields.items():
alias = getattr(field_info, "alias", None) or name
name_to_alias[name] = alias
if alias != name:
alias_to_name[alias] = name
name_to_alias = alias_info.name_to_alias

# Detect ambiguous keys: a key that is an alias for one field and a name for another.
ambiguous_keys = set(alias_to_name.keys()).intersection(set(name_to_alias.keys()))
for key in ambiguous_keys:
for key in alias_info.ambiguous_keys:
if key in data and name_to_alias[key] not in data:
raise ValueError(
f"Ambiguous input key '{key}': it is both a field name and an alias. "
Expand All @@ -267,9 +302,11 @@ def _coerce_field_names_to_aliases(cls, data: Any) -> Any:

original_keys = set(data.keys())
rewritten: Dict[str, Any] = dict(data)
for name, alias in name_to_alias.items():
if alias != name and name in original_keys and alias not in rewritten:
rewritten[alias] = rewritten.pop(name)
for name in alias_info.differing_names:
if name in original_keys:
alias = name_to_alias[name]
if alias not in rewritten:
rewritten[alias] = rewritten.pop(name)

return rewritten

Expand All @@ -293,18 +330,13 @@ def _coerce_field_names_to_aliases(cls, values: Any) -> Any:
if not isinstance(values, Mapping):
return values

fields = getattr(cls, "__fields__", {})
name_to_alias: Dict[str, str] = {}
alias_to_name: Dict[str, str] = {}
alias_info = _get_field_alias_info(cls)
if alias_info.relevant_keys.isdisjoint(values.keys()):
return values

for name, field in fields.items():
alias = getattr(field, "alias", None) or name
name_to_alias[name] = alias
if alias != name:
alias_to_name[alias] = name
name_to_alias = alias_info.name_to_alias

ambiguous_keys = set(alias_to_name.keys()).intersection(set(name_to_alias.keys()))
for key in ambiguous_keys:
for key in alias_info.ambiguous_keys:
if key in values and name_to_alias[key] not in values:
raise ValueError(
f"Ambiguous input key '{key}': it is both a field name and an alias. "
Expand All @@ -313,9 +345,11 @@ def _coerce_field_names_to_aliases(cls, values: Any) -> Any:

original_keys = set(values.keys())
rewritten: Dict[str, Any] = dict(values)
for name, alias in name_to_alias.items():
if alias != name and name in original_keys and alias not in rewritten:
rewritten[alias] = rewritten.pop(name)
for name in alias_info.differing_names:
if name in original_keys:
alias = name_to_alias[name]
if alias not in rewritten:
rewritten[alias] = rewritten.pop(name)

return rewritten

Expand Down