122 lines
3.4 KiB
Python
122 lines
3.4 KiB
Python
from copy import deepcopy
|
|
|
|
import pytest
|
|
|
|
from src.agent.state import AccidentGraphState, GeneratedTurn
|
|
from src.backends.chat import ChatInput, FormUpdate, TextDelta
|
|
from src.backends.fastgpt import FastGPTBackend
|
|
from src.backends.langgraph import LangGraphBackend
|
|
from src.core.config import Settings
|
|
from src.core.fastgpt_client import create_chat_backend
|
|
|
|
|
|
class FakeResponseGenerator:
|
|
def __init__(self):
|
|
self.states = []
|
|
self.closed = False
|
|
|
|
async def generate(self, state: AccidentGraphState) -> GeneratedTurn:
|
|
self.states.append(deepcopy(state))
|
|
turn_number = state.get("turn_count", 0) + 1
|
|
patch = {"turn": turn_number} if state["need_form_update"] else {}
|
|
return GeneratedTurn(
|
|
content=f"<state>1002</state>第{turn_number}轮",
|
|
form_update=patch,
|
|
)
|
|
|
|
async def aclose(self) -> None:
|
|
self.closed = True
|
|
|
|
|
|
def langgraph_settings():
|
|
return Settings(
|
|
_env_file=None,
|
|
environment="test",
|
|
agent_backend="langgraph",
|
|
langgraph_checkpointer="memory",
|
|
llm_api_key="test-key",
|
|
llm_model="test-model",
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_langgraph_backend_preserves_thread_scoped_turn_state():
|
|
generator = FakeResponseGenerator()
|
|
backend = create_chat_backend(
|
|
langgraph_settings(),
|
|
response_generator=generator,
|
|
)
|
|
assert isinstance(backend, LangGraphBackend)
|
|
|
|
first = await backend.complete(
|
|
ChatInput("session-1", "第一轮", need_form_update=True)
|
|
)
|
|
second = await backend.complete(
|
|
ChatInput("session-1", "第二轮", need_form_update=True)
|
|
)
|
|
other_session = await backend.complete(
|
|
ChatInput("session-2", "独立会话", need_form_update=True)
|
|
)
|
|
|
|
assert first.content == "<state>1002</state>第1轮"
|
|
assert first.form_update == {"turn": 1}
|
|
assert second.content == "<state>1002</state>第2轮"
|
|
assert second.form_update == {"turn": 2}
|
|
assert other_session.content == "<state>1002</state>第1轮"
|
|
assert generator.states[0].get("turn_count", 0) == 0
|
|
assert generator.states[1]["turn_count"] == 1
|
|
assert generator.states[2].get("turn_count", 0) == 0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_langgraph_stream_bridges_form_update_before_text():
|
|
generator = FakeResponseGenerator()
|
|
backend = create_chat_backend(
|
|
langgraph_settings(),
|
|
response_generator=generator,
|
|
)
|
|
|
|
events = [
|
|
event
|
|
async for event in backend.stream(
|
|
ChatInput("session-stream", "开始", need_form_update=True)
|
|
)
|
|
]
|
|
|
|
assert events == [
|
|
FormUpdate({"turn": 1}),
|
|
TextDelta("<state>1002</state>第1轮"),
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_langgraph_backend_closes_owned_generator():
|
|
generator = FakeResponseGenerator()
|
|
backend = create_chat_backend(
|
|
langgraph_settings(),
|
|
response_generator=generator,
|
|
)
|
|
|
|
await backend.aclose()
|
|
|
|
assert generator.closed
|
|
|
|
|
|
def test_factory_keeps_fastgpt_as_default_compatible_backend():
|
|
settings = Settings(
|
|
_env_file=None,
|
|
environment="test",
|
|
agent_backend="fastgpt",
|
|
fastgpt_api_key="test-key",
|
|
fastgpt_base_url="http://fastgpt.test",
|
|
fastgpt_app_id="test-app",
|
|
)
|
|
fake_client = object()
|
|
|
|
backend = create_chat_backend(
|
|
settings,
|
|
fastgpt_client=fake_client,
|
|
)
|
|
|
|
assert isinstance(backend, FastGPTBackend)
|