deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/agents/fixed_answer_agent.py
2026-09-07 18:24:55 +08:00

187 lines
6.4 KiB
Python

"""Small local streaming graph used for exact-match fixed answers.
The graph deliberately uses a local ``BaseChatModel`` implementation so
LangGraph emits the same token/message stream events as a normal model-backed
answer, while making no network or model-provider request.
"""
from __future__ import annotations
import asyncio
import codecs
import math
import uuid
from collections.abc import AsyncIterator, Sequence
from typing import Any
import tiktoken
from langchain_core.callbacks import AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
from langchain_core.runnables import RunnableConfig
from langgraph.graph import END, START, StateGraph
from deerflow.agents.thread_state import ThreadState
_TOKEN_ENCODING = tiktoken.get_encoding("cl100k_base")
_TARGET_CHUNKS_PER_SECOND = 10
_MAX_TOKENS_PER_CHUNK = 20
_MAX_TOKENS_PER_SECOND = 1_000
class FixedAnswerChatModel(BaseChatModel):
"""A local streaming model that only replays an administrator answer."""
answer: str
message_id: str
fixed_question_id: str
tokens_per_second: int = 100
@property
def _llm_type(self) -> str:
return "fixed-answer"
@property
def _identifying_params(self) -> dict[str, Any]:
return {
"fixed_question_id": self.fixed_question_id,
"tokens_per_second": self.tokens_per_second,
}
def _message(self) -> AIMessage:
return AIMessage(
content=self.answer,
id=self.message_id,
additional_kwargs={
"fixed_answer": True,
"fixed_question_id": self.fixed_question_id,
"tokens_per_second": self.tokens_per_second,
},
)
def _generate(
self,
messages: list[BaseMessage],
stop: list[str] | None = None,
run_manager: CallbackManagerForLLMRun | None = None,
**kwargs: Any,
) -> ChatResult:
del messages, stop, run_manager, kwargs
return ChatResult(generations=[ChatGeneration(message=self._message())])
async def _astream(
self,
messages: list[BaseMessage],
stop: list[str] | None = None,
run_manager: AsyncCallbackManagerForLLMRun | None = None,
**kwargs: Any,
) -> AsyncIterator[ChatGenerationChunk]:
del messages, stop, run_manager, kwargs
answer = self.answer
if not answer:
yield ChatGenerationChunk(message=AIMessageChunk(content="", id=self.message_id))
return
token_ids = _TOKEN_ENCODING.encode(answer, disallowed_special=())
tokens_per_second = max(
1,
min(_MAX_TOKENS_PER_SECOND, self.tokens_per_second),
)
chunk_size = max(
1,
min(
_MAX_TOKENS_PER_CHUNK,
math.ceil(tokens_per_second / _TARGET_CHUNKS_PER_SECOND),
),
)
decoder = codecs.getincrementaldecoder("utf-8")()
token_groups = [token_ids[:1], *(token_ids[offset : offset + chunk_size] for offset in range(1, len(token_ids), chunk_size))]
for index, chunk_token_ids in enumerate(token_groups):
if index > 0:
await asyncio.sleep(len(chunk_token_ids) / tokens_per_second)
chunk_bytes = b"".join(_TOKEN_ENCODING.decode_single_token_bytes(token_id) for token_id in chunk_token_ids)
is_last = index == len(token_groups) - 1
content = decoder.decode(chunk_bytes, final=is_last)
chunk_kwargs = (
{
"fixed_answer": True,
"fixed_question_id": self.fixed_question_id,
"tokens_per_second": tokens_per_second,
}
if index == 0
else {}
)
yield ChatGenerationChunk(
message=AIMessageChunk(
content=content,
id=self.message_id,
additional_kwargs=chunk_kwargs,
)
)
def _fallback_title(question: str) -> str:
if len(question) > 50:
return question[:50].rstrip() + "..."
return question or "New Conversation"
def make_fixed_answer_agent(
*,
question: str,
answer: str,
fixed_question_id: str,
tokens_per_second: int = 100,
):
"""Build a one-node graph that streams and persists a fixed answer."""
tokens_per_second = max(
1,
min(_MAX_TOKENS_PER_SECOND, int(tokens_per_second)),
)
message_id = f"fixed-answer-{uuid.uuid4()}"
model = FixedAnswerChatModel(
answer=answer,
message_id=message_id,
fixed_question_id=fixed_question_id,
tokens_per_second=tokens_per_second,
).with_config(tags=["fixed-answer"])
async def respond(
state: ThreadState,
config: RunnableConfig,
) -> dict[str, Any]:
messages: Sequence[BaseMessage] = state.get("messages", [])
# Calling the local model through ``astream`` makes the runtime publish
# normal messages-mode chunks. The full AIMessage returned below is the
# durable checkpoint value used on refresh/history reads.
async for _chunk in model.astream(list(messages), config=config):
pass
update: dict[str, Any] = {
"messages": [
AIMessage(
content=answer,
id=message_id,
additional_kwargs={
"fixed_answer": True,
"fixed_question_id": fixed_question_id,
"tokens_per_second": tokens_per_second,
},
)
]
}
if not state.get("title"):
update["title"] = _fallback_title(question)
update["title_provisional"] = False
elif state.get("title_provisional") is True:
# A fixed-answer run must not trigger the worker's background LLM
# title finalizer either; keep the already-visible local title.
update["title_provisional"] = False
return update
graph = StateGraph(ThreadState)
graph.add_node("fixed_answer", respond)
graph.add_edge(START, "fixed_answer")
graph.add_edge("fixed_answer", END)
return graph.compile()