From bb6cd8d96dde01991283c29917abccb73f287772 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=B0=95=E6=B0=99?= Date: Wed, 9 Sep 2026 23:01:22 +0800 Subject: [PATCH 1/6] feat(provider): add OpenCode Go protocol providers --- astrbot/core/agent/context/compressor.py | 5 + astrbot/core/agent/context/manager.py | 3 + .../agent/runners/tool_loop_agent_runner.py | 14 +- astrbot/core/astr_main_agent.py | 13 +- astrbot/core/config/default.py | 36 +++ astrbot/core/provider/manager.py | 8 + .../core/provider/sources/anthropic_source.py | 4 + .../sources/openai_responses_source.py | 4 + .../core/provider/sources/openai_source.py | 3 + .../provider/sources/opencode_go_source.py | 199 ++++++++++++ docs/en/providers/opencode-go.md | 24 ++ docs/zh/providers/opencode-go.md | 24 ++ tests/test_opencode_go_source.py | 305 ++++++++++++++++++ 13 files changed, 638 insertions(+), 4 deletions(-) create mode 100644 astrbot/core/provider/sources/opencode_go_source.py create mode 100644 docs/en/providers/opencode-go.md create mode 100644 docs/zh/providers/opencode-go.md create mode 100644 tests/test_opencode_go_source.py 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 3c4cab9046..4686260f7e 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -257,7 +257,10 @@ 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.req.conversation.cid + if self.req.conversation + else None, ) self.provider = provider @@ -505,6 +508,9 @@ 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.req.conversation.cid + if self.req.conversation + else None, "extra_user_content_parts": self.req.extra_user_content_parts, # list[ContentPart] "abort_signal": self._abort_signal, "request_max_retries": self.request_max_retries, @@ -1427,6 +1433,9 @@ async def _resolve_tool_exec( func_tool=param_subset, model=self.req.model, session_id=self.req.session_id, + conversation_id=self.req.conversation.cid + if self.req.conversation + else None, extra_user_content_parts=self.req.extra_user_content_parts, # tool_choice="required", abort_signal=self._abort_signal, @@ -1459,6 +1468,9 @@ async def _resolve_tool_exec( func_tool=param_subset, model=self.req.model, session_id=self.req.session_id, + conversation_id=self.req.conversation.cid + if self.req.conversation + else None, 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 4a4456b410..82247acd42 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -696,6 +696,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: @@ -715,6 +716,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 @@ -738,6 +740,7 @@ async def _ensure_img_caption( cfg, compressed_urls, plugin_context, + conversation_id=req.conversation.cid if req.conversation else None, ) if caption: req.extra_user_content_parts.append( @@ -947,6 +950,9 @@ async def _process_quote_message( llm_resp = await prov.text_chat( prompt="Please describe the image content.", image_urls=[compress_path], + conversation_id=req.conversation.cid + if req.conversation + else None, ) if llm_resp.completion_text: content_parts.append( @@ -1111,6 +1117,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, " @@ -1458,6 +1465,9 @@ async def build_main_agent( return None req.prompt = event.message_str[len(config.provider_wake_prefix) :] + conversation = await _get_session_conv(event, plugin_context) + req.conversation = conversation + req.contexts = json.loads(conversation.history) # media files attachments for comp in event.message_obj.message: @@ -1575,9 +1585,6 @@ async def build_main_agent( exc_info=True, ) - conversation = await _get_session_conv(event, plugin_context) - req.conversation = conversation - req.contexts = json.loads(conversation.history) event.set_extra("provider_request", req) if isinstance(req.contexts, str): diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 8ce60ed7d0..18f28f7c8a 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 c861ded6ba..ebcdbc6d60 100644 --- a/astrbot/core/provider/sources/anthropic_source.py +++ b/astrbot/core/provider/sources/anthropic_source.py @@ -809,6 +809,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 @@ -881,6 +883,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 f7870b7137..b38d39b096 100644 --- a/astrbot/core/provider/sources/openai_source.py +++ b/astrbot/core/provider/sources/openai_source.py @@ -956,6 +956,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: """准备聊天所需的有效载荷和上下文""" @@ -993,6 +994,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..6fa6702f9f --- /dev/null +++ b/astrbot/core/provider/sources/opencode_go_source.py @@ -0,0 +1,199 @@ +import hashlib +from collections.abc import AsyncGenerator +from uuid import uuid4 + +from astrbot import __version__ +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/") + headers = config.get("custom_headers") or {} + config["custom_headers"] = { + key: value + for key, value in headers.items() + if key.lower() not in {"user-agent", "x-opencode-session"} + } + config["custom_headers"]["User-Agent"] = f"AstrBot/{__version__}" + 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/") + delegate = self.delegate + extra_headers = { + key: value + for key, value in (kwargs.pop("extra_headers", None) or {}).items() + if key.lower() not in {"user-agent", "x-opencode-session"} + } + # 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 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="any" + if isinstance(delegate, ProviderAnthropic) and tool_choice == "required" + else 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/") + delegate = self.delegate + extra_headers = { + key: value + for key, value in (kwargs.pop("extra_headers", None) or {}).items() + if key.lower() not in {"user-agent", "x-opencode-session"} + } + extra_headers["x-opencode-session"] = hashlib.sha256( + (kwargs.pop("conversation_id", None) or uuid4().hex).encode() + ).hexdigest() + async for response in 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="any" + if isinstance(delegate, ProviderAnthropic) and tool_choice == "required" + else 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/docs/en/providers/opencode-go.md b/docs/en/providers/opencode-go.md new file mode 100644 index 0000000000..0d8638e905 --- /dev/null +++ b/docs/en/providers/opencode-go.md @@ -0,0 +1,24 @@ +# 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. + +## 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..d2167dfb42 --- /dev/null +++ b/docs/zh/providers/opencode-go.md @@ -0,0 +1,24 @@ +# 接入 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 | + +保存后,点击提供商卡片,添加需要使用的模型。 + +## 设为默认模型 + +进入 **配置文件 → 提供商设置**,将「默认聊天模型」设置为刚刚添加的 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..6691e9d2b9 --- /dev/null +++ b/tests/test_opencode_go_source.py @@ -0,0 +1,305 @@ +import asyncio +import hashlib +import json +from uuid import uuid4 + +import httpx +import pytest + +from astrbot import __version__ +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.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, +) + + +@pytest.fixture +def go_http(monkeypatch): + """Capture real SDK HTTP requests with deterministic protocol responses.""" + requests = [] + + async def handle(request): + requests.append(request) + await asyncio.sleep(0) + if request.method == "GET": + return httpx.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.Response( + 200, text=content, headers={"Content-Type": "text/event-stream"} + ) + return httpx.Response(200, json=response) + + def client(*args): + return httpx.AsyncClient(transport=httpx.MockTransport(handle)) + + 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", + [ + (ProviderOpenCodeGo, "chat/completions"), + (ProviderOpenCodeGoResponses, "responses"), + (ProviderOpenCodeGoMessages, "messages"), + ], +) +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": "wrong", + "X-OpenCode-Session": "wrong", + "X-Custom": "keep", + }, + }, + {}, + ) + sessions = [ + str(uuid4()), + str(uuid4()), + str(uuid4()), + str(uuid4()), + ] + + async def send(conversation_id): + kwargs = { + "prompt": "Write a Python function", + "conversation_id": conversation_id, + "model": f"opencode-go/{model}", + } + 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 + assert {r.headers["x-opencode-session"] for r in go_http} == { + hashlib.sha256(conversation_id.encode()).hexdigest() for conversation_id in sessions + } + assert ( + go_http[-1].headers["x-opencode-session"] + == hashlib.sha256(sessions[0].encode()).hexdigest() + ) + assert ( + sum( + r.headers["x-opencode-session"] + == go_http[-1].headers["x-opencode-session"] + for r in go_http + ) + == 2 + ) + for request in go_http: + assert request.url.path == f"/zen/go/v1/{endpoint}" + assert request.headers["user-agent"] == f"AstrBot/{__version__}" + assert request.headers["x-custom"] == "keep" + assert json.loads(request.content)["model"] == model + finally: + await provider.terminate() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "provider_class,endpoint", + [ + (ProviderOpenCodeGo, "chat/completions"), + (ProviderOpenCodeGoResponses, "responses"), + (ProviderOpenCodeGoMessages, "messages"), + ], +) +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() From c8974c0f21ce972ea99ef8eb5e64ba47e0955615 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=B0=95=E6=B0=99?= Date: Thu, 10 Sep 2026 01:11:48 +0800 Subject: [PATCH 2/6] fix: stabilize agent sessions and support SDK HTTP clients in tests --- .../agent/runners/tool_loop_agent_runner.py | 20 ++--- tests/test_opencode_go_source.py | 36 +++++++-- tests/test_tool_loop_agent_runner.py | 75 +++++++++++++++++++ 3 files changed, 111 insertions(+), 20 deletions(-) diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index 4686260f7e..2c1fdff32f 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -232,6 +232,10 @@ async def reset( **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 else uuid.uuid4().hex + ) self.streaming = streaming self.enforce_max_turns = enforce_max_turns self.llm_compress_instruction = llm_compress_instruction @@ -258,9 +262,7 @@ async def reset( ) self.request_context_manager = ContextManager( self.request_context_manager_config, - conversation_id=self.req.conversation.cid - if self.req.conversation - else None, + conversation_id=self._conversation_id, ) self.provider = provider @@ -508,9 +510,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.req.conversation.cid - if self.req.conversation - else None, + "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, @@ -1433,9 +1433,7 @@ async def _resolve_tool_exec( func_tool=param_subset, model=self.req.model, session_id=self.req.session_id, - conversation_id=self.req.conversation.cid - if self.req.conversation - else None, + conversation_id=self._conversation_id, extra_user_content_parts=self.req.extra_user_content_parts, # tool_choice="required", abort_signal=self._abort_signal, @@ -1468,9 +1466,7 @@ async def _resolve_tool_exec( func_tool=param_subset, model=self.req.model, session_id=self.req.session_id, - conversation_id=self.req.conversation.cid - if self.req.conversation - else None, + 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/tests/test_opencode_go_source.py b/tests/test_opencode_go_source.py index 6691e9d2b9..8dc21bb556 100644 --- a/tests/test_opencode_go_source.py +++ b/tests/test_opencode_go_source.py @@ -1,10 +1,13 @@ import asyncio 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 import __version__ from astrbot.core.agent.context.config import ContextConfig @@ -25,11 +28,11 @@ def go_http(monkeypatch): """Capture real SDK HTTP requests with deterministic protocol responses.""" requests = [] - async def handle(request): + async def handle(request, *, httpx_module): requests.append(request) await asyncio.sleep(0) if request.method == "GET": - return httpx.Response(200, json={"data": [{"id": "kimi-k2.6"}]}) + 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"): @@ -128,13 +131,23 @@ async def handle(request): f"event: {event.get('type', 'message')}\ndata: {json.dumps(event)}\n\n" for event in events ) - return httpx.Response( + return httpx_module.Response( 200, text=content, headers={"Content-Type": "text/event-stream"} ) - return httpx.Response(200, json=response) + return httpx_module.Response(200, json=response) - def client(*args): - return httpx.AsyncClient(transport=httpx.MockTransport(handle)) + 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) @@ -178,6 +191,7 @@ async def send(conversation_id): "prompt": "Write a Python function", "conversation_id": conversation_id, "model": f"opencode-go/{model}", + "extra_headers": {"X-Request-Test": "request-header"}, } if streaming: result = [item async for item in provider.text_chat_stream(**kwargs)] @@ -190,7 +204,8 @@ async def send(conversation_id): await send(sessions[0]) assert len(go_http) == 5 assert {r.headers["x-opencode-session"] for r in go_http} == { - hashlib.sha256(conversation_id.encode()).hexdigest() for conversation_id in sessions + hashlib.sha256(conversation_id.encode()).hexdigest() + for conversation_id in sessions } assert ( go_http[-1].headers["x-opencode-session"] @@ -208,7 +223,12 @@ async def send(conversation_id): assert request.url.path == f"/zen/go/v1/{endpoint}" assert request.headers["user-agent"] == f"AstrBot/{__version__}" assert request.headers["x-custom"] == "keep" - assert json.loads(request.content)["model"] == model + 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() diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index b21ed40d82..01b69848e9 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 +from uuid import uuid4 import pytest @@ -19,6 +20,7 @@ from astrbot.core.agent.runners.tool_loop_agent_runner import ToolLoopAgentRunner from astrbot.core.agent.tool import FunctionTool, ToolSet 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.provider.entities import LLMResponse, ProviderRequest, TokenUsage from astrbot.core.provider.provider import Provider @@ -617,6 +619,79 @@ 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_normal_completion_without_max_step( runner, mock_provider, provider_request, mock_tool_executor, mock_hooks From 233ccdb5a96f142c397d79d7805d7a6ec44b081c Mon Sep 17 00:00:00 2001 From: RuochenPan Date: Thu, 10 Sep 2026 09:07:17 +0800 Subject: [PATCH 3/6] refactor: simplify OpenCode Go adapter and tests --- .../provider/sources/opencode_go_source.py | 14 ++--- tests/test_opencode_go_source.py | 53 ++++++------------- 2 files changed, 20 insertions(+), 47 deletions(-) diff --git a/astrbot/core/provider/sources/opencode_go_source.py b/astrbot/core/provider/sources/opencode_go_source.py index 6fa6702f9f..b850b57328 100644 --- a/astrbot/core/provider/sources/opencode_go_source.py +++ b/astrbot/core/provider/sources/opencode_go_source.py @@ -85,7 +85,6 @@ async def text_chat( The normalized model response. """ model = (model or self.get_model()).removeprefix("opencode-go/") - delegate = self.delegate extra_headers = { key: value for key, value in (kwargs.pop("extra_headers", None) or {}).items() @@ -95,7 +94,7 @@ async def text_chat( extra_headers["x-opencode-session"] = hashlib.sha256( (kwargs.pop("conversation_id", None) or uuid4().hex).encode() ).hexdigest() - return await delegate.text_chat( + return await self.delegate.text_chat( prompt=prompt, session_id=session_id, image_urls=image_urls, @@ -106,9 +105,7 @@ async def text_chat( tool_calls_result=tool_calls_result, model=model, extra_user_content_parts=extra_user_content_parts, - tool_choice="any" - if isinstance(delegate, ProviderAnthropic) and tool_choice == "required" - else tool_choice, + tool_choice=tool_choice, extra_headers=extra_headers, **kwargs, ) @@ -149,7 +146,6 @@ async def text_chat_stream( Normalized response chunks. """ model = (model or self.get_model()).removeprefix("opencode-go/") - delegate = self.delegate extra_headers = { key: value for key, value in (kwargs.pop("extra_headers", None) or {}).items() @@ -158,7 +154,7 @@ async def text_chat_stream( extra_headers["x-opencode-session"] = hashlib.sha256( (kwargs.pop("conversation_id", None) or uuid4().hex).encode() ).hexdigest() - async for response in delegate.text_chat_stream( + async for response in self.delegate.text_chat_stream( prompt=prompt, session_id=session_id, image_urls=image_urls, @@ -169,9 +165,7 @@ async def text_chat_stream( tool_calls_result=tool_calls_result, model=model, extra_user_content_parts=extra_user_content_parts, - tool_choice="any" - if isinstance(delegate, ProviderAnthropic) and tool_choice == "required" - else tool_choice, + tool_choice=tool_choice, extra_headers=extra_headers, **kwargs, ): diff --git a/tests/test_opencode_go_source.py b/tests/test_opencode_go_source.py index 8dc21bb556..86098652ff 100644 --- a/tests/test_opencode_go_source.py +++ b/tests/test_opencode_go_source.py @@ -22,6 +22,12 @@ ProviderOpenCodeGoResponses, ) +GO_PROTOCOL_CASES = [ + (ProviderOpenCodeGo, "chat/completions"), + (ProviderOpenCodeGoResponses, "responses"), + (ProviderOpenCodeGoMessages, "messages"), +] + @pytest.fixture def go_http(monkeypatch): @@ -156,14 +162,7 @@ def client(provider, _config): @pytest.mark.asyncio @pytest.mark.parametrize("streaming", [False, True]) -@pytest.mark.parametrize( - "provider_class,endpoint", - [ - (ProviderOpenCodeGo, "chat/completions"), - (ProviderOpenCodeGoResponses, "responses"), - (ProviderOpenCodeGoMessages, "messages"), - ], -) +@pytest.mark.parametrize("provider_class,endpoint", GO_PROTOCOL_CASES) async def test_go_http_identity_and_concurrent_sessions( go_http, provider_class, endpoint, streaming ): @@ -179,11 +178,10 @@ async def test_go_http_identity_and_concurrent_sessions( }, {}, ) - sessions = [ - str(uuid4()), - str(uuid4()), - str(uuid4()), - str(uuid4()), + sessions = [str(uuid4()) for _ in range(4)] + expected_session_ids = [ + hashlib.sha256(conversation_id.encode()).hexdigest() + for conversation_id in sessions ] async def send(conversation_id): @@ -203,22 +201,10 @@ async def send(conversation_id): await asyncio.gather(*(send(conversation_id) for conversation_id in sessions)) await send(sessions[0]) assert len(go_http) == 5 - assert {r.headers["x-opencode-session"] for r in go_http} == { - hashlib.sha256(conversation_id.encode()).hexdigest() - for conversation_id in sessions - } - assert ( - go_http[-1].headers["x-opencode-session"] - == hashlib.sha256(sessions[0].encode()).hexdigest() - ) - assert ( - sum( - r.headers["x-opencode-session"] - == go_http[-1].headers["x-opencode-session"] - for r in go_http - ) - == 2 - ) + 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["user-agent"] == f"AstrBot/{__version__}" @@ -234,14 +220,7 @@ async def send(conversation_id): @pytest.mark.asyncio -@pytest.mark.parametrize( - "provider_class,endpoint", - [ - (ProviderOpenCodeGo, "chat/completions"), - (ProviderOpenCodeGoResponses, "responses"), - (ProviderOpenCodeGoMessages, "messages"), - ], -) +@pytest.mark.parametrize("provider_class,endpoint", GO_PROTOCOL_CASES) async def test_go_model_changes_preserve_selected_protocol( go_http, provider_class, endpoint ): From bc7db6392fc443223c314c38374d43833b9e740c Mon Sep 17 00:00:00 2001 From: RuochenPan Date: Thu, 10 Sep 2026 09:08:45 +0800 Subject: [PATCH 4/6] refactor: restore conversation initialization after attachments --- astrbot/core/astr_main_agent.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/astrbot/core/astr_main_agent.py b/astrbot/core/astr_main_agent.py index 82247acd42..2cbf45a873 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -1465,9 +1465,6 @@ async def build_main_agent( return None req.prompt = event.message_str[len(config.provider_wake_prefix) :] - conversation = await _get_session_conv(event, plugin_context) - req.conversation = conversation - req.contexts = json.loads(conversation.history) # media files attachments for comp in event.message_obj.message: @@ -1585,6 +1582,9 @@ async def build_main_agent( exc_info=True, ) + conversation = await _get_session_conv(event, plugin_context) + req.conversation = conversation + req.contexts = json.loads(conversation.history) event.set_extra("provider_request", req) if isinstance(req.contexts, str): From 2cd37cca4d9b7e351fe7941e127c8de101971ed7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=B0=95=E6=B0=99?= Date: Mon, 14 Sep 2026 18:29:24 +0800 Subject: [PATCH 5/6] fix: honor explicit conversation_id in plugin tool loop calls Context.tool_loop_agent forwards extra kwargs to the runner, but reset() never read conversation_id, so plugin-specified identities were silently replaced with random UUIDs and each call produced a new OpenCode session. Accept conversation_id in reset() with three-tier precedence: request.conversation.cid, then the explicit value, then a random UUID. --- .../agent/runners/tool_loop_agent_runner.py | 7 ++- astrbot/core/star/context.py | 1 + tests/test_tool_loop_agent_runner.py | 53 +++++++++++++++++++ 3 files changed, 60 insertions(+), 1 deletion(-) diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index ea57b02562..b5d9c3c73c 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -229,12 +229,17 @@ 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 else uuid.uuid4().hex + 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 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/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index ffd434443a..1e84ff4942 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -24,8 +24,10 @@ 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): @@ -694,6 +696,57 @@ async def test_conversation_identity_is_stable_within_each_run( 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 From 0d365e2624588a85dc33a7bf648091fe185f3c50 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=B0=95=E6=B0=99?= Date: Tue, 15 Sep 2026 10:31:09 +0800 Subject: [PATCH 6/6] fix: reuse shared User-Agent handling for OpenCode Go --- .../provider/sources/opencode_go_source.py | 19 +-- docs/en/providers/opencode-go.md | 6 + docs/zh/providers/opencode-go.md | 6 + tests/test_opencode_go_source.py | 114 +++++++++++++++++- 4 files changed, 132 insertions(+), 13 deletions(-) diff --git a/astrbot/core/provider/sources/opencode_go_source.py b/astrbot/core/provider/sources/opencode_go_source.py index b850b57328..ebeb6b2b56 100644 --- a/astrbot/core/provider/sources/opencode_go_source.py +++ b/astrbot/core/provider/sources/opencode_go_source.py @@ -2,7 +2,6 @@ from collections.abc import AsyncGenerator from uuid import uuid4 -from astrbot import __version__ from astrbot.core.provider.entities import LLMResponse from astrbot.core.provider.provider import Provider @@ -28,13 +27,11 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None: config = dict(provider_config) config["api_base"] = config.get("api_base") or OPENCODE_GO_API_BASE config["model"] = self.get_model().removeprefix("opencode-go/") - headers = config.get("custom_headers") or {} config["custom_headers"] = { key: value - for key, value in headers.items() - if key.lower() not in {"user-agent", "x-opencode-session"} + for key, value in self.request_headers.items() + if key.lower() != "x-opencode-session" } - config["custom_headers"]["User-Agent"] = f"AstrBot/{__version__}" self.delegate = self.ADAPTER(config, provider_settings) def get_current_key(self) -> str: @@ -85,10 +82,12 @@ async def text_chat( 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 = { - key: value + ("User-Agent" if key.lower() == "user-agent" else key): value for key, value in (kwargs.pop("extra_headers", None) or {}).items() - if key.lower() not in {"user-agent", "x-opencode-session"} + 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( @@ -146,10 +145,12 @@ async def text_chat_stream( 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 = { - key: value + ("User-Agent" if key.lower() == "user-agent" else key): value for key, value in (kwargs.pop("extra_headers", None) or {}).items() - if key.lower() not in {"user-agent", "x-opencode-session"} + 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() diff --git a/docs/en/providers/opencode-go.md b/docs/en/providers/opencode-go.md index 0d8638e905..a70b1e4c40 100644 --- a/docs/en/providers/opencode-go.md +++ b/docs/en/providers/opencode-go.md @@ -17,6 +17,12 @@ Open the AstrBot dashboard and go to **Providers → Add Provider**. Select **Op 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. diff --git a/docs/zh/providers/opencode-go.md b/docs/zh/providers/opencode-go.md index d2167dfb42..5aab49d745 100644 --- a/docs/zh/providers/opencode-go.md +++ b/docs/zh/providers/opencode-go.md @@ -17,6 +17,12 @@ 保存后,点击提供商卡片,添加需要使用的模型。 +## 请求头 + +默认 `User-Agent` 为 `astrbot/<版本>`。非空白值的优先级为:单次请求的 `extra_headers` → 提供商的 `custom_headers` → 默认值。User-Agent 键名不区分大小写,空白值回退到下一层。 + +`x-opencode-session` 根据对话 ID 自动生成,不允许手动覆盖。未传入对话 ID 的直接调用使用独立会话标识。 + ## 设为默认模型 进入 **配置文件 → 提供商设置**,将「默认聊天模型」设置为刚刚添加的 OpenCode Go 模型,然后保存配置。 diff --git a/tests/test_opencode_go_source.py b/tests/test_opencode_go_source.py index 86098652ff..3b494f4e57 100644 --- a/tests/test_opencode_go_source.py +++ b/tests/test_opencode_go_source.py @@ -1,4 +1,5 @@ import asyncio +import copy import hashlib import json from functools import partial @@ -9,11 +10,11 @@ from anthropic import _base_client as anthropic_base_client from openai import _base_client as openai_base_client -from astrbot import __version__ 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 ( @@ -171,7 +172,7 @@ async def test_go_http_identity_and_concurrent_sessions( { "key": ["test-key"], "custom_headers": { - "user-agent": "wrong", + "user-agent": "configured/1.0", "X-OpenCode-Session": "wrong", "X-Custom": "keep", }, @@ -183,13 +184,21 @@ async def test_go_http_identity_and_concurrent_sessions( 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"}, + "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)] @@ -207,7 +216,10 @@ async def send(conversation_id): 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["user-agent"] == f"AstrBot/{__version__}" + 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) @@ -219,6 +231,100 @@ async def send(conversation_id): 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(