Files
EvoScientist-Multi/EvoScientist/middleware/provider_context.py
T

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"]