379 lines
13 KiB
Python
379 lines
13 KiB
Python
"""Bound provider context by externalizing assistant-generated inline media."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
import hashlib
|
|
import logging
|
|
import mimetypes
|
|
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
from dataclasses import replace
|
|
from typing import Any
|
|
|
|
from langchain.agents.middleware.types import (
|
|
AgentMiddleware,
|
|
ExtendedModelResponse,
|
|
ModelRequest,
|
|
ModelResponse,
|
|
)
|
|
from langchain_core.messages import AIMessage, BaseMessage
|
|
from langgraph.types import Overwrite
|
|
|
|
from ..llm.contracts import EvoRuntimeError
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_DEFAULT_MAX_INLINE_MEDIA_BYTES = 16_777_216
|
|
_MEDIA_PREFIX = "/artifacts/model-output"
|
|
|
|
|
|
def _decode_base64_block(block: Mapping[str, Any]) -> tuple[bytes, str] | None:
|
|
payload = block.get("base64")
|
|
mime = str(block.get("mime_type") or "application/octet-stream")
|
|
if not isinstance(payload, str) or not payload:
|
|
return None
|
|
try:
|
|
return base64.b64decode(payload, validate=True), mime
|
|
except (ValueError, TypeError) as exc:
|
|
raise EvoRuntimeError(
|
|
"MEDIA_PERSIST_FAILED",
|
|
details=({"reason": "invalid_base64", "media_type": mime},),
|
|
) from exc
|
|
|
|
|
|
def _extension(mime: str) -> str:
|
|
return (mimetypes.guess_extension(mime) or ".bin").lstrip(".")
|
|
|
|
|
|
def _artifact_reference(path: str, mime: str, digest: str) -> dict[str, Any]:
|
|
major = mime.split("/", 1)[0]
|
|
if major in {"image", "audio", "video"}:
|
|
return {
|
|
"type": major,
|
|
"url": path,
|
|
"mime_type": mime,
|
|
"media_id": f"sha256:{digest}",
|
|
}
|
|
return {
|
|
"type": "text",
|
|
"text": f'<generated_file path="{path}" media_type="{mime}" sha256="{digest}" />',
|
|
}
|
|
|
|
|
|
def _provider_reference(block: Mapping[str, Any]) -> dict[str, str] | None:
|
|
block_type = str(block.get("type") or "")
|
|
path = block.get("url")
|
|
if (
|
|
block_type not in {"image", "audio", "video"}
|
|
or not isinstance(path, str)
|
|
or not path.startswith((f"{_MEDIA_PREFIX}/", f"/workspace{_MEDIA_PREFIX}/"))
|
|
):
|
|
return None
|
|
mime = str(block.get("mime_type") or f"{block_type}/unknown")
|
|
media_id = str(block.get("media_id") or "")
|
|
media_id_attribute = f' media_id="{media_id}"' if media_id else ""
|
|
return {
|
|
"type": "text",
|
|
"text": (
|
|
f'<generated_{block_type} path="{path}" media_type="{mime}"'
|
|
f"{media_id_attribute} />"
|
|
),
|
|
}
|
|
|
|
|
|
def _assistant_messages(messages: Sequence[BaseMessage]) -> list[AIMessage]:
|
|
return [message for message in messages if isinstance(message, AIMessage)]
|
|
|
|
|
|
def _collect_inline_media(
|
|
messages: Sequence[BaseMessage],
|
|
*,
|
|
max_inline_media_bytes: int,
|
|
) -> dict[str, tuple[str, bytes]]:
|
|
collected: dict[str, tuple[str, bytes]] = {}
|
|
for message in _assistant_messages(messages):
|
|
content = message.content
|
|
if not isinstance(content, list):
|
|
continue
|
|
for value in content:
|
|
if not isinstance(value, Mapping):
|
|
continue
|
|
decoded = _decode_base64_block(value)
|
|
if decoded is None:
|
|
continue
|
|
raw, mime = decoded
|
|
if len(raw) > max_inline_media_bytes:
|
|
raise EvoRuntimeError(
|
|
"MEDIA_PERSIST_FAILED",
|
|
details=(
|
|
{
|
|
"reason": "media_too_large",
|
|
"media_type": mime,
|
|
"media_bytes": len(raw),
|
|
},
|
|
),
|
|
)
|
|
digest = hashlib.sha256(raw).hexdigest()
|
|
collected.setdefault(digest, (mime, raw))
|
|
return collected
|
|
|
|
|
|
def _rewrite_messages(
|
|
messages: Sequence[BaseMessage],
|
|
paths: Mapping[str, tuple[str, str]],
|
|
*,
|
|
for_provider: bool,
|
|
) -> list[BaseMessage]:
|
|
rewritten: list[BaseMessage] = []
|
|
for message in messages:
|
|
if not isinstance(message, AIMessage) or not isinstance(message.content, list):
|
|
rewritten.append(message)
|
|
continue
|
|
modified = False
|
|
content: list[Any] = []
|
|
for value in message.content:
|
|
if not isinstance(value, Mapping):
|
|
content.append(value)
|
|
continue
|
|
if for_provider:
|
|
reference = _provider_reference(value)
|
|
if reference is not None:
|
|
content.append(reference)
|
|
modified = True
|
|
continue
|
|
decoded = _decode_base64_block(value)
|
|
if decoded is None:
|
|
content.append(value)
|
|
continue
|
|
raw, mime = decoded
|
|
digest = hashlib.sha256(raw).hexdigest()
|
|
path_entry = paths.get(digest)
|
|
if path_entry is None:
|
|
raise EvoRuntimeError(
|
|
"MEDIA_PERSIST_FAILED",
|
|
details=({"reason": "artifact_path_missing", "media_type": mime},),
|
|
)
|
|
path, stored_mime = path_entry
|
|
reference = _artifact_reference(path, stored_mime, digest)
|
|
content.append(
|
|
_provider_reference(reference) if for_provider else reference
|
|
)
|
|
modified = True
|
|
if modified:
|
|
copy = message.model_copy()
|
|
copy.content = content
|
|
rewritten.append(copy)
|
|
else:
|
|
rewritten.append(message)
|
|
return rewritten
|
|
|
|
|
|
def _response_messages(
|
|
response: Any,
|
|
) -> tuple[list[BaseMessage], Callable[[list[BaseMessage]], Any]]:
|
|
if isinstance(response, ExtendedModelResponse):
|
|
return response.model_response.result, lambda result: replace(
|
|
response,
|
|
model_response=replace(response.model_response, result=result),
|
|
)
|
|
if isinstance(response, ModelResponse):
|
|
return response.result, lambda result: replace(response, result=result)
|
|
if isinstance(response, AIMessage):
|
|
return [response], lambda result: result[0]
|
|
return [], lambda _result: response
|
|
|
|
|
|
class ProviderContextMediaMiddleware(AgentMiddleware):
|
|
"""Persist assistant media and keep base64 out of later provider calls."""
|
|
|
|
name = "provider_context_media"
|
|
|
|
def __init__(
|
|
self,
|
|
backend: Any,
|
|
*,
|
|
max_inline_media_bytes: int = _DEFAULT_MAX_INLINE_MEDIA_BYTES,
|
|
media_prefix: str = _MEDIA_PREFIX,
|
|
) -> None:
|
|
self.backend = backend
|
|
self.media_prefix = media_prefix
|
|
self.max_inline_media_bytes = max(1, int(max_inline_media_bytes))
|
|
|
|
def _paths_for(
|
|
self,
|
|
media: Mapping[str, tuple[str, bytes]],
|
|
) -> dict[str, tuple[str, str]]:
|
|
return {
|
|
digest: (f"{self.media_prefix}/{digest[:24]}.{_extension(mime)}", mime)
|
|
for digest, (mime, _raw) in media.items()
|
|
}
|
|
|
|
def _persist(self, messages: Sequence[BaseMessage]) -> dict[str, tuple[str, str]]:
|
|
media = _collect_inline_media(
|
|
messages,
|
|
max_inline_media_bytes=self.max_inline_media_bytes,
|
|
)
|
|
paths = self._paths_for(media)
|
|
for digest, (mime, raw) in media.items():
|
|
path = paths[digest][0]
|
|
responses = self.backend.upload_files([(path, raw)])
|
|
error = (
|
|
getattr(responses[0], "error", None)
|
|
if responses
|
|
else "missing upload response"
|
|
)
|
|
if error:
|
|
raise EvoRuntimeError(
|
|
"MEDIA_PERSIST_FAILED",
|
|
details=({"reason": "artifact_upload_failed", "media_type": mime},),
|
|
)
|
|
logger.info(
|
|
"provider_context_media_persisted path=%s media_type=%s bytes=%s",
|
|
path,
|
|
mime,
|
|
len(raw),
|
|
)
|
|
return paths
|
|
|
|
async def _apersist(
|
|
self, messages: Sequence[BaseMessage]
|
|
) -> dict[str, tuple[str, str]]:
|
|
media = _collect_inline_media(
|
|
messages,
|
|
max_inline_media_bytes=self.max_inline_media_bytes,
|
|
)
|
|
paths = self._paths_for(media)
|
|
for digest, (mime, raw) in media.items():
|
|
path = paths[digest][0]
|
|
responses = await self.backend.aupload_files([(path, raw)])
|
|
error = (
|
|
getattr(responses[0], "error", None)
|
|
if responses
|
|
else "missing upload response"
|
|
)
|
|
if error:
|
|
raise EvoRuntimeError(
|
|
"MEDIA_PERSIST_FAILED",
|
|
details=({"reason": "artifact_upload_failed", "media_type": mime},),
|
|
)
|
|
logger.info(
|
|
"provider_context_media_persisted path=%s media_type=%s bytes=%s",
|
|
path,
|
|
mime,
|
|
len(raw),
|
|
)
|
|
return paths
|
|
|
|
def _prepare_request(self, request: ModelRequest) -> ModelRequest:
|
|
paths = self._persist(request.messages)
|
|
messages = _rewrite_messages(
|
|
request.messages,
|
|
paths,
|
|
for_provider=True,
|
|
)
|
|
return request.override(messages=messages)
|
|
|
|
async def _aprepare_request(self, request: ModelRequest) -> ModelRequest:
|
|
paths = await self._apersist(request.messages)
|
|
messages = _rewrite_messages(
|
|
request.messages,
|
|
paths,
|
|
for_provider=True,
|
|
)
|
|
return request.override(messages=messages)
|
|
|
|
def _prepare_response(self, response: Any) -> Any:
|
|
messages, rebuild = _response_messages(response)
|
|
if not messages:
|
|
return response
|
|
paths = self._persist(messages)
|
|
return rebuild(_rewrite_messages(messages, paths, for_provider=False))
|
|
|
|
async def _aprepare_response(self, response: Any) -> Any:
|
|
messages, rebuild = _response_messages(response)
|
|
if not messages:
|
|
return response
|
|
paths = await self._apersist(messages)
|
|
return rebuild(_rewrite_messages(messages, paths, for_provider=False))
|
|
|
|
def before_model(self, state: Any, runtime: Any) -> dict[str, Any] | None:
|
|
_ = runtime
|
|
try:
|
|
messages = state.get("messages") if isinstance(state, Mapping) else None
|
|
if not isinstance(messages, Sequence) or isinstance(messages, str | bytes):
|
|
return None
|
|
original = list(messages)
|
|
paths = self._persist(original)
|
|
rewritten = _rewrite_messages(original, paths, for_provider=False)
|
|
if all(
|
|
left is right for left, right in zip(original, rewritten, strict=True)
|
|
):
|
|
return None
|
|
logger.info(
|
|
"provider_context_media_checkpoint_repaired messages=%s",
|
|
len(rewritten),
|
|
)
|
|
return {"messages": Overwrite(rewritten)}
|
|
except EvoRuntimeError:
|
|
raise
|
|
except Exception as exc:
|
|
raise self._middleware_failure("before_model", exc) from exc
|
|
|
|
async def abefore_model(self, state: Any, runtime: Any) -> dict[str, Any] | None:
|
|
_ = runtime
|
|
try:
|
|
messages = state.get("messages") if isinstance(state, Mapping) else None
|
|
if not isinstance(messages, Sequence) or isinstance(messages, str | bytes):
|
|
return None
|
|
original = list(messages)
|
|
paths = await self._apersist(original)
|
|
rewritten = _rewrite_messages(original, paths, for_provider=False)
|
|
if all(
|
|
left is right for left, right in zip(original, rewritten, strict=True)
|
|
):
|
|
return None
|
|
logger.info(
|
|
"provider_context_media_checkpoint_repaired messages=%s",
|
|
len(rewritten),
|
|
)
|
|
return {"messages": Overwrite(rewritten)}
|
|
except EvoRuntimeError:
|
|
raise
|
|
except Exception as exc:
|
|
raise self._middleware_failure("before_model", exc) from exc
|
|
|
|
@staticmethod
|
|
def _middleware_failure(node: str, exc: Exception) -> EvoRuntimeError:
|
|
return EvoRuntimeError(
|
|
"AGENT_MIDDLEWARE_FAILED",
|
|
details=(
|
|
{
|
|
"failure_stage": "agent_middleware",
|
|
"middleware": ProviderContextMediaMiddleware.name,
|
|
"middleware_node": (
|
|
f"{ProviderContextMediaMiddleware.name}.{node}"
|
|
),
|
|
"agent_error_type": type(exc).__name__[:128],
|
|
"agent_error_module": type(exc).__module__[:128],
|
|
},
|
|
),
|
|
)
|
|
|
|
def wrap_model_call(
|
|
self,
|
|
request: ModelRequest,
|
|
handler: Callable[[ModelRequest], ModelResponse],
|
|
) -> ModelResponse:
|
|
return self._prepare_response(handler(self._prepare_request(request)))
|
|
|
|
async def awrap_model_call(
|
|
self,
|
|
request: ModelRequest,
|
|
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
|
) -> ModelResponse:
|
|
prepared = await self._aprepare_request(request)
|
|
return await self._aprepare_response(await handler(prepared))
|
|
|
|
|
|
__all__ = ["ProviderContextMediaMiddleware"]
|