diff --git a/astrbot/core/agent/context/compressor.py b/astrbot/core/agent/context/compressor.py index 759604dd93..00d9c61d61 100644 --- a/astrbot/core/agent/context/compressor.py +++ b/astrbot/core/agent/context/compressor.py @@ -130,6 +130,7 @@ def __init__( instruction_text: str | None = None, compression_threshold: float = 0.82, token_counter: TokenCounter | None = None, + conversation_id: str | None = None, ) -> None: """Initialize the LLM summary compressor. @@ -139,8 +140,11 @@ def __init__( exact context. Clamped to 0-0.3. instruction_text: Custom instruction for summary generation. compression_threshold: The compression trigger threshold (default: 0.82). + token_counter: Token counter used to preserve recent context. + conversation_id: Conversation UUID for the summary request. """ self.provider = provider + self.conversation_id = conversation_id self.keep_recent_ratio = min(max(float(keep_recent_ratio), 0.0), 0.3) self.compression_threshold = compression_threshold self.token_counter = token_counter or EstimateTokenCounter() @@ -275,6 +279,7 @@ async def __call__(self, messages: list[Message]) -> list[Message]: try: response = await self.provider.text_chat( contexts=sanitized_summary_contexts, + conversation_id=self.conversation_id, ) summary_content = (response.completion_text or "").strip() except Exception as e: diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index 1a11ebff96..66b365a2dd 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -13,6 +13,7 @@ class ContextManager: def __init__( self, config: ContextConfig, + conversation_id: str | None = None, ) -> None: """Initialize the context manager. @@ -22,6 +23,7 @@ def __init__( Args: config: The context configuration. + conversation_id: Conversation UUID forwarded to summary requests. """ self.config = config @@ -36,6 +38,7 @@ def __init__( keep_recent_ratio=config.llm_compress_keep_recent_ratio, instruction_text=config.llm_compress_instruction, token_counter=self.token_counter, + conversation_id=conversation_id, ) else: self.compressor = TruncateByTurnsCompressor( diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index c9787ed6f0..b5d9c3c73c 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -229,9 +229,18 @@ async def reset( request_max_retries: int | None = None, tool_result_overflow_dir: str | None = None, read_tool: FunctionTool | None = None, + # stable identity for plugin-managed conversations when + # request.conversation is None (e.g. Context.tool_loop_agent) + conversation_id: str | None = None, **kwargs: T.Any, ) -> None: self.req = request + # Transient agents need one identity across tool calls and summary requests. + self._conversation_id = ( + request.conversation.cid + if request.conversation is not None + else (conversation_id or uuid.uuid4().hex) + ) self.streaming = streaming self.enforce_max_turns = enforce_max_turns self.llm_compress_instruction = llm_compress_instruction @@ -257,7 +266,8 @@ async def reset( custom_compressor=self.custom_compressor, ) self.request_context_manager = ContextManager( - self.request_context_manager_config + self.request_context_manager_config, + conversation_id=self._conversation_id, ) self.provider = provider @@ -505,6 +515,7 @@ async def _iter_llm_responses( "contexts": self._sanitize_contexts_for_provider(self.run_context.messages), "func_tool": self._func_tool_for_provider(), "session_id": self.req.session_id, + "conversation_id": self._conversation_id, "extra_user_content_parts": self.req.extra_user_content_parts, # list[ContentPart] "abort_signal": self._abort_signal, "request_max_retries": self.request_max_retries, @@ -1447,6 +1458,7 @@ async def _resolve_tool_exec( func_tool=param_subset, model=self.req.model, session_id=self.req.session_id, + conversation_id=self._conversation_id, extra_user_content_parts=self.req.extra_user_content_parts, # tool_choice="required", abort_signal=self._abort_signal, @@ -1479,6 +1491,7 @@ async def _resolve_tool_exec( func_tool=param_subset, model=self.req.model, session_id=self.req.session_id, + conversation_id=self._conversation_id, extra_user_content_parts=self.req.extra_user_content_parts, # tool_choice="required", abort_signal=self._abort_signal, diff --git a/astrbot/core/astr_main_agent.py b/astrbot/core/astr_main_agent.py index ce86e2c18e..1fb9681aef 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -692,6 +692,7 @@ async def _request_img_caption( cfg: dict, image_urls: list[str], plugin_context: Context, + conversation_id: str | None = None, ) -> str: prov = plugin_context.get_provider_by_id(provider_id) if prov is None: @@ -711,6 +712,7 @@ async def _request_img_caption( llm_resp = await prov.text_chat( prompt=img_cap_prompt, image_urls=image_urls, + conversation_id=conversation_id, ) return llm_resp.completion_text @@ -728,6 +730,7 @@ async def _ensure_img_caption( cfg, req.image_urls, plugin_context, + conversation_id=req.conversation.cid if req.conversation else None, ) if caption: req.extra_user_content_parts.append( @@ -881,6 +884,9 @@ async def _process_quote_message( llm_resp = await prov.text_chat( prompt="Please describe the image content.", image_urls=[image_ref], + conversation_id=req.conversation.cid + if req.conversation + else None, ) if llm_resp.completion_text: content_parts.append( @@ -1019,6 +1025,7 @@ async def _handle_webchat( try: llm_resp = await prov.text_chat( + conversation_id=req.conversation.cid if req.conversation else None, system_prompt=( "You are a conversation title generator. " "Generate a concise title in the same language as the user’s input, " diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 9cc95a13ed..aaf1f19ffb 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -1315,6 +1315,42 @@ "proxy": "", "custom_headers": {}, }, + "OpenCode Go Chat Completions": { + "id": "opencode-go", + "provider": "opencode-go", + "type": "opencode_go_chat_completion", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://opencode.ai/zen/go/v1", + "timeout": 120, + "proxy": "", + "custom_headers": {}, + }, + "OpenCode Go Responses": { + "id": "opencode-go-responses", + "provider": "opencode-go", + "type": "opencode_go_responses", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://opencode.ai/zen/go/v1", + "timeout": 120, + "proxy": "", + "custom_headers": {}, + }, + "OpenCode Go Messages": { + "id": "opencode-go-messages", + "provider": "opencode-go", + "type": "opencode_go_messages", + "provider_type": "chat_completion", + "enable": True, + "key": [], + "api_base": "https://opencode.ai/zen/go/v1", + "timeout": 120, + "proxy": "", + "custom_headers": {}, + }, "Google Gemini": { "id": "google_gemini", "provider": "google", diff --git a/astrbot/core/provider/manager.py b/astrbot/core/provider/manager.py index 60044fb863..c3378c3934 100644 --- a/astrbot/core/provider/manager.py +++ b/astrbot/core/provider/manager.py @@ -455,6 +455,14 @@ def dynamic_import_provider(self, type: str) -> None: from .sources.mirarouter_source import ( ProviderMiraRouter as ProviderMiraRouter, ) + case ( + "opencode_go_chat_completion" + | "opencode_go_responses" + | "opencode_go_messages" + ): + from .sources.opencode_go_source import ( + ProviderOpenCodeGo as ProviderOpenCodeGo, + ) case "openrouter_chat_completion": from .sources.openrouter_source import ( ProviderOpenRouter as ProviderOpenRouter, diff --git a/astrbot/core/provider/sources/anthropic_source.py b/astrbot/core/provider/sources/anthropic_source.py index 5a21737b44..fbc861646d 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -812,6 +812,8 @@ async def text_chat( model = model or self.get_model() payloads = {"messages": new_messages, "model": model} + if extra_headers := kwargs.get("extra_headers"): + payloads["extra_headers"] = extra_headers if func_tool and not func_tool.empty(): payloads["tool_choice"] = tool_choice @@ -884,6 +886,8 @@ async def text_chat_stream( model = model or self.get_model() payloads = {"messages": new_messages, "model": model} + if extra_headers := kwargs.get("extra_headers"): + payloads["extra_headers"] = extra_headers if func_tool and not func_tool.empty(): payloads["tool_choice"] = tool_choice diff --git a/astrbot/core/provider/sources/openai_responses_source.py b/astrbot/core/provider/sources/openai_responses_source.py index c5cb9bdb82..aa8cf7a1aa 100644 --- a/astrbot/core/provider/sources/openai_responses_source.py +++ b/astrbot/core/provider/sources/openai_responses_source.py @@ -241,6 +241,7 @@ async def _prepare_chat_payload( tool_calls_result: ToolCallsResult | list[ToolCallsResult] | None = None, model: str | None = None, extra_user_content_parts: list[ContentPart] | None = None, + extra_headers: dict[str, str] | None = None, **kwargs: Any, ) -> tuple[dict, list[dict]]: """Build a stateless Responses API payload and replayable context. @@ -254,6 +255,7 @@ async def _prepare_chat_payload( tool_calls_result: Function calls and their returned outputs. model: Optional per-request model override. extra_user_content_parts: Additional user content blocks. + extra_headers: HTTP headers applied only to this request. **kwargs: Reserved provider request arguments. Returns: @@ -291,6 +293,8 @@ async def _prepare_chat_payload( } if system_prompt: payloads["instructions"] = system_prompt + if extra_headers: + payloads["extra_headers"] = extra_headers return payloads, context_query diff --git a/astrbot/core/provider/sources/openai_source.py b/astrbot/core/provider/sources/openai_source.py index 15fd6b72f4..9a05600015 100644 --- a/astrbot/core/provider/sources/openai_source.py +++ b/astrbot/core/provider/sources/openai_source.py @@ -950,6 +950,7 @@ async def _prepare_chat_payload( tool_calls_result: ToolCallsResult | list[ToolCallsResult] | None = None, model: str | None = None, extra_user_content_parts: list[ContentPart] | None = None, + extra_headers: dict[str, str] | None = None, **kwargs, ) -> tuple: """准备聊天所需的有效载荷和上下文""" @@ -987,6 +988,8 @@ async def _prepare_chat_payload( model = model or self.get_model() payloads = {"messages": context_query, "model": model} + if extra_headers: + payloads["extra_headers"] = extra_headers self._finally_convert_payload(payloads) diff --git a/astrbot/core/provider/sources/opencode_go_source.py b/astrbot/core/provider/sources/opencode_go_source.py new file mode 100644 index 0000000000..ebeb6b2b56 --- /dev/null +++ b/astrbot/core/provider/sources/opencode_go_source.py @@ -0,0 +1,194 @@ +import hashlib +from collections.abc import AsyncGenerator +from uuid import uuid4 + +from astrbot.core.provider.entities import LLMResponse +from astrbot.core.provider.provider import Provider + +from ..register import register_provider_adapter +from .anthropic_source import ProviderAnthropic +from .openai_responses_source import ProviderOpenAIResponses +from .openai_source import ProviderOpenAIOfficial + +OPENCODE_GO_API_BASE = "https://opencode.ai/zen/go/v1" + + +@register_provider_adapter( + "opencode_go_chat_completion", "OpenCode Go Chat Completions Provider Adapter" +) +class ProviderOpenCodeGo(Provider): + """Send Go requests using the explicitly selected protocol adapter.""" + + ADAPTER: type[Provider] = ProviderOpenAIOfficial + + def __init__(self, provider_config: dict, provider_settings: dict) -> None: + super().__init__(provider_config, provider_settings) + self.set_model(provider_config.get("model") or "unknown") + config = dict(provider_config) + config["api_base"] = config.get("api_base") or OPENCODE_GO_API_BASE + config["model"] = self.get_model().removeprefix("opencode-go/") + config["custom_headers"] = { + key: value + for key, value in self.request_headers.items() + if key.lower() != "x-opencode-session" + } + self.delegate = self.ADAPTER(config, provider_settings) + + def get_current_key(self) -> str: + return self.delegate.get_current_key() + + def get_keys(self) -> list[str]: + return self.delegate.get_keys() + + def set_key(self, key: str) -> None: + self.delegate.set_key(key) + + async def get_models(self) -> list[str]: + return await self.delegate.get_models() + + async def text_chat( + self, + prompt=None, + session_id=None, + image_urls=None, + audio_urls=None, + func_tool=None, + contexts=None, + system_prompt=None, + tool_calls_result=None, + model=None, + extra_user_content_parts=None, + tool_choice="auto", + **kwargs, + ) -> LLMResponse: + """Send a request using the conversation UUID as the session identity. + + Args: + prompt: Current user prompt. + session_id: Deprecated provider argument, forwarded for compatibility. + image_urls: Images attached to the request. + audio_urls: Audio attached to the request. + func_tool: Available function tools. + contexts: Conversation history. + system_prompt: System instructions. + tool_calls_result: Results of previous tool calls. + model: Optional per-request model override. + extra_user_content_parts: Additional user content blocks. + tool_choice: Whether tool use is automatic or required. + **kwargs: Optional conversation_id (AstrBot conversation UUID) and + additional arguments forwarded to the protocol adapter. + + Returns: + The normalized model response. + """ + model = (model or self.get_model()).removeprefix("opencode-go/") + # Normalize UA casing for SDK merging; blank overrides keep client defaults. + extra_headers = { + ("User-Agent" if key.lower() == "user-agent" else key): value + for key, value in (kwargs.pop("extra_headers", None) or {}).items() + if key.lower() != "x-opencode-session" + and (key.lower() != "user-agent" or str(value).strip()) + } + # Calls without a conversation (such as connection tests) are independent. + extra_headers["x-opencode-session"] = hashlib.sha256( + (kwargs.pop("conversation_id", None) or uuid4().hex).encode() + ).hexdigest() + return await self.delegate.text_chat( + prompt=prompt, + session_id=session_id, + image_urls=image_urls, + audio_urls=audio_urls, + func_tool=func_tool, + contexts=contexts, + system_prompt=system_prompt, + tool_calls_result=tool_calls_result, + model=model, + extra_user_content_parts=extra_user_content_parts, + tool_choice=tool_choice, + extra_headers=extra_headers, + **kwargs, + ) + + async def text_chat_stream( + self, + prompt=None, + session_id=None, + image_urls=None, + audio_urls=None, + func_tool=None, + contexts=None, + system_prompt=None, + tool_calls_result=None, + model=None, + extra_user_content_parts=None, + tool_choice="auto", + **kwargs, + ) -> AsyncGenerator[LLMResponse, None]: + """Stream a response with the same session identity as non-streaming calls. + + Args: + prompt: Current user prompt. + session_id: Deprecated provider argument, forwarded for compatibility. + image_urls: Images attached to the request. + audio_urls: Audio attached to the request. + func_tool: Available function tools. + contexts: Conversation history. + system_prompt: System instructions. + tool_calls_result: Results of previous tool calls. + model: Optional per-request model override. + extra_user_content_parts: Additional user content blocks. + tool_choice: Whether tool use is automatic or required. + **kwargs: Optional conversation_id (AstrBot conversation UUID) and + additional arguments forwarded to the protocol adapter. + + Yields: + Normalized response chunks. + """ + model = (model or self.get_model()).removeprefix("opencode-go/") + # Normalize UA casing for SDK merging; blank overrides keep client defaults. + extra_headers = { + ("User-Agent" if key.lower() == "user-agent" else key): value + for key, value in (kwargs.pop("extra_headers", None) or {}).items() + if key.lower() != "x-opencode-session" + and (key.lower() != "user-agent" or str(value).strip()) + } + extra_headers["x-opencode-session"] = hashlib.sha256( + (kwargs.pop("conversation_id", None) or uuid4().hex).encode() + ).hexdigest() + async for response in self.delegate.text_chat_stream( + prompt=prompt, + session_id=session_id, + image_urls=image_urls, + audio_urls=audio_urls, + func_tool=func_tool, + contexts=contexts, + system_prompt=system_prompt, + tool_calls_result=tool_calls_result, + model=model, + extra_user_content_parts=extra_user_content_parts, + tool_choice=tool_choice, + extra_headers=extra_headers, + **kwargs, + ): + yield response + + async def terminate(self) -> None: + await self.delegate.terminate() + + +@register_provider_adapter( + "opencode_go_responses", "OpenCode Go Responses Provider Adapter" +) +class ProviderOpenCodeGoResponses(ProviderOpenCodeGo): + """Use Go's Responses endpoint for user-selected models.""" + + ADAPTER = ProviderOpenAIResponses + + +@register_provider_adapter( + "opencode_go_messages", "OpenCode Go Messages Provider Adapter" +) +class ProviderOpenCodeGoMessages(ProviderOpenCodeGo): + """Use Go's Messages endpoint for user-selected models.""" + + ADAPTER = ProviderAnthropic diff --git a/astrbot/core/star/context.py b/astrbot/core/star/context.py index b4f6e61c48..becde7100c 100644 --- a/astrbot/core/star/context.py +++ b/astrbot/core/star/context.py @@ -245,6 +245,7 @@ async def tool_loop_agent( stream: bool - whether to stream the LLM response agent_hooks: BaseAgentRunHooks[AstrAgentContext] - hooks to run during agent execution agent_context: AstrAgentContext - context to use for the agent + conversation_id: str - stable identity for a plugin-managed conversation; without it each call gets a random one other kwargs will be DIRECTLY passed to the runner.reset() method diff --git a/docs/en/providers/opencode-go.md b/docs/en/providers/opencode-go.md new file mode 100644 index 0000000000..a70b1e4c40 --- /dev/null +++ b/docs/en/providers/opencode-go.md @@ -0,0 +1,30 @@ +# Connect OpenCode Go + +[OpenCode Go](https://opencode.ai/docs/go/) is a model subscription service for coding agents. + +## Get an API Key + +Open the [OpenCode console](https://opencode.ai/auth), subscribe to Go, and copy your API key. + +## Configure AstrBot + +Open the AstrBot dashboard and go to **Providers → Add Provider**. Select **OpenCode Go Chat Completions**, **OpenCode Go Responses**, or **OpenCode Go Messages** according to the model's API format in the [OpenCode Go documentation](https://opencode.ai/docs/go/#endpoints). + +| Field | Value | +| --- | --- | +| API Base URL | `https://opencode.ai/zen/go/v1` | +| API Key | The API key obtained from the OpenCode console | + +Save the provider, then open its card and add the models you want to use. + +## Request Headers + +The default `User-Agent` is `astrbot/`. Non-blank values in per-request `extra_headers` take priority over provider `custom_headers`, which take priority over the default. User-Agent header names are case-insensitive; blank values fall back to the next level. + +`x-opencode-session` is generated automatically from the conversation ID and cannot be overridden. Direct calls without a conversation ID get an independent session. + +## Set as Default + +Go to **Settings → Provider Settings**, select the OpenCode Go model you just added as the default chat model, and save the configuration. + +For supported models, usage requirements, and limits, see the [OpenCode Go documentation](https://opencode.ai/docs/go/). diff --git a/docs/zh/providers/opencode-go.md b/docs/zh/providers/opencode-go.md new file mode 100644 index 0000000000..5aab49d745 --- /dev/null +++ b/docs/zh/providers/opencode-go.md @@ -0,0 +1,30 @@ +# 接入 OpenCode Go + +[OpenCode Go](https://opencode.ai/docs/go/) 是面向编程代理的模型订阅服务。 + +## 获取 API Key + +前往 [OpenCode 控制台](https://opencode.ai/auth),订阅 Go 并复制 API Key。 + +## 在 AstrBot 中配置 + +打开 AstrBot 管理面板,进入 **服务提供商 → 新增提供商**,根据 [OpenCode Go 文档](https://opencode.ai/docs/go/#endpoints)中模型的接口类型选择 **OpenCode Go Chat Completions**、**OpenCode Go Responses** 或 **OpenCode Go Messages**。 + +| 配置项 | 值 | +| --- | --- | +| API Base URL | `https://opencode.ai/zen/go/v1` | +| API Key | 在 OpenCode 控制台获取的 API Key | + +保存后,点击提供商卡片,添加需要使用的模型。 + +## 请求头 + +默认 `User-Agent` 为 `astrbot/<版本>`。非空白值的优先级为:单次请求的 `extra_headers` → 提供商的 `custom_headers` → 默认值。User-Agent 键名不区分大小写,空白值回退到下一层。 + +`x-opencode-session` 根据对话 ID 自动生成,不允许手动覆盖。未传入对话 ID 的直接调用使用独立会话标识。 + +## 设为默认模型 + +进入 **配置文件 → 提供商设置**,将「默认聊天模型」设置为刚刚添加的 OpenCode Go 模型,然后保存配置。 + +支持的模型、使用要求和额度请参阅 [OpenCode Go 文档](https://opencode.ai/docs/go/)。 diff --git a/tests/test_opencode_go_source.py b/tests/test_opencode_go_source.py new file mode 100644 index 0000000000..3b494f4e57 --- /dev/null +++ b/tests/test_opencode_go_source.py @@ -0,0 +1,410 @@ +import asyncio +import copy +import hashlib +import json +from functools import partial +from uuid import uuid4 + +import httpx +import pytest +from anthropic import _base_client as anthropic_base_client +from openai import _base_client as openai_base_client + +from astrbot.core.agent.context.config import ContextConfig +from astrbot.core.agent.context.manager import ContextManager +from astrbot.core.agent.message import Message +from astrbot.core.config.default import CONFIG_METADATA_2 +from astrbot.core.provider.headers import DEFAULT_USER_AGENT +from astrbot.core.provider.sources.anthropic_source import ProviderAnthropic +from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial +from astrbot.core.provider.sources.opencode_go_source import ( + ProviderOpenCodeGo, + ProviderOpenCodeGoMessages, + ProviderOpenCodeGoResponses, +) + +GO_PROTOCOL_CASES = [ + (ProviderOpenCodeGo, "chat/completions"), + (ProviderOpenCodeGoResponses, "responses"), + (ProviderOpenCodeGoMessages, "messages"), +] + + +@pytest.fixture +def go_http(monkeypatch): + """Capture real SDK HTTP requests with deterministic protocol responses.""" + requests = [] + + async def handle(request, *, httpx_module): + requests.append(request) + await asyncio.sleep(0) + if request.method == "GET": + return httpx_module.Response(200, json={"data": [{"id": "kimi-k2.6"}]}) + body = json.loads(request.content) + model = body["model"] + if request.url.path.endswith("/chat/completions"): + response = { + "id": "chat-1", + "object": "chat.completion", + "created": 1, + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + } + events = [ + { + **response, + "object": "chat.completion.chunk", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + } + ] + elif request.url.path.endswith("/responses"): + response = { + "id": "resp-1", + "object": "response", + "created_at": 1, + "model": model, + "status": "completed", + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "output": [ + { + "type": "message", + "id": "msg-1", + "role": "assistant", + "status": "completed", + "content": [ + {"type": "output_text", "text": "ok", "annotations": []} + ], + } + ], + } + events = [ + { + "type": "response.completed", + "response": response, + "sequence_number": 0, + } + ] + else: + assert request.url.path == "/zen/go/v1/messages" + response = { + "id": "msg-1", + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + } + events = [ + { + "type": "message_start", + "message": {**response, "content": [], "stop_reason": None}, + }, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "ok"}, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 1}, + }, + {"type": "message_stop"}, + ] + if body.get("stream"): + content = "".join( + f"event: {event.get('type', 'message')}\ndata: {json.dumps(event)}\n\n" + for event in events + ) + return httpx_module.Response( + 200, text=content, headers={"Content-Type": "text/event-stream"} + ) + return httpx_module.Response(200, json=response) + + def client(provider, _config): + sdk = ( + anthropic_base_client + if isinstance(provider, ProviderAnthropic) + else openai_base_client + ) + httpx_module = getattr(sdk, "httpx", getattr(sdk, "httpx2", httpx)) + return httpx_module.AsyncClient( + transport=httpx_module.MockTransport( + partial(handle, httpx_module=httpx_module) + ) + ) + + monkeypatch.setattr(ProviderOpenAIOfficial, "_create_http_client", client) + monkeypatch.setattr(ProviderAnthropic, "_create_http_client", client) + return requests + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("provider_class,endpoint", GO_PROTOCOL_CASES) +async def test_go_http_identity_and_concurrent_sessions( + go_http, provider_class, endpoint, streaming +): + model = "new-model-without-routing-metadata" + provider = provider_class( + { + "key": ["test-key"], + "custom_headers": { + "user-agent": "configured/1.0", + "X-OpenCode-Session": "wrong", + "X-Custom": "keep", + }, + }, + {}, + ) + sessions = [str(uuid4()) for _ in range(4)] + expected_session_ids = [ + hashlib.sha256(conversation_id.encode()).hexdigest() + for conversation_id in sessions + ] + expected_user_agents = { + session_id: f"request/{conversation_id}" + for session_id, conversation_id in zip(expected_session_ids, sessions) + } + + async def send(conversation_id): + kwargs = { + "prompt": "Write a Python function", + "conversation_id": conversation_id, + "model": f"opencode-go/{model}", + "extra_headers": { + "X-Request-Test": "request-header", + "uSeR-aGeNt": f"request/{conversation_id}", + "X-OPENCODE-SESSION": "wrong-request-session", + }, + } + if streaming: + result = [item async for item in provider.text_chat_stream(**kwargs)] + assert any(item.completion_text == "ok" for item in result) + else: + assert (await provider.text_chat(**kwargs)).completion_text == "ok" + + try: + await asyncio.gather(*(send(conversation_id) for conversation_id in sessions)) + await send(sessions[0]) + assert len(go_http) == 5 + actual_session_ids = [r.headers["x-opencode-session"] for r in go_http] + assert set(actual_session_ids) == set(expected_session_ids) + assert actual_session_ids[-1] == expected_session_ids[0] + assert actual_session_ids.count(expected_session_ids[0]) == 2 + for request in go_http: + assert request.url.path == f"/zen/go/v1/{endpoint}" + assert request.headers.get_list("user-agent") == [ + expected_user_agents[request.headers["x-opencode-session"]] + ] + assert len(request.headers.get_list("x-opencode-session")) == 1 + assert request.headers["x-custom"] == "keep" + assert request.headers["x-request-test"] == "request-header" + body = json.loads(request.content) + assert body["model"] == model + assert "extra_headers" not in body + assert "x-opencode-session" not in body + assert "X-Request-Test" not in body + finally: + await provider.terminate() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("provider_class,endpoint", GO_PROTOCOL_CASES) +@pytest.mark.parametrize( + "custom_headers,extra_headers,expected_user_agent", + [ + (None, None, DEFAULT_USER_AGENT), + ({}, {}, DEFAULT_USER_AGENT), + ({"user-agent": "configured/1.0"}, {}, "configured/1.0"), + ({"USER-AGENT": " "}, {}, DEFAULT_USER_AGENT), + ({}, {"user-agent": "request/1.0"}, "request/1.0"), + ( + {"uSeR-aGeNt": "configured/1.0"}, + {"User-Agent": "request/1.0"}, + "request/1.0", + ), + ( + {"User-Agent": "configured/1.0"}, + {"USER-AGENT": "request/1.0"}, + "request/1.0", + ), + ( + {"User-Agent": "configured/1.0"}, + {"User-Agent": " "}, + "configured/1.0", + ), + ( + {"User-Agent": "configured/1.0"}, + {"uSeR-aGeNt": ""}, + "configured/1.0", + ), + ({"user-agent": " "}, {"USER-AGENT": " "}, DEFAULT_USER_AGENT), + ( + {"user-agent": "configured/1.0", "X-OpenCode-Session": "wrong"}, + {"X-Request-Test": "keep"}, + "configured/1.0", + ), + ( + {"User-Agent": "first/1.0", "USER-AGENT": "configured/1.0"}, + { + "user-agent": "first/1.0", + "User-Agent": "request/1.0", + "USER-AGENT": " ", + }, + "request/1.0", + ), + ], +) +async def test_go_user_agent_defaults_and_overrides( + go_http, + provider_class, + endpoint, + streaming, + custom_headers, + extra_headers, + expected_user_agent, +): + """Keep one UA per request without changing configuration or client defaults.""" + config = {"key": ["test-key"], "custom_headers": custom_headers} + original_config = copy.deepcopy(config) + original_extra_headers = copy.deepcopy(extra_headers) + provider = provider_class(config, {}) + default_headers = dict(provider.delegate.request_headers) + try: + kwargs = { + "prompt": "Write code", + "conversation_id": "test-conversation", + "extra_headers": extra_headers, + } + if streaming: + result = [item async for item in provider.text_chat_stream(**kwargs)] + assert any(item.completion_text == "ok" for item in result) + else: + assert (await provider.text_chat(**kwargs)).completion_text == "ok" + assert len(go_http) == 1 + request = go_http[0] + assert request.url.path == f"/zen/go/v1/{endpoint}" + assert request.headers.get_list("user-agent") == [expected_user_agent] + assert request.headers.get_list("x-opencode-session") == [ + hashlib.sha256(b"test-conversation").hexdigest() + ] + assert "extra_headers" not in json.loads(request.content) + await provider.get_models() + assert go_http[-1].headers.get_list("user-agent") == [ + provider.request_headers["User-Agent"] + ] + assert "x-opencode-session" not in go_http[-1].headers + assert provider.delegate.request_headers == default_headers + assert config == original_config + assert extra_headers == original_extra_headers + finally: + await provider.terminate() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider_class,endpoint", GO_PROTOCOL_CASES) +async def test_go_model_changes_preserve_selected_protocol( + go_http, provider_class, endpoint +): + provider = provider_class({"key": ["test-key"]}, {}) + try: + assert await provider.get_models() == ["kimi-k2.6"] + for model in ["kimi-k2.6", "gpt-5.6-luna", "minimax-m3"]: + provider.set_model(model) + await provider.text_chat(prompt="Write code") + assert [r.url.path for r in go_http[1:]] == [f"/zen/go/v1/{endpoint}"] * 3 + assert [json.loads(r.content)["model"] for r in go_http[1:]] == [ + "kimi-k2.6", + "gpt-5.6-luna", + "minimax-m3", + ] + assert len({r.headers["x-opencode-session"] for r in go_http[1:]}) == 3 + finally: + await provider.terminate() + + +@pytest.mark.asyncio +async def test_go_summary_preserves_conversation_conversation_id(go_http): + provider = ProviderOpenCodeGo({"key": ["test-key"]}, {}) + conversation_id = str(uuid4()) + manager = ContextManager( + ContextConfig(llm_compress_provider=provider, llm_compress_keep_recent_ratio=0), + conversation_id=conversation_id, + ) + try: + await provider.text_chat(prompt="Write code", conversation_id=conversation_id) + await manager.compressor( + [ + Message(role="user", content="Write code"), + Message(role="assistant", content="Here is the code"), + ] + ) + assert len(go_http) == 2 + assert ( + go_http[0].headers["x-opencode-session"] + == go_http[1].headers["x-opencode-session"] + ) + finally: + await provider.terminate() + + +@pytest.mark.parametrize( + "name,provider_type,provider_class", + [ + ("Chat Completions", "opencode_go_chat_completion", ProviderOpenCodeGo), + ("Responses", "opencode_go_responses", ProviderOpenCodeGoResponses), + ("Messages", "opencode_go_messages", ProviderOpenCodeGoMessages), + ], +) +def test_go_templates(name, provider_type, provider_class): + from astrbot.core.provider.register import provider_cls_map + + template = CONFIG_METADATA_2["provider_group"]["metadata"]["provider"][ + "config_template" + ][f"OpenCode Go {name}"] + assert template["type"] == provider_type + assert provider_cls_map[provider_type].cls_type is provider_class + assert template["api_base"] == "https://opencode.ai/zen/go/v1" + + +@pytest.mark.asyncio +async def test_go_conversation_switch_within_same_umo(go_http): + provider = ProviderOpenCodeGo({"key": ["test-key"]}, {}) + first_cid, second_cid = str(uuid4()), str(uuid4()) + try: + for cid in [first_cid, second_cid, first_cid]: + await provider.text_chat( + prompt="Write code", + session_id="qq:GroupMessage:456", + conversation_id=cid, + ) + identities = [request.headers["x-opencode-session"] for request in go_http] + assert identities[0] == identities[2] + assert identities[0] != identities[1] + assert identities[0] == hashlib.sha256(first_cid.encode()).hexdigest() + finally: + await provider.terminate() diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index 180e0edf2d..1e84ff4942 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -5,6 +5,7 @@ from types import SimpleNamespace from typing import Any, cast from unittest.mock import AsyncMock, MagicMock +from uuid import uuid4 import pytest @@ -20,10 +21,13 @@ from astrbot.core.agent.tool import FunctionTool, ToolSet from astrbot.core.astr_agent_run_util import run_agent from astrbot.core.astr_agent_tool_exec import FunctionToolExecutor +from astrbot.core.db.po import Conversation from astrbot.core.exceptions import EmptyModelOutputError from astrbot.core.message.message_event_result import MessageChain +from astrbot.core.platform.astr_message_event import AstrMessageEvent from astrbot.core.provider.entities import LLMResponse, ProviderRequest, TokenUsage from astrbot.core.provider.provider import Provider +from astrbot.core.star.context import Context class MockProvider(Provider): @@ -619,6 +623,130 @@ async def snapshot_context_manager(messages, trusted_token_usage=0): assert "工具执行结果" in tool_messages[0].content +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [False, True]) +@pytest.mark.parametrize("persistent", [False, True]) +@pytest.mark.parametrize("tool_schema_mode", ["full", "skills_like"]) +async def test_conversation_identity_is_stable_within_each_run( + runner, + mock_provider, + tool_set, + mock_tool_executor, + mock_hooks, + streaming, + persistent, + tool_schema_mode, +): + """Keep one identity through tools, requery repair, compression, and reset.""" + conversation = ( + Conversation(platform_id="test", user_id="user", cid=str(uuid4())) + if persistent + else None + ) + identities = [] + for _ in range(2): + tool_response = LLMResponse( + role="assistant", + tools_call_name=["test_tool"], + tools_call_args=[{"query": "test"}], + tools_call_ids=["call_identity"], + ) + responses = [tool_response] + if tool_schema_mode == "skills_like": + responses.extend([LLMResponse(role="assistant"), tool_response]) + responses.extend( + [ + LLMResponse(role="assistant", completion_text="final"), + LLMResponse(role="assistant", completion_text="summary"), + ] + ) + mock_provider.text_chat = AsyncMock(side_effect=responses) + request = ProviderRequest( + prompt="Run the tool", + func_tool=tool_set, + conversation=conversation, + ) + await runner.reset( + provider=mock_provider, + request=request, + run_context=ContextWrapper(context=None), + tool_executor=mock_tool_executor, + agent_hooks=mock_hooks, + streaming=streaming, + tool_schema_mode=tool_schema_mode, + llm_compress_provider=mock_provider, + llm_compress_keep_recent_ratio=0, + ) + async for _ in runner.step_until_done(3): + pass + assert runner.done() + assert any(message.role == "tool" for message in runner.run_context.messages) + await runner.request_context_manager.compressor(runner.run_context.messages) + + calls = mock_provider.text_chat.call_args_list + assert len(calls) == (5 if tool_schema_mode == "skills_like" else 3) + identity = calls[0].kwargs["conversation_id"] + assert identity + assert all(call.kwargs["conversation_id"] == identity for call in calls) + assert request.conversation is conversation + if conversation: + assert identity == conversation.cid + identities.append(identity) + + assert (identities[0] == identities[1]) is persistent + + +@pytest.mark.asyncio +async def test_tool_loop_agent_passes_explicit_conversation_id(): + """Context.tool_loop_agent forwards conversation_id through runner to provider.""" + provider = MockProvider() + context = Context( + event_queue=AsyncMock(), + config=MagicMock(), + db=MagicMock(), + provider_manager=SimpleNamespace( + get_provider_by_id=AsyncMock(return_value=provider) + ), + platform_manager=MagicMock(), + conversation_manager=MagicMock(), + message_history_manager=MagicMock(), + persona_manager=MagicMock(), + astrbot_config_mgr=MagicMock(), + knowledge_base_manager=MagicMock(), + cron_manager=MagicMock(), + ) + event = MagicMock(spec=AstrMessageEvent) + event.unified_msg_origin = "test_umo" + seen: list[str | None] = [] + + async def text_chat(**kwargs): + seen.append(kwargs.get("conversation_id")) + return LLMResponse(role="assistant", completion_text="done") + + provider.text_chat = text_chat + + async def run_once(conversation_id: str | None = None) -> None: + kwargs = {"conversation_id": conversation_id} if conversation_id else {} + resp = await context.tool_loop_agent( + event=event, + chat_provider_id="provider-id", + prompt="hi", + **kwargs, + ) + assert resp.completion_text == "done" + + await run_once("stable-id") + await run_once("stable-id") + assert seen == ["stable-id", "stable-id"] + await run_once("other-id") + assert seen[-1] == "other-id" + transient = len(seen) + await run_once() + await run_once() + assert seen[transient] and seen[transient + 1] + assert seen[transient] != seen[transient + 1] + + @pytest.mark.asyncio async def test_normal_completion_without_max_step( runner, mock_provider, provider_request, mock_tool_executor, mock_hooks