mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-06 06:15:18 +00:00
feat(v1): stateless tool continuation + sampling-param passthrough
- Stateless tool continuation. OpenAI-compatible clients (opencode, etc.) resend the full messages array — system, user, assistant(tool_calls), tool(results) — but no conversation_id, so the prior "conversation_id required for tool continuation" 400 broke every tool call. When no conversation_id is present, rebuild the agent + pending tool calls + tool results directly from the resent messages (StreamProcessor.build_continuation_from_messages) instead of loading server-side pending_tool_state, and call gen_continuation. - Forward OpenAI sampling params (temperature, max_tokens, max_completion_tokens, top_p, frequency_penalty, presence_penalty, stop, seed) from the request to the LLM gen call; the agent otherwise uses its configured defaults.
This commit is contained in:
1 parent
b1f220ec05
commit
534ffc4e5b
4 files changed
+107
-14
No files matched your search
@@ -35,6 +35,7 @@ class BaseAgent(ABC):
|
||||
json_schema: Optional[Dict] = None,
|
||||
json_schema_strict: bool = True,
|
||||
json_object: bool = False,
|
||||
llm_params: Optional[Dict] = None,
|
||||
limited_token_mode: Optional[bool] = False,
|
||||
token_limit: Optional[int] = settings.DEFAULT_AGENT_LIMITS["token_limit"],
|
||||
limited_request_mode: Optional[bool] = False,
|
||||
@@ -115,6 +116,9 @@ class BaseAgent(ABC):
|
||||
# ``json_object`` mirrors response_format {"type":"json_object"}.
|
||||
self.json_schema_strict = json_schema_strict
|
||||
self.json_object = json_object
|
||||
# OpenAI sampling params forwarded from the request (temperature,
|
||||
# max_tokens, top_p, ...). Empty when the caller sent none.
|
||||
self.llm_params = llm_params or {}
|
||||
self.limited_token_mode = limited_token_mode
|
||||
self.token_limit = token_limit
|
||||
self.limited_request_mode = limited_request_mode
|
||||
@@ -602,6 +606,10 @@ class BaseAgent(ABC):
|
||||
previous_response_id = self._previous_response_id()
|
||||
if previous_response_id:
|
||||
gen_kwargs["previous_response_id"] = previous_response_id
|
||||
|
||||
# Forward OpenAI sampling params (temperature, max_tokens, top_p, ...).
|
||||
if self.llm_params:
|
||||
gen_kwargs.update(self.llm_params)
|
||||
resp = self.llm.gen_stream(**gen_kwargs)
|
||||
|
||||
if log_context:
|
||||
|
||||
@@ -172,6 +172,63 @@ class StreamProcessor:
|
||||
tools_data=tools_data,
|
||||
)
|
||||
|
||||
def build_continuation_from_messages(self, messages, tool_actions):
|
||||
"""Rebuild a tool continuation from the request messages (STATELESS).
|
||||
|
||||
OpenAI-compatible clients (opencode, etc.) resend the full conversation
|
||||
-- system, user, assistant(tool_calls), tool(results) -- but carry no
|
||||
conversation_id, so there is no server-side ``pending_tool_state`` to
|
||||
load. Reconstruct the agent + continuation context directly from the
|
||||
resent messages and return the same tuple as ``resume_from_tool_actions``:
|
||||
(agent, messages, tools_dict, pending_tool_calls, tool_actions,
|
||||
reasoning_content).
|
||||
"""
|
||||
# Locate the last assistant message that issued tool calls.
|
||||
pending_idx = None
|
||||
for i in range(len(messages) - 1, -1, -1):
|
||||
m = messages[i]
|
||||
if m.get("role") == "assistant" and m.get("tool_calls"):
|
||||
pending_idx = i
|
||||
break
|
||||
if pending_idx is None:
|
||||
raise ValueError(
|
||||
"No assistant message with tool_calls found for continuation"
|
||||
)
|
||||
|
||||
pending_tool_calls = []
|
||||
for tc in messages[pending_idx].get("tool_calls") or []:
|
||||
fn = tc.get("function") or {}
|
||||
raw_args = fn.get("arguments")
|
||||
try:
|
||||
args = (
|
||||
json.loads(raw_args)
|
||||
if isinstance(raw_args, str)
|
||||
else (raw_args or {})
|
||||
)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
args = {}
|
||||
name = fn.get("name", "")
|
||||
pending_tool_calls.append(
|
||||
{
|
||||
"call_id": tc.get("id", ""),
|
||||
"name": name,
|
||||
"tool_name": name,
|
||||
"action_name": name,
|
||||
"llm_name": name,
|
||||
"arguments": args,
|
||||
}
|
||||
)
|
||||
|
||||
# The conversation up to (but not including) the assistant tool_calls;
|
||||
# gen_continuation re-appends the assistant message + tool results.
|
||||
prior_messages = [dict(m) for m in messages[:pending_idx]]
|
||||
|
||||
# Build a normal agent (config / LLM / client tools), no new question.
|
||||
agent = self.build_agent("")
|
||||
tools_dict = agent.tool_executor.get_tools()
|
||||
|
||||
return agent, prior_messages, tools_dict, pending_tool_calls, tool_actions, ""
|
||||
|
||||
def _load_conversation_history(self):
|
||||
"""Load conversation history either from DB or request"""
|
||||
if self.conversation_id and self.initial_user_id:
|
||||
@@ -1268,6 +1325,7 @@ class StreamProcessor:
|
||||
"json_schema": self.agent_config.get("json_schema"),
|
||||
"json_schema_strict": self.agent_config.get("json_schema_strict", True),
|
||||
"json_object": self.agent_config.get("json_object", False),
|
||||
"llm_params": self.data.get("llm_params") or {},
|
||||
"compressed_summary": self.compressed_summary,
|
||||
"llm": llm,
|
||||
"llm_handler": llm_handler,
|
||||
|
||||
@@ -106,21 +106,32 @@ def chat_completions():
|
||||
if internal_data.get("tool_actions"):
|
||||
# Continuation mode
|
||||
conversation_id = internal_data.get("conversation_id")
|
||||
if not conversation_id:
|
||||
return make_response(
|
||||
jsonify({"error": {"message": "conversation_id required for tool continuation", "type": "invalid_request"}}),
|
||||
400,
|
||||
if conversation_id:
|
||||
(
|
||||
agent,
|
||||
messages,
|
||||
tools_dict,
|
||||
pending_tool_calls,
|
||||
tool_actions,
|
||||
reasoning_content,
|
||||
) = processor.resume_from_tool_actions(
|
||||
internal_data["tool_actions"], conversation_id
|
||||
)
|
||||
else:
|
||||
# Stateless continuation: OpenAI-compatible clients (opencode,
|
||||
# etc.) resend the full messages array but no conversation_id,
|
||||
# so rebuild the agent + pending calls from the request itself.
|
||||
(
|
||||
agent,
|
||||
messages,
|
||||
tools_dict,
|
||||
pending_tool_calls,
|
||||
tool_actions,
|
||||
reasoning_content,
|
||||
) = processor.build_continuation_from_messages(
|
||||
internal_data.get("messages", []),
|
||||
internal_data["tool_actions"],
|
||||
)
|
||||
(
|
||||
agent,
|
||||
messages,
|
||||
tools_dict,
|
||||
pending_tool_calls,
|
||||
tool_actions,
|
||||
reasoning_content,
|
||||
) = processor.resume_from_tool_actions(
|
||||
internal_data["tool_actions"], conversation_id
|
||||
)
|
||||
continuation = {
|
||||
"messages": messages,
|
||||
"tools_dict": tools_dict,
|
||||
|
||||
@@ -212,6 +212,16 @@ def translate_request(
|
||||
json_schema_strict = bool((_rf.get("json_schema") or {}).get("strict", True))
|
||||
json_object_mode = _rf.get("type") == "json_object"
|
||||
|
||||
# OpenAI sampling params, forwarded to the LLM gen call (the agent otherwise
|
||||
# uses its configured defaults).
|
||||
sampling_params = {}
|
||||
for _k in (
|
||||
"temperature", "max_tokens", "max_completion_tokens",
|
||||
"top_p", "frequency_penalty", "presence_penalty", "stop", "seed",
|
||||
):
|
||||
if data.get(_k) is not None:
|
||||
sampling_params[_k] = data[_k]
|
||||
|
||||
# Check for continuation (tool results after assistant tool_calls)
|
||||
if is_continuation(messages):
|
||||
tool_actions = extract_tool_results(messages)
|
||||
@@ -222,6 +232,10 @@ def translate_request(
|
||||
"conversation_id": conversation_id,
|
||||
"tool_actions": tool_actions,
|
||||
"api_key": api_key,
|
||||
# Full messages array for STATELESS continuation: OpenAI clients
|
||||
# (opencode, etc.) don't carry a conversation_id, so the agent is
|
||||
# rebuilt from the resent messages instead of server-side state.
|
||||
"messages": messages,
|
||||
}
|
||||
# A continuation only exists if turn 1 was saved, so default to True —
|
||||
# otherwise the final turn and its WAL row are never persisted. An
|
||||
@@ -276,6 +290,8 @@ def translate_request(
|
||||
result["json_schema_strict"] = json_schema_strict
|
||||
if json_object_mode:
|
||||
result["json_object"] = True
|
||||
if sampling_params:
|
||||
result["llm_params"] = sampling_params
|
||||
|
||||
return result
|
||||
|
||||
|
||||
Reference in new issue
Block a user