feat(models): add cache_config to Mistral (#3986)

This commit is contained in:
opieter-aws
2026-08-28 11:20:02 -04:00
committed by GitHub
parent e448bea9e4
commit eabcb137ad
4 changed files with 120 additions and 10 deletions
@@ -8,6 +8,7 @@ OpenAI caches prompt prefixes automatically server-side and routes reads on a ca
import warnings
from typing import Any
from ._validation import warn_on_cache_config_not_supported
from .model import CacheConfig
# OpenAI's prompt_cache_retention accepts only these literals. ttl maps through only on an exact
@@ -47,9 +48,4 @@ def apply_cache_config(request: dict[str, Any], cache_config: CacheConfig | None
stacklevel=4,
)
if cache_config.strategy != "auto" or cache_config.system_prompt_ttl is not True:
warnings.warn(
"openai caches prompt prefixes automatically server-side; cache_config.strategy and "
"system_prompt_ttl have no effect and will be ignored",
stacklevel=4,
)
warn_on_cache_config_not_supported(cache_config, "OpenAI", supported={"cache_key", "ttl"})
+35 -1
View File
@@ -1,14 +1,16 @@
"""Configuration validation utilities for model providers."""
import dataclasses
import re
import warnings
from collections.abc import Mapping
from collections.abc import Collection, Mapping
from typing import Any
from typing_extensions import get_type_hints
from ..types.content import ContentBlock
from ..types.tools import ToolChoice
from .model import CacheConfig
# Matches AWS region identifiers such as us-east-1, ap-southeast-1, and us-gov-east-1.
# ``\A``/``\Z`` anchor the whole string (``$`` would allow a trailing newline) and ``[0-9]``
@@ -73,6 +75,38 @@ def warn_on_tool_choice_not_supported(tool_choice: ToolChoice | None) -> None:
)
def _cache_config_fields_set(cache_config: CacheConfig) -> set[str]:
"""Return the names of cache_config fields the caller set to a non-default value."""
return {
field.name for field in dataclasses.fields(cache_config) if getattr(cache_config, field.name) != field.default
}
def warn_on_cache_config_not_supported(
cache_config: CacheConfig,
provider: str,
*,
supported: Collection[str],
stacklevel: int = 4,
) -> None:
"""Warn once about supplied cache_config fields a provider does not support.
Args:
cache_config: The provider's configured cache settings.
provider: Human-readable provider name for the warning message.
supported: Names of ``CacheConfig`` fields the provider applies; the rest are the no-ops.
stacklevel: Frames to skip so the warning points at the caller. Defaults to 4, correct when a
mapper one frame below ``format_request`` invokes this.
"""
unsupported = sorted(field for field in _cache_config_fields_set(cache_config) if field not in supported)
if unsupported:
warnings.warn(
f"cache_config fields {unsupported} have no effect on {provider}, which does not support them; "
"they will be ignored.",
stacklevel=stacklevel,
)
def _has_location_source(content: ContentBlock) -> bool:
"""Check if a content block contains a location source.
+27 -2
View File
@@ -18,14 +18,33 @@ from ..types.exceptions import ContextWindowOverflowException, ModelThrottledExc
from ..types.streaming import StopReason, StreamEvent
from ..types.tools import ToolChoice, ToolResult, ToolSpec, ToolUse
from ._defaults import resolve_config_metadata
from ._validation import _has_location_source, validate_config_keys, warn_on_tool_choice_not_supported
from .model import BaseModelConfig, Model
from ._validation import (
_has_location_source,
validate_config_keys,
warn_on_cache_config_not_supported,
warn_on_tool_choice_not_supported,
)
from .model import BaseModelConfig, CacheConfig, Model
logger = logging.getLogger(__name__)
T = TypeVar("T", bound=BaseModel)
def _apply_cache_config(request: dict[str, Any], cache_config: CacheConfig | None) -> None:
"""Map a CacheConfig onto a Mistral request.
Args:
request: The Mistral request dict to mutate.
cache_config: The provider's configured cache settings, if any.
"""
if cache_config is None:
return
if cache_config.cache_key is not None and "prompt_cache_key" not in request:
request["prompt_cache_key"] = cache_config.cache_key
warn_on_cache_config_not_supported(cache_config, "Mistral", supported={"cache_key"})
class MistralModel(Model):
"""Mistral API model provider implementation.
@@ -51,6 +70,9 @@ class MistralModel(Model):
temperature: Controls randomness in generation (0.0 to 1.0).
top_p: Controls diversity via nucleus sampling.
stream: Whether to enable streaming responses.
cache_config: Prompt-caching configuration. Mistral routes cache reads on
cache_config.cache_key (mapped to the request's prompt_cache_key); it exposes no
retention or placement controls, so other fields are ignored.
"""
model_id: str
@@ -58,6 +80,7 @@ class MistralModel(Model):
temperature: float | None
top_p: float | None
stream: bool | None
cache_config: CacheConfig | None
def __init__(
self,
@@ -295,6 +318,8 @@ class MistralModel(Model):
for tool_spec in tool_specs
]
_apply_cache_config(request, self.config.get("cache_config"))
return request
def format_chunk(self, event: dict[str, Any]) -> StreamEvent:
@@ -5,6 +5,7 @@ import pydantic
import pytest
import strands
from strands.models import CacheConfig
from strands.models.mistral import MistralModel
from strands.types.exceptions import ContextWindowOverflowException, ModelThrottledException
@@ -125,6 +126,61 @@ def test_format_request_default(model, messages, model_id):
assert actual_request == exp_request
def test_cache_config_round_trips(model_id, max_tokens, captured_warnings):
"""cache_config is a valid config field and survives get_config/update_config unchanged."""
cache_config = CacheConfig(cache_key="tenant-42")
model = MistralModel(model_id=model_id, max_tokens=max_tokens, cache_config=cache_config)
assert model.get_config()["cache_config"] is cache_config
updated = CacheConfig(cache_key="tenant-99")
model.update_config(cache_config=updated)
assert model.get_config()["cache_config"] is updated
assert not any("Invalid configuration parameters" in str(warning.message) for warning in captured_warnings)
def test_cache_key_maps_to_prompt_cache_key(model_id, max_tokens, messages):
"""Mistral is key-routed: cache_config.cache_key maps to the request's prompt_cache_key."""
model = MistralModel(model_id=model_id, max_tokens=max_tokens, cache_config=CacheConfig(cache_key="tenant-42"))
request = model.format_request(messages)
assert request["prompt_cache_key"] == "tenant-42"
def test_prompt_cache_key_absent_when_unset(model_id, max_tokens, messages):
"""A cache_config without cache_key adds no prompt_cache_key to the request."""
model = MistralModel(model_id=model_id, max_tokens=max_tokens, cache_config=CacheConfig())
request = model.format_request(messages)
assert "prompt_cache_key" not in request
def test_ttl_is_ignored_and_warned(model_id, max_tokens, messages):
"""Mistral has no retention control: ttl is dropped with a warning, never sent to the wire."""
model = MistralModel(model_id=model_id, max_tokens=max_tokens, cache_config=CacheConfig(cache_key="k", ttl="1h"))
with pytest.warns(UserWarning, match=r"fields \['ttl'\] have no effect"):
request = model.format_request(messages)
assert request["prompt_cache_key"] == "k"
assert "prompt_cache_retention" not in request
def test_placement_fields_are_no_ops_warned(model_id, max_tokens, messages):
"""strategy / system_prompt_ttl are placement controls Mistral cannot honor; they warn."""
model = MistralModel(
model_id=model_id, max_tokens=max_tokens, cache_config=CacheConfig(strategy="anthropic", cache_key="k")
)
with pytest.warns(UserWarning, match=r"fields \['strategy'\] have no effect"):
request = model.format_request(messages)
assert request["prompt_cache_key"] == "k"
def test_format_request_with_temperature(model, messages, model_id):
model.update_config(temperature=0.8)
@@ -781,7 +837,6 @@ def test_format_request_filters_s3_source_image(model, caplog):
def test_format_request_skips_message_cache_point(model, caplog):
caplog.set_level(logging.WARNING, logger="strands.models.mistral")
messages = [{"role": "user", "content": [{"text": "durable prefix"}, {"cachePoint": {"type": "default"}}]}]
formatted_messages = model._format_request_messages(messages)