187 lines
6.4 KiB
Python
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()
|