mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
feat(models): add cache_config to Mistral (#3986)
This commit is contained in:
@@ -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"})
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user