5a581c78a2
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
Introduce provider, model, and invocation contracts with encrypted configuration persistence. Add web runtime fencing, route fallback, recovery middleware, workspace scoping, and comprehensive tests.
380 lines
14 KiB
Python
380 lines
14 KiB
Python
"""LangChain chat-model bridge for the stateless Gemini Interactions API."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from collections.abc import AsyncIterator, Mapping, Sequence
|
|
from typing import Any
|
|
|
|
from langchain_core.language_models.chat_models import BaseChatModel
|
|
from langchain_core.messages import (
|
|
AIMessage,
|
|
AIMessageChunk,
|
|
BaseMessage,
|
|
HumanMessage,
|
|
SystemMessage,
|
|
ToolMessage,
|
|
)
|
|
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
|
|
from langchain_core.tools import BaseTool
|
|
from langchain_core.utils.function_calling import convert_to_openai_tool
|
|
from pydantic import Field, SecretStr
|
|
|
|
|
|
class GeminiInteractionsChatModel(BaseChatModel):
|
|
"""Minimal native bridge that preserves signed Provider content blocks."""
|
|
|
|
model_name: str
|
|
api_key: SecretStr
|
|
base_url: str = "https://generativelanguage.googleapis.com"
|
|
max_output_tokens: int = 8192
|
|
temperature: float | None = None
|
|
top_p: float | None = None
|
|
thinking: bool | None = None
|
|
store: bool = False
|
|
bound_tools: tuple[dict[str, Any], ...] = Field(default_factory=tuple)
|
|
|
|
@property
|
|
def _llm_type(self) -> str:
|
|
return "google-gemini-interactions"
|
|
|
|
@property
|
|
def _identifying_params(self) -> dict[str, Any]:
|
|
return {"model_name": self.model_name, "api_mode": "interactions"}
|
|
|
|
def bind_tools(
|
|
self,
|
|
tools: Sequence[dict[str, Any] | type | BaseTool | Any],
|
|
*,
|
|
tool_choice: str | None = None,
|
|
**kwargs: Any,
|
|
) -> Any:
|
|
_ = tool_choice, kwargs
|
|
compiled = []
|
|
for tool in tools:
|
|
value = convert_to_openai_tool(tool)
|
|
function = value.get("function", value)
|
|
compiled.append(
|
|
{
|
|
"type": "function",
|
|
"name": function["name"],
|
|
"description": function.get("description", ""),
|
|
"parameters": function.get(
|
|
"parameters", {"type": "object", "properties": {}}
|
|
),
|
|
}
|
|
)
|
|
return self.model_copy(update={"bound_tools": tuple(compiled)})
|
|
|
|
def _client(self) -> Any:
|
|
from google import genai
|
|
from google.genai import types
|
|
|
|
return genai.Client(
|
|
api_key=self.api_key.get_secret_value(),
|
|
http_options=types.HttpOptions(base_url=self.base_url.rstrip("/")),
|
|
)
|
|
|
|
def _request(self, messages: Sequence[BaseMessage]) -> dict[str, Any]:
|
|
turns, system_instruction = _compile_messages(messages)
|
|
generation_config: dict[str, Any] = {
|
|
"max_output_tokens": self.max_output_tokens,
|
|
}
|
|
if self.temperature is not None:
|
|
generation_config["temperature"] = self.temperature
|
|
if self.top_p is not None:
|
|
generation_config["top_p"] = self.top_p
|
|
if self.thinking is not None:
|
|
generation_config["thinking_level"] = "high" if self.thinking else "minimal"
|
|
generation_config["thinking_summaries"] = (
|
|
"auto" if self.thinking else "none"
|
|
)
|
|
return {
|
|
"model": self.model_name,
|
|
"input": turns,
|
|
"system_instruction": system_instruction or "",
|
|
"generation_config": generation_config,
|
|
"tools": list(self.bound_tools),
|
|
"store": False,
|
|
"stream": False,
|
|
}
|
|
|
|
def _generate(
|
|
self,
|
|
messages: list[BaseMessage],
|
|
stop: list[str] | None = None,
|
|
**kwargs: Any,
|
|
) -> ChatResult:
|
|
request = self._request(messages)
|
|
if stop:
|
|
request["generation_config"]["stop_sequences"] = stop
|
|
response = self._client().interactions.create(**request)
|
|
return _chat_result(response)
|
|
|
|
async def _agenerate(
|
|
self,
|
|
messages: list[BaseMessage],
|
|
stop: list[str] | None = None,
|
|
**kwargs: Any,
|
|
) -> ChatResult:
|
|
request = self._request(messages)
|
|
if stop:
|
|
request["generation_config"]["stop_sequences"] = stop
|
|
response = await self._client().aio.interactions.create(**request)
|
|
return _chat_result(response)
|
|
|
|
async def _astream(
|
|
self,
|
|
messages: list[BaseMessage],
|
|
stop: list[str] | None = None,
|
|
**kwargs: Any,
|
|
) -> AsyncIterator[ChatGenerationChunk]:
|
|
request = self._request(messages)
|
|
request["stream"] = True
|
|
if stop:
|
|
request["generation_config"]["stop_sequences"] = stop
|
|
stream = await self._client().aio.interactions.create(**request)
|
|
blocks: dict[int, dict[str, Any]] = {}
|
|
async for event in stream:
|
|
payload = _dump(event)
|
|
event_type = payload.get("event_type")
|
|
if event_type == "content.start":
|
|
blocks[int(payload["index"])] = dict(payload.get("content") or {})
|
|
continue
|
|
if event_type == "content.delta":
|
|
index = int(payload["index"])
|
|
delta = dict(payload.get("delta") or {})
|
|
block = blocks.setdefault(index, {})
|
|
_merge_stream_delta(block, delta)
|
|
if delta.get("type") == "text" and delta.get("text"):
|
|
yield ChatGenerationChunk(
|
|
message=AIMessageChunk(content=str(delta["text"]))
|
|
)
|
|
continue
|
|
if event_type == "content.stop":
|
|
block = blocks.get(int(payload["index"]), {})
|
|
if block.get("type") == "function_call":
|
|
yield ChatGenerationChunk(
|
|
message=AIMessageChunk(
|
|
content="",
|
|
tool_call_chunks=[
|
|
{
|
|
"id": str(block.get("id") or ""),
|
|
"name": str(block.get("name") or ""),
|
|
"args": json.dumps(
|
|
block.get("arguments") or {},
|
|
separators=(",", ":"),
|
|
),
|
|
"index": int(payload["index"]),
|
|
"type": "tool_call_chunk",
|
|
}
|
|
],
|
|
)
|
|
)
|
|
continue
|
|
if event_type == "error":
|
|
raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR")
|
|
if event_type == "interaction.complete":
|
|
interaction = payload.get("interaction") or {}
|
|
ordered_blocks = [blocks[index] for index in sorted(blocks)]
|
|
yield ChatGenerationChunk(
|
|
message=AIMessageChunk(
|
|
content="",
|
|
additional_kwargs={
|
|
"gemini_interaction_content": ordered_blocks
|
|
},
|
|
usage_metadata=_usage_metadata(
|
|
interaction.get("usage"),
|
|
provider_request_id=interaction.get("id"),
|
|
),
|
|
response_metadata={
|
|
"model_name": str(
|
|
(interaction.get("model") or {}).get("id") or ""
|
|
),
|
|
"finish_reason": str(
|
|
interaction.get("status") or "unknown"
|
|
),
|
|
},
|
|
)
|
|
)
|
|
|
|
|
|
def _message_content(message: BaseMessage) -> list[dict[str, Any]]:
|
|
if isinstance(message, AIMessage):
|
|
preserved = message.additional_kwargs.get("gemini_interaction_content")
|
|
if isinstance(preserved, list):
|
|
return [dict(item) for item in preserved]
|
|
content = message.content
|
|
if isinstance(content, str):
|
|
blocks: list[dict[str, Any]] = [{"type": "text", "text": content}]
|
|
elif isinstance(content, list):
|
|
blocks = []
|
|
for item in content:
|
|
if isinstance(item, str):
|
|
blocks.append({"type": "text", "text": item})
|
|
elif isinstance(item, dict) and item.get("type") in {
|
|
"text",
|
|
"thought",
|
|
"function_call",
|
|
"function_result",
|
|
"image",
|
|
"audio",
|
|
"video",
|
|
"document",
|
|
}:
|
|
blocks.append(dict(item))
|
|
else:
|
|
raise ValueError("MODEL_CONTENT_BLOCK_UNSUPPORTED")
|
|
else:
|
|
blocks = [{"type": "text", "text": str(content)}]
|
|
if isinstance(message, AIMessage):
|
|
for call in message.tool_calls:
|
|
blocks.append(
|
|
{
|
|
"type": "function_call",
|
|
"id": str(call["id"]),
|
|
"name": str(call["name"]),
|
|
"arguments": dict(call.get("args") or {}),
|
|
}
|
|
)
|
|
return blocks
|
|
|
|
|
|
def _compile_messages(
|
|
messages: Sequence[BaseMessage],
|
|
) -> tuple[list[dict[str, Any]], str]:
|
|
turns: list[dict[str, Any]] = []
|
|
system_parts: list[str] = []
|
|
for message in messages:
|
|
if isinstance(message, SystemMessage):
|
|
system_parts.append(str(message.content))
|
|
continue
|
|
if isinstance(message, ToolMessage):
|
|
turns.append(
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "function_result",
|
|
"call_id": str(message.tool_call_id),
|
|
"name": str(getattr(message, "name", "") or ""),
|
|
"result": message.content,
|
|
}
|
|
],
|
|
}
|
|
)
|
|
continue
|
|
role = "model" if isinstance(message, AIMessage) else "user"
|
|
if not isinstance(message, (AIMessage, HumanMessage)):
|
|
role = "user"
|
|
turns.append({"role": role, "content": _message_content(message)})
|
|
return turns, "\n\n".join(system_parts)
|
|
|
|
|
|
def _chat_result(response: Any) -> ChatResult:
|
|
blocks = [
|
|
item.model_dump(mode="json", by_alias=True, exclude_none=True)
|
|
if hasattr(item, "model_dump")
|
|
else dict(item)
|
|
for item in (getattr(response, "outputs", None) or [])
|
|
]
|
|
tool_calls = [
|
|
{
|
|
"id": str(item.get("id") or ""),
|
|
"name": str(item.get("name") or ""),
|
|
"args": dict(item.get("arguments") or {}),
|
|
"type": "tool_call",
|
|
}
|
|
for item in blocks
|
|
if item.get("type") == "function_call"
|
|
]
|
|
usage_metadata = _usage_metadata(
|
|
_dump(getattr(response, "usage", None)),
|
|
provider_request_id=getattr(response, "id", None),
|
|
)
|
|
message = AIMessage(
|
|
content=blocks,
|
|
tool_calls=tool_calls,
|
|
additional_kwargs={"gemini_interaction_content": blocks},
|
|
usage_metadata=usage_metadata,
|
|
response_metadata={
|
|
"model_name": str(getattr(getattr(response, "model", None), "id", "")),
|
|
"finish_reason": str(getattr(response, "status", "unknown")),
|
|
},
|
|
)
|
|
return ChatResult(generations=[ChatGeneration(message=message)])
|
|
|
|
|
|
def _dump(value: Any) -> dict[str, Any]:
|
|
if value is None:
|
|
return {}
|
|
if hasattr(value, "model_dump"):
|
|
return value.model_dump(mode="json", by_alias=True, exclude_none=True)
|
|
if isinstance(value, Mapping):
|
|
return dict(value)
|
|
return {}
|
|
|
|
|
|
def _merge_stream_delta(block: dict[str, Any], delta: Mapping[str, Any]) -> None:
|
|
kind = str(delta.get("type") or "")
|
|
if kind == "text":
|
|
block["type"] = "text"
|
|
block["text"] = str(block.get("text") or "") + str(delta.get("text") or "")
|
|
elif kind == "thought_signature":
|
|
block.setdefault("type", "thought")
|
|
block["signature"] = delta.get("signature")
|
|
elif kind == "thought_summary":
|
|
block.setdefault("type", "thought")
|
|
content = delta.get("content")
|
|
if content is not None:
|
|
block.setdefault("summary", []).append(content)
|
|
elif kind == "text_annotation":
|
|
block.setdefault("annotations", []).extend(delta.get("annotations") or [])
|
|
else:
|
|
block.update(delta)
|
|
|
|
|
|
def _usage_metadata(
|
|
usage: Mapping[str, Any] | None, *, provider_request_id: Any
|
|
) -> dict[str, Any] | None:
|
|
if not usage:
|
|
return None
|
|
required = (
|
|
usage.get("total_input_tokens"),
|
|
usage.get("total_cached_tokens"),
|
|
usage.get("total_output_tokens"),
|
|
)
|
|
if any(value is None for value in required):
|
|
return None
|
|
result: dict[str, Any] = {
|
|
"input_tokens": int(required[0]),
|
|
"cached_input_tokens": int(required[1]),
|
|
"output_tokens": int(required[2]),
|
|
"total_tokens": int(
|
|
usage.get("total_tokens")
|
|
if usage.get("total_tokens") is not None
|
|
else int(required[0]) + int(required[2])
|
|
),
|
|
"usage_finality": "confirmed",
|
|
}
|
|
optional = {
|
|
"reasoning_tokens": usage.get("total_thought_tokens"),
|
|
"provider_request_id": provider_request_id,
|
|
}
|
|
result.update({key: value for key, value in optional.items() if value is not None})
|
|
return result
|
|
|
|
|
|
def create_gemini_interactions_model(**kwargs: Any) -> GeminiInteractionsChatModel:
|
|
return GeminiInteractionsChatModel(
|
|
model_name=str(kwargs["model"]),
|
|
api_key=SecretStr(str(kwargs["api_key"])),
|
|
base_url=str(
|
|
kwargs.get("base_url") or "https://generativelanguage.googleapis.com"
|
|
),
|
|
max_output_tokens=int(kwargs.get("max_output_tokens") or 8192),
|
|
temperature=kwargs.get("temperature"),
|
|
top_p=kwargs.get("top_p"),
|
|
thinking=kwargs.get("thinking"),
|
|
)
|