mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(bedrock): drop lookaround regex patterns from tool schemas for Converse models that reject them (#44138)
* feat(bedrock): drop lookaround regex patterns from tool schemas for Converse models that reject them * fix(bedrock): rename the lookaround flag to supports_regex_lookaround and keep dropped patternProperties names allowed The cost-map flag becomes a generic supports_regex_lookaround capability, which the cost-map schema admits as a supports_* boolean, and the Converse transform now owns the drop decision instead of the shared tools factory. A patternProperties key dropped from an object closed by additionalProperties: false leaves its value schema as that object's additionalProperties, so the names it allowed stay allowed, on the OpenAI non-Python-regex drop too. tool_with_sanitized_parameters also cleans Anthropic-shape tools (input_schema). * fix(router): keep a deployment's supports_regex_lookaround off the shared cost-map entry A deployment's model_info.supports_regex_lookaround was written to the shared bedrock/<model> cost-map key, so every sibling deployment of that model id inherited one deployment's choice. The flag now stays under the deployment's own id, which is what the Converse lookaround check reads first, and the shared entry keeps the cost map's value * test(bedrock): audit the Converse lookaround drop on the wire across endpoints, SDKs, flags and chaos --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
fe910889f7
commit
5dbe4f95e8
14 changed files with 1738 additions and 55 deletions
|
|
@ -1354,12 +1354,32 @@ def drop_non_python_regex_patterns(schema: Mapping[str, object]) -> Mapping[str,
|
|||
at more schema levels than a JSON parser admits, so a cyclic schema built in
|
||||
code cannot spin it.
|
||||
"""
|
||||
return _schema_without_rejected_regex(schema, _is_not_python_regex)
|
||||
|
||||
|
||||
def drop_lookaround_regex_patterns(schema: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""Drop every regex in a schema position that uses a lookaround assertion.
|
||||
|
||||
Some Bedrock Converse families compile tool schema regexes with an engine that
|
||||
has no lookahead or lookbehind and refuse the whole request over one. The ``(?=``,
|
||||
``(?!``, ``(?<=`` and ``(?<!`` openers are matched textually, so an escaped literal
|
||||
that spells one is dropped too, trading a hint for a request that goes through.
|
||||
A ``patternProperties`` key dropped from an object closed by ``additionalProperties:
|
||||
false`` leaves its value schema as that object's ``additionalProperties``, so the
|
||||
names it allowed stay allowed; :func:`drop_non_python_regex_patterns` shares the walk.
|
||||
"""
|
||||
return _schema_without_rejected_regex(schema, _uses_regex_lookaround)
|
||||
|
||||
|
||||
def _schema_without_rejected_regex(
|
||||
schema: Mapping[str, object], rejected: Callable[[str], bool]
|
||||
) -> Mapping[str, object]:
|
||||
rebuilt: dict[int, Mapping[str, object]] = {} # mutable-ok: per-call memo of rewritten nodes, deepest level first
|
||||
for level in reversed(tuple(islice(_schema_levels(schema), _MAX_SCHEMA_NESTING))):
|
||||
rebuilt.update(
|
||||
(id(node), rewritten)
|
||||
for node in level
|
||||
if (rewritten := _node_without_non_python_regex(node, rebuilt)) is not node
|
||||
if (rewritten := _node_without_rejected_regex(node, rebuilt, rejected)) is not node
|
||||
)
|
||||
return rebuilt.get(id(schema), schema)
|
||||
|
||||
|
|
@ -1381,23 +1401,56 @@ def _subschemas(node: Mapping[str, object]) -> Iterator[Mapping[str, object]]:
|
|||
yield value
|
||||
|
||||
|
||||
def _node_without_non_python_regex(
|
||||
node: Mapping[str, object], rebuilt: Mapping[int, Mapping[str, object]]
|
||||
def _node_without_rejected_regex(
|
||||
node: Mapping[str, object],
|
||||
rebuilt: Mapping[int, Mapping[str, object]],
|
||||
rejected: Callable[[str], bool],
|
||||
) -> Mapping[str, object]:
|
||||
kept: Final = {
|
||||
key: _keyword_value_rebuilt(key, value, rebuilt)
|
||||
key: _keyword_value_rebuilt(key, value, rebuilt, rejected)
|
||||
for key, value in node.items()
|
||||
if key != "pattern" or not isinstance(value, str) or _is_python_regex(value)
|
||||
if key != "pattern" or not isinstance(value, str) or not rejected(value)
|
||||
}
|
||||
return node if len(kept) == len(node) and all(kept[key] is node[key] for key in kept) else kept
|
||||
if len(kept) == len(node) and all(kept[key] is node[key] for key in kept):
|
||||
return node
|
||||
dropped_pattern_properties: Final = _dropped_pattern_properties(node, kept, rebuilt)
|
||||
if not dropped_pattern_properties or kept.get("additionalProperties") is not False:
|
||||
return kept
|
||||
return {**kept, "additionalProperties": _any_of(dropped_pattern_properties)}
|
||||
|
||||
|
||||
def _keyword_value_rebuilt(key: str, value: object, rebuilt: Mapping[int, Mapping[str, object]]) -> object:
|
||||
def _dropped_pattern_properties(
|
||||
node: Mapping[str, object],
|
||||
kept: Mapping[str, object],
|
||||
rebuilt: Mapping[int, Mapping[str, object]],
|
||||
) -> tuple[object, ...]:
|
||||
before: Final = _schema_at(node, "patternProperties")
|
||||
after: Final = _schema_at(kept, "patternProperties")
|
||||
if before is None or after is None:
|
||||
return ()
|
||||
return tuple(rebuilt.get(id(sub), sub) for name, sub in before.items() if name not in after)
|
||||
|
||||
|
||||
def _schema_at(container: Mapping[str, object], key: str) -> Mapping[str, object] | None:
|
||||
value: Final = container.get(key)
|
||||
return value if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _any_of(schemas: tuple[object, ...]) -> object:
|
||||
return schemas[0] if len(schemas) == 1 else {"anyOf": list(schemas)}
|
||||
|
||||
|
||||
def _keyword_value_rebuilt(
|
||||
key: str,
|
||||
value: object,
|
||||
rebuilt: Mapping[int, Mapping[str, object]],
|
||||
rejected: Callable[[str], bool],
|
||||
) -> object:
|
||||
if key in _SUBSCHEMA_MAP_KEYWORDS and isinstance(value, dict):
|
||||
kept: Final = {
|
||||
name: rebuilt.get(id(sub), sub)
|
||||
for name, sub in value.items()
|
||||
if key != "patternProperties" or not isinstance(name, str) or _is_python_regex(name)
|
||||
if key != "patternProperties" or not isinstance(name, str) or not rejected(name)
|
||||
}
|
||||
return value if len(kept) == len(value) and all(kept[name] is value[name] for name in kept) else kept
|
||||
if key in _SUBSCHEMA_LIST_KEYWORDS and isinstance(value, list):
|
||||
|
|
@ -1408,12 +1461,19 @@ def _keyword_value_rebuilt(key: str, value: object, rebuilt: Mapping[int, Mappin
|
|||
return value
|
||||
|
||||
|
||||
def _is_python_regex(pattern: str) -> bool:
|
||||
def _is_not_python_regex(pattern: str) -> bool:
|
||||
try:
|
||||
re.compile(pattern)
|
||||
except (re.error, RecursionError):
|
||||
return False
|
||||
return True
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
_REGEX_LOOKAROUND_RE: Final = re.compile(r"\(\?<?[=!]")
|
||||
|
||||
|
||||
def _uses_regex_lookaround(pattern: str) -> bool:
|
||||
return _REGEX_LOOKAROUND_RE.search(pattern) is not None
|
||||
|
||||
|
||||
def flatten_combinators_and_drop_non_python_regex_patterns(schema: Mapping[str, object]) -> Mapping[str, object]:
|
||||
|
|
@ -1424,16 +1484,23 @@ def tool_with_sanitized_parameters(
|
|||
tool: Mapping[str, object],
|
||||
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
|
||||
) -> Mapping[str, object]:
|
||||
function: Final = tool.get("function")
|
||||
if not isinstance(function, dict):
|
||||
"""Run the tool's JSON schema through ``sanitize``: ``function.parameters`` on an
|
||||
OpenAI tool, ``input_schema`` on an Anthropic one. The same object comes back when
|
||||
nothing changed."""
|
||||
function: Final = _schema_at(tool, "function")
|
||||
if function is not None:
|
||||
parameters: Final = _schema_at(function, "parameters")
|
||||
if parameters is None:
|
||||
return tool
|
||||
sanitized_parameters: Final = sanitize(parameters)
|
||||
if sanitized_parameters is parameters:
|
||||
return tool
|
||||
return {**tool, "function": {**function, "parameters": sanitized_parameters}}
|
||||
input_schema: Final = _schema_at(tool, "input_schema")
|
||||
if input_schema is None:
|
||||
return tool
|
||||
parameters: Final = function.get("parameters")
|
||||
if not isinstance(parameters, dict):
|
||||
return tool
|
||||
sanitized: Final = sanitize(parameters)
|
||||
if sanitized is parameters:
|
||||
return tool
|
||||
return {**tool, "function": {**function, "parameters": sanitized}}
|
||||
sanitized_schema: Final = sanitize(input_schema)
|
||||
return tool if sanitized_schema is input_schema else {**tool, "input_schema": sanitized_schema}
|
||||
|
||||
|
||||
def _get_image_mime_type_from_url(url: str) -> str | None:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from itertools import chain
|
|||
from typing import TYPE_CHECKING, Final, Literal, cast, overload
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -28,6 +29,8 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
_parse_content_for_reasoning,
|
||||
drop_lookaround_regex_patterns,
|
||||
tool_with_sanitized_parameters,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
BedrockConverseMessagesProcessor,
|
||||
|
|
@ -49,6 +52,7 @@ from litellm.llms.anthropic.chat.transformation import (
|
|||
)
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.llms.bedrock.common_utils import bedrock_model_supports_regex_lookaround
|
||||
from litellm.llms.bedrock.request_metadata import (
|
||||
bedrock_request_metadata_headers,
|
||||
bedrock_request_metadata_is_owned,
|
||||
|
|
@ -128,6 +132,17 @@ UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS: Final = [
|
|||
]
|
||||
|
||||
|
||||
_TOOLS_AS_SENT: Final = TypeAdapter(tuple[Mapping[str, object], ...])
|
||||
|
||||
|
||||
def _tools_the_model_accepts(
|
||||
tools: Sequence[Mapping[str, object]], model: str, litellm_params: Mapping[str, object] | None
|
||||
) -> list[Mapping[str, object]]:
|
||||
if bedrock_model_supports_regex_lookaround(model, litellm_params):
|
||||
return list(tools)
|
||||
return [tool_with_sanitized_parameters(tool, drop_lookaround_regex_patterns) for tool in tools]
|
||||
|
||||
|
||||
class AmazonConverseConfig(BaseConfig):
|
||||
"""
|
||||
Reference - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_Converse.html
|
||||
|
|
@ -1689,6 +1704,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
model: str,
|
||||
headers: dict | None,
|
||||
additional_request_params: dict,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
) -> tuple[list[ToolBlock], list]:
|
||||
"""Process tools and collect anthropic_beta values."""
|
||||
bedrock_tools: list[ToolBlock] = []
|
||||
|
|
@ -1729,7 +1745,9 @@ class AmazonConverseConfig(BaseConfig):
|
|||
computer_use_tools, regular_tools = self._separate_computer_use_tools(filtered_tools, model)
|
||||
|
||||
# Process regular function tools using existing logic
|
||||
bedrock_tools = _bedrock_tools_pt(regular_tools, model=model)
|
||||
bedrock_tools = _bedrock_tools_pt(
|
||||
_tools_the_model_accepts(regular_tools, model, litellm_params), model=model
|
||||
)
|
||||
|
||||
# Add computer use tools and anthropic_beta if needed (only when computer use tools are present)
|
||||
if computer_use_tools:
|
||||
|
|
@ -1793,7 +1811,10 @@ class AmazonConverseConfig(BaseConfig):
|
|||
additional_request_params["tools"] = transformed_computer_tools
|
||||
else:
|
||||
# No computer use tools, process all tools as regular tools
|
||||
bedrock_tools = _bedrock_tools_pt(filtered_tools, model=model)
|
||||
bedrock_tools = _bedrock_tools_pt(
|
||||
_tools_the_model_accepts(_TOOLS_AS_SENT.validate_python(filtered_tools), model, litellm_params),
|
||||
model=model,
|
||||
)
|
||||
|
||||
# Append pre-formatted tools (systemTool etc.) after transformation
|
||||
bedrock_tools.extend(pre_formatted_tools)
|
||||
|
|
@ -1905,7 +1926,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
# Process tools and collect beta values
|
||||
bedrock_tools, anthropic_beta_list = self._process_tools_and_beta(
|
||||
original_tools, model, headers, additional_request_params
|
||||
original_tools, model, headers, additional_request_params, litellm_params
|
||||
)
|
||||
|
||||
# Append cachePoint to tools if cache_control_injection_points has tool_config
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import os
|
|||
import re
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
|
@ -972,6 +972,7 @@ def is_claude_4_5_on_bedrock(model: str) -> bool:
|
|||
|
||||
|
||||
_BEDROCK_MODEL_VERSION_SUFFIX_RE: Final = re.compile(r"-v\d+(?::\d+)?$")
|
||||
_DEPLOYMENT_MODEL_INFO: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def bedrock_converse_supports_strict_tools(model: str) -> bool:
|
||||
|
|
@ -989,12 +990,38 @@ def bedrock_converse_supports_strict_tools(model: str) -> bool:
|
|||
base: Final = get_bedrock_base_model(model)
|
||||
if not base.startswith("anthropic"):
|
||||
return False
|
||||
flag: Final = _get_bedrock_converse_strict_tools_flag(base)
|
||||
flag: Final = _bedrock_converse_model_flag(base, "bedrock_converse_supports_strict_tools")
|
||||
return flag if flag is not None else True
|
||||
|
||||
|
||||
def _get_bedrock_converse_strict_tools_flag(base_model: str) -> bool | None:
|
||||
candidates: Final = dict.fromkeys((base_model, _BEDROCK_MODEL_VERSION_SUFFIX_RE.sub("", base_model)))
|
||||
def bedrock_model_supports_regex_lookaround(model: str, litellm_params: Mapping[str, object] | None = None) -> bool:
|
||||
"""
|
||||
Whether ``model`` accepts lookahead and lookbehind assertions in tool schema regexes.
|
||||
|
||||
The deployment's ``model_info.supports_regex_lookaround`` wins, then the
|
||||
``model_prices_and_context_window.json`` entry of its ``base_model``, then the
|
||||
entry of ``model`` itself. A model nobody flagged keeps its schema as sent.
|
||||
"""
|
||||
params: Final = litellm_params or {}
|
||||
model_info: Final = _DEPLOYMENT_MODEL_INFO.validate_python(params.get("model_info") or {})
|
||||
deployment_flag: Final = model_info.get("supports_regex_lookaround")
|
||||
if isinstance(deployment_flag, bool):
|
||||
return deployment_flag
|
||||
base_model: Final = params.get("base_model")
|
||||
candidates: Final = (*((base_model,) if isinstance(base_model, str) else ()), model)
|
||||
flags: Final = (_bedrock_converse_model_flag(candidate, "supports_regex_lookaround") for candidate in candidates)
|
||||
return next((flag for flag in flags if flag is not None), True)
|
||||
|
||||
|
||||
_BedrockConverseModelFlag: TypeAlias = Literal[
|
||||
"bedrock_converse_supports_strict_tools",
|
||||
"supports_regex_lookaround",
|
||||
]
|
||||
|
||||
|
||||
def _bedrock_converse_model_flag(model: str, key: _BedrockConverseModelFlag) -> bool | None:
|
||||
base: Final = get_bedrock_base_model(model)
|
||||
candidates: Final = dict.fromkeys((model, base, _BEDROCK_MODEL_VERSION_SUFFIX_RE.sub("", base)))
|
||||
for candidate in candidates:
|
||||
with contextlib.suppress(Exception):
|
||||
model_info = get_cached_model_info()(
|
||||
|
|
@ -1002,15 +1029,13 @@ def _get_bedrock_converse_strict_tools_flag(base_model: str) -> bool | None:
|
|||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
flag = model_info.get("bedrock_converse_supports_strict_tools")
|
||||
flag = model_info.get(key)
|
||||
if isinstance(flag, bool):
|
||||
return flag
|
||||
|
||||
model_cost_key = model_info.get("key")
|
||||
if isinstance(model_cost_key, str):
|
||||
local_flag = (
|
||||
_get_local_model_cost_map().get(model_cost_key, {}).get("bedrock_converse_supports_strict_tools")
|
||||
)
|
||||
local_flag = _get_local_model_cost_map().get(model_cost_key, {}).get(key)
|
||||
if isinstance(local_flag, bool):
|
||||
return local_flag
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -47460,6 +47460,7 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"us-gov.xai.grok-4.6": {
|
||||
"supports_regex_lookaround": false,
|
||||
"input_cost_per_token": 2.64e-06,
|
||||
"output_cost_per_token": 7.92e-06,
|
||||
"cache_read_input_token_cost": 6.6e-07,
|
||||
|
|
@ -59073,6 +59074,7 @@
|
|||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html"
|
||||
},
|
||||
"us.xai.grok-4.6": {
|
||||
"supports_regex_lookaround": false,
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
"output_cost_per_token": 6.6e-06,
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
|
|
@ -59089,6 +59091,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"global.xai.grok-4.6": {
|
||||
"supports_regex_lookaround": false,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"output_cost_per_token": 6e-06,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
|
|
@ -76781,6 +76784,7 @@
|
|||
"supports_web_search": true
|
||||
},
|
||||
"moonshotai.kimi-k3": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
|
|
@ -76801,6 +76805,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"global.moonshotai.kimi-k3": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 3e-06,
|
||||
|
|
@ -76821,6 +76826,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"us.moonshotai.kimi-k3": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
|
|
@ -79377,6 +79383,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"global.xai.grok-4.7": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
@ -79393,6 +79400,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"us.xai.grok-4.7": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
@ -79409,6 +79417,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"xai.grok-4.7": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
|
|||
|
|
@ -211,6 +211,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
|
|||
vertex_ai_audio_api: ReadOnly[Literal["lyria_predict", "lyria_interactions"] | None]
|
||||
bedrock_output_config_effort_ceiling: Literal["low", "medium", "high", "max", "xhigh"] | None
|
||||
bedrock_converse_supports_strict_tools: bool | None
|
||||
supports_regex_lookaround: ReadOnly[bool | None]
|
||||
|
||||
|
||||
class SearchContextCostPerQuery(TypedDict, total=False):
|
||||
|
|
@ -3844,10 +3845,13 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
|
||||
DEPLOYMENT_SCOPED_PRICING_FIELDS: Final[frozenset[str]] = frozenset({"off_peak_pricing"})
|
||||
|
||||
DEPLOYMENT_SCOPED_CAPABILITY_FIELDS: Final[frozenset[str]] = frozenset({"supports_regex_lookaround"})
|
||||
|
||||
SHARED_BACKEND_MODEL_INFO_FIELDS: Final[frozenset[str]] = (
|
||||
frozenset(ModelInfoBase.__required_keys__ | ModelInfoBase.__optional_keys__)
|
||||
- frozenset(CustomPricingLiteLLMParams.model_fields)
|
||||
- DEPLOYMENT_SCOPED_PRICING_FIELDS
|
||||
- DEPLOYMENT_SCOPED_CAPABILITY_FIELDS
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -6321,6 +6321,7 @@ def _get_model_info_helper(
|
|||
default_reasoning_effort=_model_info.get("default_reasoning_effort", None),
|
||||
bedrock_output_config_effort_ceiling=_model_info.get("bedrock_output_config_effort_ceiling", None),
|
||||
bedrock_converse_supports_strict_tools=_model_info.get("bedrock_converse_supports_strict_tools", None),
|
||||
supports_regex_lookaround=_model_info.get("supports_regex_lookaround", None),
|
||||
supports_computer_use=_model_info.get("supports_computer_use", None),
|
||||
search_context_cost_per_query=_model_info.get("search_context_cost_per_query", None),
|
||||
web_search_billing_unit=_model_info.get("web_search_billing_unit", None),
|
||||
|
|
|
|||
|
|
@ -47460,6 +47460,7 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"us-gov.xai.grok-4.6": {
|
||||
"supports_regex_lookaround": false,
|
||||
"input_cost_per_token": 2.64e-06,
|
||||
"output_cost_per_token": 7.92e-06,
|
||||
"cache_read_input_token_cost": 6.6e-07,
|
||||
|
|
@ -59073,6 +59074,7 @@
|
|||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html"
|
||||
},
|
||||
"us.xai.grok-4.6": {
|
||||
"supports_regex_lookaround": false,
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
"output_cost_per_token": 6.6e-06,
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
|
|
@ -59089,6 +59091,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"global.xai.grok-4.6": {
|
||||
"supports_regex_lookaround": false,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"output_cost_per_token": 6e-06,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
|
|
@ -76781,6 +76784,7 @@
|
|||
"supports_web_search": true
|
||||
},
|
||||
"moonshotai.kimi-k3": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
|
|
@ -76801,6 +76805,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"global.moonshotai.kimi-k3": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_creation_input_token_cost": 3.75e-06,
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 3e-06,
|
||||
|
|
@ -76821,6 +76826,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"us.moonshotai.kimi-k3": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_creation_input_token_cost": 4.125e-06,
|
||||
"cache_read_input_token_cost": 3.3e-07,
|
||||
"input_cost_per_token": 3.3e-06,
|
||||
|
|
@ -79377,6 +79383,7 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"global.xai.grok-4.7": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
@ -79393,6 +79400,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"us.xai.grok-4.7": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
@ -79409,6 +79417,7 @@
|
|||
"supports_vision": true
|
||||
},
|
||||
"xai.grok-4.7": {
|
||||
"supports_regex_lookaround": false,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
|
|
|
|||
|
|
@ -1062,6 +1062,9 @@
|
|||
"supports_reasoning": {
|
||||
"type": "boolean"
|
||||
},
|
||||
"supports_regex_lookaround": {
|
||||
"type": "boolean"
|
||||
},
|
||||
"supports_response_schema": {
|
||||
"type": "boolean"
|
||||
},
|
||||
|
|
|
|||
|
|
@ -0,0 +1,379 @@
|
|||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import unquote, urlsplit
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.upstream import _aws_event_frame
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_KIMI: Final = "global.moonshotai.kimi-k3"
|
||||
_NOVA: Final = "us.amazon.nova-lite-v1:0"
|
||||
_AWS: Final[dict[str, JsonValue]] = {
|
||||
"aws_access_key_id": "AKIASCRIPTEDPROVIDER",
|
||||
"aws_secret_access_key": "scripted-secret",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
_LOOKAHEAD: Final = r"^(?!\.\.?(?:\/|$))[A-Za-z0-9_\-.~:@+]{1,200}$"
|
||||
_PLAIN: Final = r"^[a-z][a-z0-9_]*$"
|
||||
_TOOL: Final = "ArtifactData"
|
||||
_EVENT_STREAM: Final = "application/vnd.amazon.eventstream"
|
||||
_JSON: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_LIST: Final = TypeAdapter(list[JsonValue])
|
||||
_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})")
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_USAGE: Final[dict[str, JsonValue]] = {"inputTokens": 21, "outputTokens": 7, "totalTokens": 28}
|
||||
_WIRE_AS_SENT: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {"type": "string", "pattern": _LOOKAHEAD},
|
||||
"doc_id": {"type": "string", "pattern": _PLAIN},
|
||||
},
|
||||
"required": ["collection"],
|
||||
}
|
||||
_WIRE_LOOKAROUND_FREE: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {"collection": {"type": "string"}, "doc_id": {"type": "string", "pattern": _PLAIN}},
|
||||
"required": ["collection"],
|
||||
}
|
||||
_SCHEMA_AS_SENT: Final[dict[str, JsonValue]] = {**_WIRE_AS_SENT, "additionalProperties": False}
|
||||
|
||||
Endpoint = Literal["chat", "messages", "responses"]
|
||||
_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Fleet:
|
||||
kimi_bare: str
|
||||
kimi_flagged_true: str
|
||||
nova_off: str
|
||||
nova_bare: str
|
||||
|
||||
def names(self) -> tuple[str, ...]:
|
||||
return (self.kimi_bare, self.kimi_flagged_true, self.nova_off, self.nova_bare)
|
||||
|
||||
def expected_schema(self, model: str) -> dict[str, JsonValue]:
|
||||
return _WIRE_LOOKAROUND_FREE if model in (self.kimi_bare, self.nova_off) else _WIRE_AS_SENT
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Call:
|
||||
model: str
|
||||
endpoint: Endpoint
|
||||
stream: bool
|
||||
marker: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
call: _Call
|
||||
status: int
|
||||
text: str
|
||||
|
||||
|
||||
def _answer(marker: str) -> str:
|
||||
return f"answer marker-{marker}"
|
||||
|
||||
|
||||
def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes:
|
||||
return _aws_event_frame(event_type, payload, "sc", "u")
|
||||
|
||||
|
||||
def _stream_frames(marker: str) -> tuple[bytes, ...]:
|
||||
return (
|
||||
_frame("messageStart", {"role": "assistant"}),
|
||||
_frame("contentBlockDelta", {"delta": {"text": "answer "}, "contentBlockIndex": 0}),
|
||||
_frame("contentBlockDelta", {"delta": {"text": f"marker-{marker}"}, "contentBlockIndex": 0}),
|
||||
_frame("contentBlockStop", {"contentBlockIndex": 0}),
|
||||
_frame("messageStop", {"stopReason": "end_turn"}),
|
||||
_frame("metadata", {"usage": _USAGE}),
|
||||
)
|
||||
|
||||
|
||||
def _text_reply(marker: str, stream: bool, abort_after: int | None = None) -> Reply:
|
||||
if stream:
|
||||
return Reply(content_type=_EVENT_STREAM, chunks=_stream_frames(marker), abort_after=abort_after)
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"output": {"message": {"role": "assistant", "content": [{"text": _answer(marker)}]}},
|
||||
"stopReason": "end_turn",
|
||||
"usage": _USAGE,
|
||||
"metrics": {"latencyMs": 1},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _marker_of(request: Request) -> str:
|
||||
found: Final = _MARKER.search(request.body.decode())
|
||||
assert found is not None, request.body
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def _is_stream(request: Request) -> bool:
|
||||
return unquote(request.target).endswith("/converse-stream")
|
||||
|
||||
|
||||
def _echo(request: Request) -> Reply:
|
||||
return _text_reply(_marker_of(request), _is_stream(request))
|
||||
|
||||
|
||||
def _path(endpoint: Endpoint) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return "/v1/chat/completions"
|
||||
case "messages":
|
||||
return "/v1/messages"
|
||||
case "responses":
|
||||
return "/v1/responses"
|
||||
|
||||
|
||||
def _body(call: _Call) -> dict[str, JsonValue]:
|
||||
question: Final = f"Question marker-{call.marker}"
|
||||
common: Final[dict[str, JsonValue]] = {
|
||||
"model": call.model,
|
||||
"stream": call.stream,
|
||||
"num_retries": 0,
|
||||
"cache": {"no-cache": True},
|
||||
}
|
||||
tool: Final[dict[str, JsonValue]] = {"description": f"{_TOOL} tool"}
|
||||
match call.endpoint:
|
||||
case "chat":
|
||||
return {
|
||||
**common,
|
||||
"messages": [{"role": "user", "content": question}],
|
||||
"max_tokens": 64,
|
||||
"tools": [{"type": "function", "function": {"name": _TOOL, **tool, "parameters": _SCHEMA_AS_SENT}}],
|
||||
}
|
||||
case "messages":
|
||||
return {
|
||||
**common,
|
||||
"messages": [{"role": "user", "content": question}],
|
||||
"max_tokens": 64,
|
||||
"tools": [{"name": _TOOL, **tool, "input_schema": _SCHEMA_AS_SENT}],
|
||||
}
|
||||
case "responses":
|
||||
return {
|
||||
**common,
|
||||
"input": question,
|
||||
"max_output_tokens": 64,
|
||||
"tools": [{"type": "function", "name": _TOOL, **tool, "parameters": _SCHEMA_AS_SENT}],
|
||||
}
|
||||
|
||||
|
||||
def _received_schema(request: Request) -> dict[str, JsonValue]:
|
||||
body: Final = _JSON.validate_json(request.body)
|
||||
(tool,) = _LIST.validate_python(object_value(body["toolConfig"])["tools"])
|
||||
spec: Final = object_value(object_value(tool)["toolSpec"])
|
||||
assert spec["name"] == _TOOL, spec
|
||||
return object_value(object_value(spec["inputSchema"])["json"])
|
||||
|
||||
|
||||
def _assert_schemas_by_marker(received: tuple[Request, ...], calls: tuple[_Call, ...], fleet: _Fleet) -> None:
|
||||
by_marker: Final = MappingProxyType({call.marker: call for call in calls})
|
||||
assert sorted(_marker_of(request) for request in received) == sorted(by_marker), len(received)
|
||||
for request in received:
|
||||
call: Final = by_marker[_marker_of(request)]
|
||||
assert _is_stream(request) == call.stream, (call, request.target)
|
||||
assert _received_schema(request) == fleet.expected_schema(call.model), (call, request.body)
|
||||
|
||||
|
||||
def _spend_statuses(model: str, expected: int) -> list[JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
lambda found: len(found) >= expected,
|
||||
seconds=60,
|
||||
)
|
||||
assert len({row["request_id"] for row in rows}) == len(rows), rows
|
||||
return [row["status"] for row in rows]
|
||||
|
||||
|
||||
async def _send(client: httpx.AsyncClient, key: str, call: _Call) -> _Served:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
_path(call.endpoint),
|
||||
json=_body(call),
|
||||
headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"},
|
||||
) as response:
|
||||
raw: Final = await response.aread()
|
||||
return _Served(call=call, status=response.status_code, text=raw.decode())
|
||||
|
||||
|
||||
async def _burst(
|
||||
base_url: str, key: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
|
||||
) -> tuple[_Served, ...]:
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client:
|
||||
results: Final = await asyncio.gather(
|
||||
*(_send(client, key, call) for call in calls), return_exceptions=tolerate_transport_errors
|
||||
)
|
||||
for result in results:
|
||||
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
|
||||
return tuple(result for result in results if isinstance(result, _Served))
|
||||
|
||||
|
||||
def _mixed_calls(fleet: _Fleet, count: int) -> tuple[_Call, ...]:
|
||||
names: Final = fleet.names()
|
||||
return tuple(
|
||||
_Call(
|
||||
model=names[index % len(names)],
|
||||
endpoint=_ENDPOINTS[(index // len(names)) % len(_ENDPOINTS)],
|
||||
stream=(index // (len(names) * len(_ENDPOINTS))) % 2 == 0,
|
||||
marker=uuid.uuid4().hex,
|
||||
)
|
||||
for index in range(count)
|
||||
)
|
||||
|
||||
|
||||
def _assert_answered_with_its_own_marker(served: _Served) -> None:
|
||||
assert served.status == 200, served.text
|
||||
assert set(_MARKER.findall(served.text)) == {served.call.marker}, served.text
|
||||
|
||||
|
||||
def _fleet_config(wire: Wire, tmp_path: Path) -> tuple[Path, _Fleet]:
|
||||
run_id: Final = uuid.uuid4().hex[:8]
|
||||
fleet: Final = _Fleet(
|
||||
kimi_bare=f"kimi-bare-{run_id}",
|
||||
kimi_flagged_true=f"kimi-flagged-true-{run_id}",
|
||||
nova_off=f"nova-off-{run_id}",
|
||||
nova_bare=f"nova-bare-{run_id}",
|
||||
)
|
||||
params: Final[dict[str, JsonValue]] = {"api_base": wire.url, **_AWS}
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{"model_name": fleet.kimi_bare, "litellm_params": {"model": f"bedrock/{_KIMI}", **params}},
|
||||
{
|
||||
"model_name": fleet.kimi_flagged_true,
|
||||
"litellm_params": {"model": f"bedrock/{_KIMI}", **params},
|
||||
"model_info": {"supports_regex_lookaround": True},
|
||||
},
|
||||
{
|
||||
"model_name": fleet.nova_off,
|
||||
"litellm_params": {"model": f"bedrock/converse/{_NOVA}", **params},
|
||||
"model_info": {"supports_regex_lookaround": False},
|
||||
},
|
||||
{"model_name": fleet.nova_bare, "litellm_params": {"model": f"bedrock/converse/{_NOVA}", **params}},
|
||||
]
|
||||
path: Final = tmp_path / "bedrock-lookaround-chaos.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path, fleet
|
||||
|
||||
|
||||
def _open_upstream_connections(pid: int, upstream: str) -> int:
|
||||
port: Final = urlsplit(upstream).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(600)
|
||||
async def test_a_mixed_burst_across_two_workers_cleans_only_the_flagged_deployments(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
with wire_server(_echo) as wire:
|
||||
path, fleet = _fleet_config(wire, tmp_path)
|
||||
calls: Final = _mixed_calls(fleet, 36)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
served: Final = await _burst(str(candidate.client.base_url), candidate.key, calls)
|
||||
assert len(served) == 36
|
||||
for item in served:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
_assert_schemas_by_marker(wire.drain(), calls, fleet)
|
||||
for name in fleet.names():
|
||||
assert _spend_statuses(name, 9) == ["success"] * 9
|
||||
|
||||
|
||||
@pytest.mark.timeout(600)
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_cleaning_schemas(gateway: Gateway, tmp_path: Path) -> None:
|
||||
release: Final = threading.Event()
|
||||
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
def held(request: Request) -> Reply:
|
||||
held_markers.put(_marker_of(request))
|
||||
assert release.wait(timeout=60), "The burst was never released"
|
||||
return _echo(request)
|
||||
|
||||
with wire_server(held) as wire:
|
||||
path, fleet = _fleet_config(wire, tmp_path)
|
||||
calls: Final = tuple(
|
||||
_Call(model=fleet.kimi_bare, endpoint="chat", stream=False, marker=uuid.uuid4().hex) for _ in range(20)
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
workers: Final = eventually(
|
||||
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
burst: Final = asyncio.create_task(
|
||||
_burst(str(candidate.client.base_url), candidate.key, calls, tolerate_transport_errors=True)
|
||||
)
|
||||
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60)
|
||||
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers})
|
||||
assert sum(held_by.values()) == 20, held_by
|
||||
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
|
||||
victim: Final = psutil.Process(victim_pid)
|
||||
victim.suspend()
|
||||
victim.send_signal(signal.SIGKILL)
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
assert held_by[survivor_pid] >= 10, held_by
|
||||
assert len(served) == held_by[survivor_pid], (held_by, len(served))
|
||||
for item in served:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
follow_up: Final = _Call(model=fleet.kimi_bare, endpoint="chat", stream=False, marker=uuid.uuid4().hex)
|
||||
(answered,) = await _burst(str(candidate.client.base_url), candidate.key, (follow_up,))
|
||||
_assert_answered_with_its_own_marker(answered)
|
||||
_assert_schemas_by_marker(wire.drain(), (*calls, follow_up), fleet)
|
||||
|
||||
|
||||
@pytest.mark.timeout(600)
|
||||
async def test_peer_stream_aborts_reach_callers_while_the_rest_of_the_burst_is_cleaned(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
markers: Final = tuple(uuid.uuid4().hex for _ in range(12))
|
||||
aborted: Final = frozenset(marker for index, marker in enumerate(markers) if index % 3 == 0)
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
marker: Final = _marker_of(request)
|
||||
return _text_reply(marker, stream=True, abort_after=0 if marker in aborted else None)
|
||||
|
||||
with wire_server(respond) as wire:
|
||||
path, fleet = _fleet_config(wire, tmp_path)
|
||||
calls: Final = tuple(
|
||||
_Call(model=fleet.kimi_bare, endpoint=_ENDPOINTS[index % 3], stream=True, marker=marker)
|
||||
for index, marker in enumerate(markers)
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
served: Final = await _burst(str(candidate.client.base_url), candidate.key, calls)
|
||||
assert len(served) == 12
|
||||
for item in served:
|
||||
if item.call.marker in aborted:
|
||||
assert "marker-" not in item.text, item.text
|
||||
assert item.status >= 500 or "error" in item.text.lower(), (item.status, item.text)
|
||||
else:
|
||||
_assert_answered_with_its_own_marker(item)
|
||||
recovery: Final = _Call(model=fleet.kimi_bare, endpoint="chat", stream=True, marker=uuid.uuid4().hex)
|
||||
(recovered,) = await _burst(str(candidate.client.base_url), candidate.key, (recovery,))
|
||||
_assert_answered_with_its_own_marker(recovered)
|
||||
_assert_schemas_by_marker(wire.drain(), (*calls, recovery), fleet)
|
||||
|
|
@ -0,0 +1,833 @@
|
|||
import json
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import unquote
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.upstream import _aws_event_frame
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_KIMI: Final = "global.moonshotai.kimi-k3"
|
||||
_GROK: Final = "us.xai.grok-4.7"
|
||||
_NOVA: Final = "us.amazon.nova-lite-v1:0"
|
||||
_CLAUDE: Final = "global.anthropic.claude-opus-4-8"
|
||||
_PROFILE_ARN: Final = "arn:aws:bedrock:us-east-1:000000000000:application-inference-profile/lookaround0"
|
||||
_AWS: Final[dict[str, JsonValue]] = {
|
||||
"aws_access_key_id": "AKIASCRIPTEDPROVIDER",
|
||||
"aws_secret_access_key": "scripted-secret",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
_NO_CACHE: Final[dict[str, JsonValue]] = {"cache": {"no-cache": True}, "num_retries": 0}
|
||||
_LOOKAHEAD: Final = r"^(?!\.\.?(?:\/|$))[A-Za-z0-9_\-.~:@+]{1,200}$"
|
||||
_NEGATIVE_LOOKBEHIND: Final = r"^(?<!tmp_)[a-z]+$"
|
||||
_POSITIVE_LOOKBEHIND: Final = r"(?<=v)[0-9]+"
|
||||
_POSITIVE_LOOKAHEAD_KEY: Final = r"^x_(?=[a-z])"
|
||||
_PLAIN: Final = r"^[a-z][a-z0-9_]*$"
|
||||
_TOOL: Final = "ArtifactData"
|
||||
_PLAIN_TOOL: Final = "ListNotes"
|
||||
_PROMPT: Final = "Read the notes document from the notes collection."
|
||||
_ANSWER: Final = "lookaround regex control answer"
|
||||
_TOOL_INPUT: Final[dict[str, JsonValue]] = {"collection": "notes", "doc_id": "notes"}
|
||||
_EVENT_STREAM: Final = "application/vnd.amazon.eventstream"
|
||||
_BEDROCK_REJECTION: Final = "structured output schema uses unsupported regex negative look-ahead"
|
||||
_USAGE: Final[dict[str, JsonValue]] = {"inputTokens": 21, "outputTokens": 7, "totalTokens": 28}
|
||||
_JSON: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_LIST: Final = TypeAdapter(list[JsonValue])
|
||||
_SCHEMA_AS_SENT: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {"type": "string", "description": "Collection name", "pattern": _LOOKAHEAD},
|
||||
"doc_id": {"type": "string", "description": "Document id", "pattern": _PLAIN},
|
||||
"filters": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"anyOf": [
|
||||
{"type": "string", "pattern": _NEGATIVE_LOOKBEHIND},
|
||||
{"type": "string", "pattern": _POSITIVE_LOOKBEHIND},
|
||||
]
|
||||
},
|
||||
},
|
||||
"labels": {
|
||||
"type": "object",
|
||||
"patternProperties": {_POSITIVE_LOOKAHEAD_KEY: {"type": "string"}, "^v_": {"type": "integer"}},
|
||||
"additionalProperties": False,
|
||||
},
|
||||
"meta": {"type": "object", "default": {"pattern": _LOOKAHEAD}},
|
||||
},
|
||||
"required": ["collection"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
_SCHEMA_LOOKAROUND_FREE: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {"type": "string", "description": "Collection name"},
|
||||
"doc_id": {"type": "string", "description": "Document id", "pattern": _PLAIN},
|
||||
"filters": {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "string"}]}},
|
||||
"labels": {
|
||||
"type": "object",
|
||||
"patternProperties": {"^v_": {"type": "integer"}},
|
||||
"additionalProperties": {"type": "string"},
|
||||
},
|
||||
"meta": {"type": "object", "default": {"pattern": _LOOKAHEAD}},
|
||||
},
|
||||
"required": ["collection"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
_PLAIN_SCHEMA: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {"limit": {"type": "integer", "minimum": 1}, "prefix": {"type": "string", "pattern": _PLAIN}},
|
||||
"required": ["limit"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
|
||||
def _converse_root(schema: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"type": schema["type"],
|
||||
"properties": schema.get("properties", {}),
|
||||
"required": schema.get("required", []),
|
||||
}
|
||||
|
||||
|
||||
_WIRE_AS_SENT: Final = _converse_root(_SCHEMA_AS_SENT)
|
||||
_WIRE_LOOKAROUND_FREE: Final = _converse_root(_SCHEMA_LOOKAROUND_FREE)
|
||||
_WIRE_PLAIN: Final = _converse_root(_PLAIN_SCHEMA)
|
||||
|
||||
Endpoint = Literal["chat", "messages", "responses"]
|
||||
_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses")
|
||||
|
||||
|
||||
def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes:
|
||||
return _aws_event_frame(event_type, payload, "sc", "u")
|
||||
|
||||
|
||||
_TOOL_USE_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [{"toolUse": {"toolUseId": "tooluse_lookaround_1", "name": _TOOL, "input": _TOOL_INPUT}}],
|
||||
}
|
||||
},
|
||||
"stopReason": "tool_use",
|
||||
"usage": _USAGE,
|
||||
"metrics": {"latencyMs": 1},
|
||||
}
|
||||
).encode()
|
||||
_STREAM_FRAMES: Final = b"".join(
|
||||
(
|
||||
_frame("messageStart", {"role": "assistant"}),
|
||||
_frame("contentBlockDelta", {"delta": {"text": _ANSWER}, "contentBlockIndex": 0}),
|
||||
_frame("contentBlockStop", {"contentBlockIndex": 0}),
|
||||
_frame("messageStop", {"stopReason": "end_turn"}),
|
||||
_frame("metadata", {"usage": _USAGE}),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_peer(request: Request) -> Reply:
|
||||
if unquote(request.target).endswith("/converse-stream"):
|
||||
return Reply(body=_STREAM_FRAMES, content_type=_EVENT_STREAM)
|
||||
return Reply(body=_TOOL_USE_RESPONSE)
|
||||
|
||||
|
||||
def _rejecting_peer(request: Request) -> Reply:
|
||||
return Reply(status=400, body=json.dumps({"message": _BEDROCK_REJECTION}).encode())
|
||||
|
||||
|
||||
def _openai_tool(name: str, schema: Mapping[str, JsonValue], **extra: JsonValue) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {"name": name, "description": f"{name} tool", "parameters": dict(schema), **extra},
|
||||
}
|
||||
|
||||
|
||||
def _anthropic_tool(name: str, schema: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {"name": name, "description": f"{name} tool", "input_schema": dict(schema)}
|
||||
|
||||
|
||||
def _responses_tool(name: str, schema: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {"type": "function", "name": name, "description": f"{name} tool", "parameters": dict(schema)}
|
||||
|
||||
|
||||
def _tool_for(endpoint: Endpoint, name: str, schema: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return _openai_tool(name, schema)
|
||||
case "messages":
|
||||
return _anthropic_tool(name, schema)
|
||||
case "responses":
|
||||
return _responses_tool(name, schema)
|
||||
|
||||
|
||||
def _path(endpoint: Endpoint) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return "/v1/chat/completions"
|
||||
case "messages":
|
||||
return "/v1/messages"
|
||||
case "responses":
|
||||
return "/v1/responses"
|
||||
|
||||
|
||||
def _body(
|
||||
endpoint: Endpoint,
|
||||
model: str,
|
||||
tools: Sequence[Mapping[str, JsonValue]],
|
||||
*,
|
||||
stream: bool = False,
|
||||
**extra: JsonValue,
|
||||
) -> dict[str, JsonValue]:
|
||||
tool_list: Final[list[JsonValue]] = [dict(tool) for tool in tools]
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _PROMPT}],
|
||||
"max_tokens": 64,
|
||||
"stream": stream,
|
||||
"tools": tool_list,
|
||||
**_NO_CACHE,
|
||||
**extra,
|
||||
}
|
||||
case "messages":
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _PROMPT}],
|
||||
"max_tokens": 64,
|
||||
"stream": stream,
|
||||
"tools": tool_list,
|
||||
**_NO_CACHE,
|
||||
**extra,
|
||||
}
|
||||
case "responses":
|
||||
return {
|
||||
"model": model,
|
||||
"input": _PROMPT,
|
||||
"max_output_tokens": 64,
|
||||
"stream": stream,
|
||||
"tools": tool_list,
|
||||
**_NO_CACHE,
|
||||
**extra,
|
||||
}
|
||||
|
||||
|
||||
def _deployment(
|
||||
scenario: Scenario,
|
||||
wire: Wire,
|
||||
model: str,
|
||||
*,
|
||||
model_info: Mapping[str, JsonValue] | None = None,
|
||||
**params: JsonValue,
|
||||
) -> str:
|
||||
return scenario.model(model=model, api_base=wire.url, **_AWS, **params, model_info=model_info)
|
||||
|
||||
|
||||
def _received_specs(wire: Wire) -> tuple[dict[str, JsonValue], ...]:
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1, [request.target for request in received]
|
||||
body: Final = _JSON.validate_json(received[0].body)
|
||||
tools: Final = _LIST.validate_python(object_value(body["toolConfig"])["tools"])
|
||||
return tuple(object_value(object_value(tool)["toolSpec"]) for tool in tools)
|
||||
|
||||
|
||||
def _schema_of(spec: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return object_value(object_value(spec["inputSchema"])["json"])
|
||||
|
||||
|
||||
def _only_schema(wire: Wire) -> dict[str, JsonValue]:
|
||||
(spec,) = _received_specs(wire)
|
||||
assert spec["name"] == _TOOL, spec
|
||||
return _schema_of(spec)
|
||||
|
||||
|
||||
def _assert_tool_call_relayed(endpoint: Endpoint, response: httpx.Response) -> None:
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = _JSON.validate_json(response.content)
|
||||
match endpoint:
|
||||
case "chat":
|
||||
message: Final = object_value(object_value(_LIST.validate_python(body["choices"])[0])["message"])
|
||||
(call,) = _LIST.validate_python(message["tool_calls"])
|
||||
function: Final = object_value(object_value(call)["function"])
|
||||
assert function["name"] == _TOOL and json.loads(string_value(function["arguments"])) == _TOOL_INPUT, (
|
||||
response.text
|
||||
)
|
||||
case "messages":
|
||||
blocks: Final = tuple(object_value(block) for block in _LIST.validate_python(body["content"]))
|
||||
(tool_use,) = tuple(block for block in blocks if block.get("type") == "tool_use")
|
||||
assert tool_use["name"] == _TOOL and tool_use["input"] == _TOOL_INPUT, response.text
|
||||
case "responses":
|
||||
items: Final = tuple(object_value(item) for item in _LIST.validate_python(body["output"]))
|
||||
(call_item,) = tuple(item for item in items if item.get("type") == "function_call")
|
||||
assert call_item["name"] == _TOOL and json.loads(string_value(call_item["arguments"])) == _TOOL_INPUT, (
|
||||
response.text
|
||||
)
|
||||
|
||||
|
||||
def _stream_text(gateway: Gateway, endpoint: Endpoint, body: Mapping[str, JsonValue]) -> str:
|
||||
headers: Final = {"Authorization": f"Bearer {gateway.key}"}
|
||||
with gateway.client.stream("POST", _path(endpoint), json=body, headers=headers) as response:
|
||||
lines: Final = tuple(line for line in response.iter_lines() if line)
|
||||
assert response.status_code == 200, "\n".join(lines)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _openai_client(gateway: Gateway) -> openai.OpenAI:
|
||||
return openai.OpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
|
||||
return openai.AsyncOpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _schema_sent_through(
|
||||
gateway: Gateway, wire: Wire, endpoint: Endpoint, model: str, tool: Mapping[str, JsonValue], **extra: JsonValue
|
||||
) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request("POST", _path(endpoint), _body(endpoint, model, (tool,), **extra))
|
||||
_assert_tool_call_relayed(endpoint, response)
|
||||
return _only_schema(wire)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", _ENDPOINTS)
|
||||
def test_flagged_model_receives_a_lookaround_free_schema_and_the_tool_call_comes_back(
|
||||
gateway: Gateway, endpoint: Endpoint
|
||||
) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
tool: Final = _tool_for(endpoint, _TOOL, _SCHEMA_AS_SENT)
|
||||
assert _schema_sent_through(gateway, wire, endpoint, model, tool) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", _ENDPOINTS)
|
||||
def test_flagged_model_streams_after_the_schema_lost_its_lookarounds(gateway: Gateway, endpoint: Endpoint) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
tool: Final = _tool_for(endpoint, _TOOL, _SCHEMA_AS_SENT)
|
||||
streamed: Final = _stream_text(gateway, endpoint, _body(endpoint, model, (tool,), stream=True))
|
||||
assert _ANSWER in streamed, streamed
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1 and unquote(received[0].target).endswith("/converse-stream"), received
|
||||
(tool_block,) = _LIST.validate_python(
|
||||
object_value(_JSON.validate_json(received[0].body)["toolConfig"])["tools"]
|
||||
)
|
||||
assert _schema_of(object_value(object_value(tool_block)["toolSpec"])) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
def test_openai_sdk_sync_chat_sends_a_lookaround_free_schema(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
client: Final = _openai_client(gateway)
|
||||
completion: Final = client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_openai_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_tokens=64,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
(call,) = completion.choices[0].message.tool_calls or ()
|
||||
assert call.function.name == _TOOL and json.loads(call.function.arguments) == _TOOL_INPUT
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
chunks: Final = tuple(
|
||||
client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_openai_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_tokens=64,
|
||||
stream=True,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
)
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == _ANSWER
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
async def test_openai_sdk_async_chat_sends_a_lookaround_free_schema(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
client: Final = _async_openai_client(gateway)
|
||||
completion: Final = await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_openai_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_tokens=64,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
(call,) = completion.choices[0].message.tool_calls or ()
|
||||
assert call.function.name == _TOOL and json.loads(call.function.arguments) == _TOOL_INPUT
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
stream: Final = await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_openai_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_tokens=64,
|
||||
stream=True,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
text: Final = "".join([chunk.choices[0].delta.content or "" async for chunk in stream if chunk.choices])
|
||||
assert text == _ANSWER
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
def test_anthropic_sdk_sync_messages_send_a_lookaround_free_schema(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
client: Final = anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
|
||||
message: Final = client.messages.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_anthropic_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
(tool_use,) = tuple(block for block in message.content if block.type == "tool_use")
|
||||
assert tool_use.name == _TOOL and tool_use.input == _TOOL_INPUT
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
with client.messages.stream(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_anthropic_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
extra_body=_NO_CACHE,
|
||||
) as stream:
|
||||
text: Final = "".join(stream.text_stream)
|
||||
assert text == _ANSWER
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
async def test_anthropic_sdk_async_messages_send_a_lookaround_free_schema(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
client: Final = anthropic.AsyncAnthropic(
|
||||
base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0
|
||||
)
|
||||
message: Final = await client.messages.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_anthropic_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
(tool_use,) = tuple(block for block in message.content if block.type == "tool_use")
|
||||
assert tool_use.name == _TOOL and tool_use.input == _TOOL_INPUT
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
async with client.messages.stream(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": _PROMPT}],
|
||||
tools=[_anthropic_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
extra_body=_NO_CACHE,
|
||||
) as stream:
|
||||
text: Final = "".join([piece async for piece in stream.text_stream])
|
||||
assert text == _ANSWER
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
def test_openai_sdk_sync_responses_send_a_lookaround_free_schema(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
client: Final = _openai_client(gateway)
|
||||
response: Final = client.responses.create(
|
||||
model=model,
|
||||
input=_PROMPT,
|
||||
tools=[_responses_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_output_tokens=64,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
(call,) = tuple(item for item in response.output if item.type == "function_call")
|
||||
assert call.name == _TOOL and json.loads(call.arguments) == _TOOL_INPUT
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
events: Final = tuple(
|
||||
client.responses.create(
|
||||
model=model,
|
||||
input=_PROMPT,
|
||||
tools=[_responses_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_output_tokens=64,
|
||||
stream=True,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
)
|
||||
deltas: Final = "".join(event.delta for event in events if event.type == "response.output_text.delta")
|
||||
assert deltas == _ANSWER
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
async def test_openai_sdk_async_responses_send_a_lookaround_free_schema(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
client: Final = _async_openai_client(gateway)
|
||||
response: Final = await client.responses.create(
|
||||
model=model,
|
||||
input=_PROMPT,
|
||||
tools=[_responses_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_output_tokens=64,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
(call,) = tuple(item for item in response.output if item.type == "function_call")
|
||||
assert call.name == _TOOL and json.loads(call.arguments) == _TOOL_INPUT
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
stream: Final = await client.responses.create(
|
||||
model=model,
|
||||
input=_PROMPT,
|
||||
tools=[_responses_tool(_TOOL, _SCHEMA_AS_SENT)],
|
||||
max_output_tokens=64,
|
||||
stream=True,
|
||||
extra_body=_NO_CACHE,
|
||||
)
|
||||
deltas: Final = "".join([event.delta async for event in stream if event.type == "response.output_text.delta"])
|
||||
assert deltas == _ANSWER
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
def test_grok_on_the_explicit_converse_route_is_flagged_too(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/converse/{_GROK}")
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT)
|
||||
assert _schema_sent_through(gateway, wire, "chat", model, tool) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
def test_a_tool_without_lookarounds_beside_a_cleaned_one_is_forwarded_untouched(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
tools: Final = (_openai_tool(_TOOL, _SCHEMA_AS_SENT), _openai_tool(_PLAIN_TOOL, _PLAIN_SCHEMA))
|
||||
response: Final = gateway.request("POST", _path("chat"), _body("chat", model, tools))
|
||||
_assert_tool_call_relayed("chat", response)
|
||||
cleaned, plain = _received_specs(wire)
|
||||
assert (cleaned["name"], _schema_of(cleaned)) == (_TOOL, _WIRE_LOOKAROUND_FREE)
|
||||
assert plain == {
|
||||
"name": _PLAIN_TOOL,
|
||||
"description": f"{_PLAIN_TOOL} tool",
|
||||
"inputSchema": {"json": _WIRE_PLAIN},
|
||||
}, plain
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_id", (_NOVA, _CLAUDE))
|
||||
def test_models_without_the_flag_keep_their_schema_as_sent(gateway: Gateway, model_id: str) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/converse/{model_id}")
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT)
|
||||
assert _schema_sent_through(gateway, wire, "chat", model, tool) == _WIRE_AS_SENT
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model_id", "model_info", "params", "expected"),
|
||||
(
|
||||
(_KIMI, {"supports_regex_lookaround": True}, {}, _WIRE_AS_SENT),
|
||||
(_NOVA, {"supports_regex_lookaround": False}, {}, _WIRE_LOOKAROUND_FREE),
|
||||
(_PROFILE_ARN, None, {"base_model": f"bedrock/{_KIMI}"}, _WIRE_LOOKAROUND_FREE),
|
||||
(_PROFILE_ARN, None, {}, _WIRE_AS_SENT),
|
||||
(_KIMI, {"supports_regex_lookaround": None}, {}, _WIRE_LOOKAROUND_FREE),
|
||||
(_NOVA, {"supports_regex_lookaround": "false"}, {}, _WIRE_AS_SENT),
|
||||
(_PROFILE_ARN, {"supports_regex_lookaround": True}, {"base_model": f"bedrock/{_KIMI}"}, _WIRE_AS_SENT),
|
||||
(_KIMI, None, {"base_model": ""}, _WIRE_LOOKAROUND_FREE),
|
||||
),
|
||||
ids=(
|
||||
"deployment-true-wins-over-map",
|
||||
"deployment-false-flags-an-unflagged-model",
|
||||
"base-model-flags-a-profile-arn",
|
||||
"bare-profile-arn-keeps-the-schema",
|
||||
"null-falls-back-to-the-map",
|
||||
"string-false-is-not-a-flag",
|
||||
"deployment-true-wins-over-base-model",
|
||||
"empty-base-model-falls-back-to-the-model",
|
||||
),
|
||||
)
|
||||
def test_deployment_settings_decide_before_the_cost_map(
|
||||
gateway: Gateway,
|
||||
model_id: str,
|
||||
model_info: Mapping[str, JsonValue] | None,
|
||||
params: Mapping[str, JsonValue],
|
||||
expected: Mapping[str, JsonValue],
|
||||
) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{model_id}", model_info=model_info, **params)
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT)
|
||||
assert _schema_sent_through(gateway, wire, "chat", model, tool) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model_id", "flag", "expected_for_the_bare_sibling"),
|
||||
((_KIMI, True, _WIRE_LOOKAROUND_FREE), (_NOVA, False, _WIRE_AS_SENT)),
|
||||
ids=("kimi-sibling-keeps-the-map-false", "nova-sibling-keeps-the-map-absence"),
|
||||
)
|
||||
@pytest.mark.parametrize("flagged_first", (True, False), ids=("flagged-registered-first", "bare-registered-first"))
|
||||
def test_a_deployment_flag_never_reaches_its_sibling_on_the_same_model(
|
||||
gateway: Gateway,
|
||||
model_id: str,
|
||||
flag: bool,
|
||||
expected_for_the_bare_sibling: Mapping[str, JsonValue],
|
||||
flagged_first: bool,
|
||||
) -> None:
|
||||
flag_info: Final[dict[str, JsonValue]] = {"supports_regex_lookaround": flag}
|
||||
first_info, second_info = (flag_info, None) if flagged_first else (None, flag_info)
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
first: Final = _deployment(scenario, wire, f"bedrock/{model_id}", model_info=first_info)
|
||||
second: Final = _deployment(scenario, wire, f"bedrock/{model_id}", model_info=second_info)
|
||||
bare: Final = second if flagged_first else first
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT)
|
||||
assert _schema_sent_through(gateway, wire, "chat", bare, tool) == expected_for_the_bare_sibling
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model_id", "body_base_model", "expected"),
|
||||
((_NOVA, f"bedrock/{_KIMI}", _WIRE_LOOKAROUND_FREE), (_KIMI, f"bedrock/{_NOVA}", _WIRE_LOOKAROUND_FREE)),
|
||||
ids=("client-base-model-can-loosen-an-unflagged-deployment", "client-base-model-cannot-restore-a-flagged-one"),
|
||||
)
|
||||
def test_a_base_model_in_the_request_body_only_ever_loosens(
|
||||
gateway: Gateway, model_id: str, body_base_model: str, expected: Mapping[str, JsonValue]
|
||||
) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{model_id}")
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT)
|
||||
assert _schema_sent_through(gateway, wire, "chat", model, tool, base_model=body_base_model) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("subschema", "expected"),
|
||||
(
|
||||
(
|
||||
{
|
||||
"type": "object",
|
||||
"patternProperties": {_POSITIVE_LOOKAHEAD_KEY: {"type": "string"}, r"^y_(?!z)": {"type": "integer"}},
|
||||
"additionalProperties": False,
|
||||
},
|
||||
{
|
||||
"type": "object",
|
||||
"patternProperties": {},
|
||||
"additionalProperties": {"anyOf": [{"type": "string"}, {"type": "integer"}]},
|
||||
},
|
||||
),
|
||||
(
|
||||
{"type": "object", "patternProperties": {_POSITIVE_LOOKAHEAD_KEY: {"type": "string"}}},
|
||||
{"type": "object", "patternProperties": {}},
|
||||
),
|
||||
(
|
||||
{"type": "object", "properties": {"name": {"type": "string", "pattern": r"\(?=x"}}},
|
||||
{"type": "object", "properties": {"name": {"type": "string"}}},
|
||||
),
|
||||
(
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"dependencies": {"name": {"properties": {"alias": {"type": "string", "pattern": _LOOKAHEAD}}}},
|
||||
},
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"dependencies": {"name": {"properties": {"alias": {"type": "string", "pattern": _LOOKAHEAD}}}},
|
||||
},
|
||||
),
|
||||
),
|
||||
ids=(
|
||||
"two-dropped-pattern-properties-become-an-anyof",
|
||||
"an-open-object-just-loses-the-key",
|
||||
"an-escaped-literal-spelling-an-opener-is-dropped-too",
|
||||
"draft-07-dependencies-are-not-walked",
|
||||
),
|
||||
)
|
||||
def test_schema_shapes_at_the_edges_of_the_walk(
|
||||
gateway: Gateway, subschema: Mapping[str, JsonValue], expected: Mapping[str, JsonValue]
|
||||
) -> None:
|
||||
schema: Final[dict[str, JsonValue]] = {"type": "object", "properties": {"labels": dict(subschema)}}
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
assert _schema_sent_through(gateway, wire, "chat", model, _openai_tool(_TOOL, schema)) == {
|
||||
"type": "object",
|
||||
"properties": {"labels": dict(expected)},
|
||||
"required": [],
|
||||
}
|
||||
|
||||
|
||||
def test_strict_is_still_withheld_from_a_flagged_non_anthropic_model(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT, strict=True)
|
||||
response: Final = gateway.request("POST", _path("chat"), _body("chat", model, (tool,)))
|
||||
_assert_tool_call_relayed("chat", response)
|
||||
(spec,) = _received_specs(wire)
|
||||
assert spec == {"name": _TOOL, "description": f"{_TOOL} tool", "inputSchema": {"json": _WIRE_LOOKAROUND_FREE}}
|
||||
|
||||
|
||||
def test_a_json_schema_response_format_rides_the_same_tool_path(gateway: Gateway) -> None:
|
||||
schema: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {"collection": {"type": "string", "pattern": _LOOKAHEAD}},
|
||||
"required": ["collection"],
|
||||
}
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
_path("chat"),
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _PROMPT}],
|
||||
"max_tokens": 64,
|
||||
"response_format": {"type": "json_schema", "json_schema": {"name": "document", "schema": schema}},
|
||||
**_NO_CACHE,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
(spec,) = _received_specs(wire)
|
||||
assert spec["name"] == "json_tool_call", spec
|
||||
assert _schema_of(spec) == {
|
||||
"type": "object",
|
||||
"properties": {"collection": {"type": "string"}},
|
||||
"required": ["collection"],
|
||||
}, spec
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("pattern", "expected_property"),
|
||||
(
|
||||
(5, {"type": "string", "pattern": 5}),
|
||||
([_LOOKAHEAD], {"type": "string", "pattern": [_LOOKAHEAD]}),
|
||||
("", {"type": "string", "pattern": ""}),
|
||||
("a" * 5120, {"type": "string", "pattern": "a" * 5120}),
|
||||
("a" * 5120 + "(?=b)", {"type": "string"}),
|
||||
),
|
||||
ids=("int", "list", "empty", "5kb-plain", "5kb-ending-in-a-lookahead"),
|
||||
)
|
||||
def test_odd_pattern_values_are_forwarded_unless_they_are_a_lookaround_string(
|
||||
gateway: Gateway, pattern: JsonValue, expected_property: Mapping[str, JsonValue]
|
||||
) -> None:
|
||||
schema: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {"type": "string", "pattern": pattern},
|
||||
"doc_id": {"type": "string", "pattern": pattern},
|
||||
},
|
||||
}
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
assert _schema_sent_through(gateway, wire, "chat", model, _openai_tool(_TOOL, schema)) == {
|
||||
"type": "object",
|
||||
"properties": {"collection": dict(expected_property), "doc_id": dict(expected_property)},
|
||||
"required": [],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"parameters",
|
||||
(None, {"type": "object", "properties": [{"name": "collection", "pattern": _LOOKAHEAD}]}),
|
||||
ids=("null-parameters", "properties-as-a-list"),
|
||||
)
|
||||
def test_malformed_tool_parameters_never_take_the_proxy_down(gateway: Gateway, parameters: JsonValue) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
tool: Final[dict[str, JsonValue]] = {
|
||||
"type": "function",
|
||||
"function": {"name": _TOOL, "description": f"{_TOOL} tool", "parameters": parameters},
|
||||
}
|
||||
response: Final = gateway.request("POST", _path("chat"), _body("chat", model, (tool,)))
|
||||
assert response.status_code in (200, 400), response.text
|
||||
if response.status_code == 400:
|
||||
assert "error" in _JSON.validate_json(response.content), response.text
|
||||
wire.drain()
|
||||
control: Final = gateway.request(
|
||||
"POST", _path("chat"), _body("chat", model, (_openai_tool(_TOOL, _SCHEMA_AS_SENT),))
|
||||
)
|
||||
_assert_tool_call_relayed("chat", control)
|
||||
assert _only_schema(wire) == _WIRE_LOOKAROUND_FREE
|
||||
|
||||
|
||||
def test_an_unauthenticated_request_never_reaches_the_peer(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
response: Final = gateway.request(
|
||||
"POST", _path("chat"), _body("chat", model, (_openai_tool(_TOOL, _SCHEMA_AS_SENT),)), key="sk-not-a-key"
|
||||
)
|
||||
assert response.status_code == 401, response.text
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
def test_a_bedrock_rejection_of_an_unflagged_model_reaches_the_caller(gateway: Gateway) -> None:
|
||||
with wire_server(_rejecting_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/converse/{_CLAUDE}")
|
||||
response: Final = gateway.request(
|
||||
"POST", _path("chat"), _body("chat", model, (_openai_tool(_TOOL, _SCHEMA_AS_SENT),))
|
||||
)
|
||||
assert response.status_code == 400, response.text
|
||||
assert _BEDROCK_REJECTION in response.text, response.text
|
||||
assert _only_schema(wire) == _WIRE_AS_SENT
|
||||
|
||||
|
||||
@pytest.mark.timeout(120)
|
||||
def test_the_worst_case_lookaround_input_scans_in_linear_time(gateway: Gateway) -> None:
|
||||
pattern: Final = "(?<" * (2 * 1024 * 1024 // 3)
|
||||
schema: Final[dict[str, JsonValue]] = {
|
||||
"type": "object",
|
||||
"properties": {"collection": {"type": "string", "pattern": pattern}},
|
||||
}
|
||||
liveliness: Final[list[tuple[float, int]]] = []
|
||||
stop: Final = threading.Event()
|
||||
|
||||
def poll() -> None:
|
||||
while not stop.is_set():
|
||||
liveliness.append(_timed_liveliness(gateway))
|
||||
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}")
|
||||
poller: Final = threading.Thread(target=poll)
|
||||
poller.start()
|
||||
started: Final = time.perf_counter()
|
||||
response: Final = gateway.request("POST", _path("chat"), _body("chat", model, (_openai_tool(_TOOL, schema),)))
|
||||
elapsed: Final = time.perf_counter() - started
|
||||
stop.set()
|
||||
poller.join()
|
||||
_assert_tool_call_relayed("chat", response)
|
||||
assert elapsed < 30, elapsed
|
||||
assert liveliness and max(latency for latency, _ in liveliness) < 5, liveliness
|
||||
assert {status for _, status in liveliness} == {200}, liveliness
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
def _timed_liveliness(gateway: Gateway) -> tuple[float, int]:
|
||||
started: Final = time.perf_counter()
|
||||
probe: Final = gateway.client.get("/health/liveliness")
|
||||
return time.perf_counter() - started, probe.status_code
|
||||
|
||||
|
||||
def _model_id(gateway: Gateway, name: str) -> str:
|
||||
entries: Final = gateway.get("/model/info")["data"]
|
||||
assert isinstance(entries, list), entries
|
||||
(identity,) = (
|
||||
string_value(object_value(object_value(entry)["model_info"])["id"])
|
||||
for entry in entries
|
||||
if object_value(entry)["model_name"] == name
|
||||
)
|
||||
return identity
|
||||
|
||||
|
||||
def _settled_schema(gateway: Gateway, wire: Wire, model: str, expected: Mapping[str, JsonValue]) -> None:
|
||||
tool: Final = _openai_tool(_TOOL, _SCHEMA_AS_SENT)
|
||||
eventually(
|
||||
lambda: tuple(_schema_sent_through(gateway, wire, "chat", model, tool) for _ in range(8)),
|
||||
lambda schemas: all(schema == expected for schema in schemas),
|
||||
seconds=90,
|
||||
)
|
||||
|
||||
|
||||
def _patch_flag(gateway: Gateway, identity: str, flag: bool) -> None:
|
||||
patched: Final = gateway.request(
|
||||
"PATCH", f"/model/{identity}/update", {"model_info": {"supports_regex_lookaround": flag}}
|
||||
)
|
||||
assert patched.status_code == 200, patched.text
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_updating_the_flag_on_a_live_deployment_takes_effect_without_a_restart(gateway: Gateway) -> None:
|
||||
with wire_server(_bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, wire, f"bedrock/{_KIMI}", model_info={"supports_regex_lookaround": True})
|
||||
_settled_schema(gateway, wire, model, _WIRE_AS_SENT)
|
||||
identity: Final = _model_id(gateway, model)
|
||||
_patch_flag(gateway, identity, False)
|
||||
_settled_schema(gateway, wire, model, _WIRE_LOOKAROUND_FREE)
|
||||
_patch_flag(gateway, identity, True)
|
||||
_settled_schema(gateway, wire, model, _WIRE_AS_SENT)
|
||||
|
|
@ -30,6 +30,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
|
||||
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
|
||||
_ARTIFACT_DATA_ID_PATTERN: Final = r"^(?!\.\.?(?:\/|$))[A-Za-z0-9_\-.~:@+]{1,200}$"
|
||||
|
||||
|
||||
def test_get_format_from_file_id():
|
||||
|
|
@ -1620,39 +1621,74 @@ class TestToolWithSanitizedParameters:
|
|||
|
||||
assert tool_with_sanitized_parameters(tool, flatten_combinators_and_drop_non_python_regex_patterns) is tool
|
||||
|
||||
def test_sanitizes_the_input_schema_of_an_anthropic_tool(self):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_lookaround_regex_patterns,
|
||||
tool_with_sanitized_parameters,
|
||||
)
|
||||
|
||||
tool = {
|
||||
"name": "ArtifactData",
|
||||
"description": "Read a shared database",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"doc_id": {"type": "string", "pattern": _ARTIFACT_DATA_ID_PATTERN}},
|
||||
},
|
||||
}
|
||||
|
||||
result = tool_with_sanitized_parameters(tool, drop_lookaround_regex_patterns)
|
||||
|
||||
assert result == {
|
||||
"name": "ArtifactData",
|
||||
"description": "Read a shared database",
|
||||
"input_schema": {"type": "object", "properties": {"doc_id": {"type": "string"}}},
|
||||
}
|
||||
assert tool["input_schema"]["properties"]["doc_id"]["pattern"] == _ARTIFACT_DATA_ID_PATTERN
|
||||
|
||||
def test_returns_the_same_anthropic_tool_when_its_schema_has_nothing_to_drop(self):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_lookaround_regex_patterns,
|
||||
tool_with_sanitized_parameters,
|
||||
)
|
||||
|
||||
tool = {"name": "Read", "input_schema": {"type": "object", "properties": {"path": {"type": "string"}}}}
|
||||
|
||||
assert tool_with_sanitized_parameters(tool, drop_lookaround_regex_patterns) is tool
|
||||
|
||||
|
||||
def _regex_schema(pattern):
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"field": {"type": "string", "pattern": pattern},
|
||||
"writes": {
|
||||
"type": "array",
|
||||
"items": {"properties": {"doc_id": {"type": "string", "pattern": pattern}}},
|
||||
},
|
||||
"query": {"anyOf": [{"type": "string", "pattern": pattern}, {"type": "null"}]},
|
||||
"pair": {"type": "array", "prefixItems": [{"type": "string", "pattern": pattern}]},
|
||||
"extra": {"type": "object", "additionalProperties": {"type": "string", "pattern": pattern}},
|
||||
"tagged": {
|
||||
"type": "object",
|
||||
"patternProperties": {pattern: {"type": "string"}, "^x_": {"type": "integer"}},
|
||||
},
|
||||
},
|
||||
"$defs": {"segment": {"type": "string", "pattern": pattern}},
|
||||
"required": ["field"],
|
||||
}
|
||||
|
||||
|
||||
class TestDropNonPythonRegexPatterns:
|
||||
"""Claude Code's Artifact tool declares ECMA-262 ``\\p{..}`` escapes that OpenAI's
|
||||
validator, which compiles ``pattern`` values and ``patternProperties`` keys with
|
||||
Python ``re``, refuses as "not a 'regex'"."""
|
||||
|
||||
def _schema(self, pattern):
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"field": {"type": "string", "pattern": pattern},
|
||||
"writes": {
|
||||
"type": "array",
|
||||
"items": {"properties": {"doc_id": {"type": "string", "pattern": pattern}}},
|
||||
},
|
||||
"query": {"anyOf": [{"type": "string", "pattern": pattern}, {"type": "null"}]},
|
||||
"pair": {"type": "array", "prefixItems": [{"type": "string", "pattern": pattern}]},
|
||||
"extra": {"type": "object", "additionalProperties": {"type": "string", "pattern": pattern}},
|
||||
"tagged": {
|
||||
"type": "object",
|
||||
"patternProperties": {pattern: {"type": "string"}, "^x_": {"type": "integer"}},
|
||||
},
|
||||
},
|
||||
"$defs": {"segment": {"type": "string", "pattern": pattern}},
|
||||
"required": ["field"],
|
||||
}
|
||||
|
||||
def test_drops_every_regex_python_re_rejects_from_every_schema_position(self):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_non_python_regex_patterns,
|
||||
)
|
||||
|
||||
schema = self._schema(_ARTIFACT_FIELD_PATTERN)
|
||||
schema = _regex_schema(_ARTIFACT_FIELD_PATTERN)
|
||||
|
||||
result = drop_non_python_regex_patterns(schema)
|
||||
|
||||
|
|
@ -1666,14 +1702,14 @@ class TestDropNonPythonRegexPatterns:
|
|||
assert properties["tagged"]["patternProperties"] == {"^x_": {"type": "integer"}}
|
||||
assert result["$defs"]["segment"] == {"type": "string"}
|
||||
assert result["required"] == ["field"]
|
||||
assert schema == self._schema(_ARTIFACT_FIELD_PATTERN)
|
||||
assert schema == _regex_schema(_ARTIFACT_FIELD_PATTERN)
|
||||
|
||||
def test_keeps_regexes_python_re_compiles_and_returns_the_same_object(self):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_non_python_regex_patterns,
|
||||
)
|
||||
|
||||
schema = self._schema(r'^(?!__.*__$)[^"\\./[\]]{1,200}$')
|
||||
schema = _regex_schema(r'^(?!__.*__$)[^"\\./[\]]{1,200}$')
|
||||
|
||||
assert drop_non_python_regex_patterns(schema) is schema
|
||||
|
||||
|
|
@ -1737,6 +1773,137 @@ class TestDropNonPythonRegexPatterns:
|
|||
assert drop_non_python_regex_patterns(schema) is schema
|
||||
|
||||
|
||||
class TestDropLookaroundRegexPatterns:
|
||||
"""Kimi K3 and Grok 4.6/4.7 on Bedrock Converse reject every tool schema regex that
|
||||
uses a lookaround assertion, Claude Code's ``ArtifactData`` ``pattern`` included."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"pattern",
|
||||
[r"^(?!x).*$", r"^(?=.*a).*$", r"^.*(?<!x)$", r"^.*(?<=a)$", _ARTIFACT_DATA_ID_PATTERN],
|
||||
ids=["negative-lookahead", "positive-lookahead", "negative-lookbehind", "positive-lookbehind", "ArtifactData"],
|
||||
)
|
||||
def test_drops_every_lookaround_regex_from_every_schema_position(self, pattern):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_lookaround_regex_patterns,
|
||||
)
|
||||
|
||||
schema = _regex_schema(pattern)
|
||||
|
||||
result = drop_lookaround_regex_patterns(schema)
|
||||
|
||||
assert '"pattern"' not in json.dumps(result)
|
||||
properties = result["properties"]
|
||||
assert properties["field"] == {"type": "string"}
|
||||
assert properties["writes"]["items"]["properties"]["doc_id"] == {"type": "string"}
|
||||
assert properties["query"]["anyOf"] == [{"type": "string"}, {"type": "null"}]
|
||||
assert properties["pair"]["prefixItems"] == [{"type": "string"}]
|
||||
assert properties["extra"]["additionalProperties"] == {"type": "string"}
|
||||
assert properties["tagged"]["patternProperties"] == {"^x_": {"type": "integer"}}
|
||||
assert result["$defs"]["segment"] == {"type": "string"}
|
||||
assert result["required"] == ["field"]
|
||||
assert schema == _regex_schema(pattern)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"pattern",
|
||||
[r"^[A-Za-z0-9_\-.~:@+]{1,200}$", r"^(?:a|b)+$", r"^(?P<name>\w+)$", r"^(?i)abc$", r"^[^\p{Cc}\p{Cf}]{1,200}$"],
|
||||
ids=["plain", "non-capturing-group", "named-group", "inline-flag", "non-python-without-lookaround"],
|
||||
)
|
||||
def test_keeps_regexes_without_lookaround_and_returns_the_same_object(self, pattern):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_lookaround_regex_patterns,
|
||||
)
|
||||
|
||||
schema = _regex_schema(pattern)
|
||||
|
||||
assert drop_lookaround_regex_patterns(schema) is schema
|
||||
|
||||
def test_lookaround_inside_data_positions_is_not_a_regex(self):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_lookaround_regex_patterns,
|
||||
)
|
||||
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"pattern": {"type": "string"},
|
||||
"template": {"type": "object", "default": {"pattern": _ARTIFACT_DATA_ID_PATTERN}},
|
||||
"hint": {"type": "string", "description": "ids match " + _ARTIFACT_DATA_ID_PATTERN},
|
||||
},
|
||||
"required": ["pattern"],
|
||||
}
|
||||
|
||||
assert drop_lookaround_regex_patterns(schema) is schema
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("dropper", "patterns"),
|
||||
[
|
||||
("drop_non_python_regex_patterns", (_ARTIFACT_FIELD_PATTERN, r"^\p{L}+$")),
|
||||
("drop_lookaround_regex_patterns", (_ARTIFACT_DATA_ID_PATTERN, r"^(?=.*[a-z])\w+$")),
|
||||
],
|
||||
ids=["non-python", "lookaround"],
|
||||
)
|
||||
class TestDroppedPatternPropertiesKeepTheirNamesAllowed:
|
||||
"""Dropping a ``patternProperties`` key from an object closed by ``additionalProperties:
|
||||
false`` must not ban the names that key allowed: its value schema takes over as the
|
||||
object's ``additionalProperties``."""
|
||||
|
||||
@staticmethod
|
||||
def _drop(dropper):
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_lookaround_regex_patterns,
|
||||
drop_non_python_regex_patterns,
|
||||
)
|
||||
|
||||
return {
|
||||
"drop_non_python_regex_patterns": drop_non_python_regex_patterns,
|
||||
"drop_lookaround_regex_patterns": drop_lookaround_regex_patterns,
|
||||
}[dropper]
|
||||
|
||||
def test_closed_object_takes_the_dropped_value_schema(self, dropper, patterns):
|
||||
schema = {
|
||||
"type": "object",
|
||||
"patternProperties": {patterns[0]: {"type": "string", "pattern": patterns[0]}},
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
assert self._drop(dropper)(schema) == {
|
||||
"type": "object",
|
||||
"patternProperties": {},
|
||||
"additionalProperties": {"type": "string"},
|
||||
}
|
||||
|
||||
def test_closed_object_losing_two_entries_accepts_either_value_schema(self, dropper, patterns):
|
||||
schema = {
|
||||
"type": "object",
|
||||
"patternProperties": {
|
||||
patterns[0]: {"type": "string"},
|
||||
patterns[1]: {"type": "integer"},
|
||||
"^x_": {"type": "boolean"},
|
||||
},
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
assert self._drop(dropper)(schema) == {
|
||||
"type": "object",
|
||||
"patternProperties": {"^x_": {"type": "boolean"}},
|
||||
"additionalProperties": {"anyOf": [{"type": "string"}, {"type": "integer"}]},
|
||||
}
|
||||
|
||||
def test_object_with_its_own_additional_properties_schema_keeps_it(self, dropper, patterns):
|
||||
schema = {
|
||||
"type": "object",
|
||||
"patternProperties": {patterns[0]: {"type": "string"}},
|
||||
"additionalProperties": {"type": "integer"},
|
||||
}
|
||||
|
||||
assert self._drop(dropper)(schema) == {
|
||||
"type": "object",
|
||||
"patternProperties": {},
|
||||
"additionalProperties": {"type": "integer"},
|
||||
}
|
||||
|
||||
|
||||
class TestRequestContainsImageContent:
|
||||
"""One detector for every dialect that reaches pre-routing hooks untranslated."""
|
||||
|
||||
|
|
|
|||
|
|
@ -637,6 +637,142 @@ def test_output_config_effort_forwarded_into_additional_request_fields(model):
|
|||
assert additional.get("output_config") == {"effort": "high"}
|
||||
|
||||
|
||||
_ARTIFACT_DATA_ID_PATTERN: Final = r"^(?!\.\.?(?:\/|$))[A-Za-z0-9_\-.~:@+]{1,200}$"
|
||||
_ARTIFACT_DATA_INPUT_SCHEMA: Final = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"collection": {"type": "string", "pattern": _ARTIFACT_DATA_ID_PATTERN, "description": "Collection"},
|
||||
"doc_id": {"type": "string", "pattern": _ARTIFACT_DATA_ID_PATTERN},
|
||||
"writes": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "object",
|
||||
"properties": {"doc_id": {"type": "string", "pattern": _ARTIFACT_DATA_ID_PATTERN}},
|
||||
},
|
||||
},
|
||||
"limit": {"type": "integer", "minimum": 1},
|
||||
},
|
||||
"required": ["collection"],
|
||||
}
|
||||
_ARTIFACT_DATA_ANTHROPIC_TOOL: Final = {
|
||||
"name": "ArtifactData",
|
||||
"description": "Read a shared database",
|
||||
"input_schema": _ARTIFACT_DATA_INPUT_SCHEMA,
|
||||
}
|
||||
_ARTIFACT_DATA_OPENAI_TOOL: Final = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "ArtifactData",
|
||||
"description": "Read a shared database",
|
||||
"parameters": _ARTIFACT_DATA_INPUT_SCHEMA,
|
||||
},
|
||||
}
|
||||
_LOOKAROUND_FREE_PROPERTIES: Final = {
|
||||
"collection": {"type": "string", "description": "Collection"},
|
||||
"doc_id": {"type": "string"},
|
||||
"writes": {"type": "array", "items": {"type": "object", "properties": {"doc_id": {"type": "string"}}}},
|
||||
"limit": {"type": "integer", "minimum": 1},
|
||||
}
|
||||
|
||||
|
||||
def _converse_tools(model, tools, litellm_params=None):
|
||||
request = AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"tools": copy.deepcopy(tools)},
|
||||
litellm_params=litellm_params or {},
|
||||
headers={},
|
||||
)
|
||||
return request["toolConfig"]["tools"]
|
||||
|
||||
|
||||
def _tool_schema_properties(model, tool, litellm_params=None):
|
||||
return _converse_tools(model, [tool], litellm_params)[0]["toolSpec"]["inputSchema"]["json"]["properties"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tool", [_ARTIFACT_DATA_ANTHROPIC_TOOL, _ARTIFACT_DATA_OPENAI_TOOL], ids=["anthropic-shape", "openai-shape"]
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"global.moonshotai.kimi-k3",
|
||||
"us.moonshotai.kimi-k3",
|
||||
"moonshotai.kimi-k3",
|
||||
"us-east-1/us.moonshotai.kimi-k3",
|
||||
"us.xai.grok-4.6",
|
||||
"us-gov.xai.grok-4.6",
|
||||
"global.xai.grok-4.7",
|
||||
"xai.grok-4.7",
|
||||
],
|
||||
)
|
||||
def test_transform_request_drops_lookaround_regex_for_models_the_cost_map_flags(tool, model):
|
||||
"""Kimi K3 and Grok 4.6/4.7 refuse the whole request over a lookaround in a tool schema regex."""
|
||||
tools = _converse_tools(model, [tool])
|
||||
|
||||
json_schema = tools[0]["toolSpec"]["inputSchema"]["json"]
|
||||
assert json_schema["properties"] == _LOOKAROUND_FREE_PROPERTIES
|
||||
assert json_schema["required"] == ["collection"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"us.anthropic.claude-sonnet-4-6",
|
||||
"us.amazon.nova-pro-v1:0",
|
||||
"us.meta.llama4-maverick-17b-instruct-v1:0",
|
||||
"us.openai.gpt-5.6-sol",
|
||||
],
|
||||
)
|
||||
def test_transform_request_keeps_lookaround_regex_for_models_that_accept_it(model):
|
||||
assert _tool_schema_properties(model, _ARTIFACT_DATA_ANTHROPIC_TOOL) == _ARTIFACT_DATA_INPUT_SCHEMA["properties"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"us.amazon.nova-lite-v1:0",
|
||||
"us.moonshotai.kimi-k4",
|
||||
"arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123",
|
||||
],
|
||||
)
|
||||
def test_transform_request_drops_lookaround_regex_when_the_deployment_model_info_opts_in(model):
|
||||
"""A deployment's ``model_info`` flag covers a model the cost map does not know, an inference profile included."""
|
||||
properties = _tool_schema_properties(
|
||||
model, _ARTIFACT_DATA_ANTHROPIC_TOOL, {"model_info": {"supports_regex_lookaround": False}}
|
||||
)
|
||||
|
||||
assert properties == _LOOKAROUND_FREE_PROPERTIES
|
||||
|
||||
|
||||
def test_transform_request_keeps_lookaround_regex_when_the_deployment_model_info_opts_out():
|
||||
properties = _tool_schema_properties(
|
||||
"global.moonshotai.kimi-k3", _ARTIFACT_DATA_ANTHROPIC_TOOL, {"model_info": {"supports_regex_lookaround": True}}
|
||||
)
|
||||
|
||||
assert properties["doc_id"]["pattern"] == _ARTIFACT_DATA_ID_PATTERN
|
||||
|
||||
|
||||
def test_transform_request_resolves_an_inference_profile_through_its_base_model():
|
||||
properties = _tool_schema_properties(
|
||||
"arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123",
|
||||
_ARTIFACT_DATA_ANTHROPIC_TOOL,
|
||||
{"base_model": "bedrock/global.moonshotai.kimi-k3"},
|
||||
)
|
||||
|
||||
assert properties == _LOOKAROUND_FREE_PROPERTIES
|
||||
|
||||
|
||||
def test_transform_request_drops_lookaround_regex_around_pre_formatted_tool_blocks():
|
||||
"""Blocks that arrive already in Bedrock shape, like Nova's grounding ``systemTool``, pass through as sent."""
|
||||
grounding: Final = {"systemTool": {"name": "nova_grounding"}}
|
||||
|
||||
tools = _converse_tools("global.moonshotai.kimi-k3", [_ARTIFACT_DATA_OPENAI_TOOL, grounding])
|
||||
|
||||
assert tools[0]["toolSpec"]["inputSchema"]["json"]["properties"] == _LOOKAROUND_FREE_PROPERTIES
|
||||
assert tools[1] == grounding
|
||||
|
||||
|
||||
def test_reasoning_effort_requests_summarized_display_converse():
|
||||
"""Regression LIT-5714: adaptive thinking synthesized from reasoning_effort must
|
||||
request the summarized display, otherwise the provider returns a blank thinking
|
||||
|
|
|
|||
|
|
@ -514,6 +514,34 @@ def test_should_not_pollute_shared_key_with_custom_nonzero_pricing():
|
|||
)
|
||||
|
||||
|
||||
def test_regex_lookaround_flag_stays_on_the_deployment_that_set_it() -> None:
|
||||
"""A deployment's ``supports_regex_lookaround`` override must not land on the shared
|
||||
``{provider}/{model}`` key, or every sibling deployment of that model would inherit it."""
|
||||
backend_model = "bedrock/us.xai.grok-4.6"
|
||||
deploy_id = "grok-deploy-keep-regex"
|
||||
|
||||
builtin_flag = litellm.get_model_info(model=backend_model).get("supports_regex_lookaround")
|
||||
model_keys = {
|
||||
deploy_id: litellm.model_cost.get(deploy_id),
|
||||
backend_model: copy.deepcopy(litellm.model_cost.get(backend_model)),
|
||||
}
|
||||
try:
|
||||
Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "grok-keep-regex",
|
||||
"litellm_params": {"model": backend_model},
|
||||
"model_info": {"id": deploy_id, "supports_regex_lookaround": not builtin_flag},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
assert litellm.model_cost[deploy_id]["supports_regex_lookaround"] is (not builtin_flag)
|
||||
assert litellm.get_model_info(model=backend_model).get("supports_regex_lookaround") is builtin_flag
|
||||
finally:
|
||||
_restore_model_cost_entries(model_keys)
|
||||
|
||||
|
||||
def test_should_store_full_pricing_under_deployment_model_id():
|
||||
"""
|
||||
Per-deployment pricing (including zero) should be stored and
|
||||
|
|
|
|||
|
|
@ -996,6 +996,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"enum": ["low", "medium", "high", "max", "xhigh"],
|
||||
},
|
||||
"bedrock_converse_supports_strict_tools": {"type": "boolean"},
|
||||
"supports_regex_lookaround": {"type": "boolean"},
|
||||
"tpm": {"type": "number"},
|
||||
"supported_endpoints": {
|
||||
"type": "array",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue