diff --git a/backend/package/yuxi/agents/middlewares/steer.py b/backend/package/yuxi/agents/middlewares/steer.py index 003cebefb..81a1d1c83 100644 --- a/backend/package/yuxi/agents/middlewares/steer.py +++ b/backend/package/yuxi/agents/middlewares/steer.py @@ -1,14 +1,24 @@ -"""主会话 Steer Middleware。""" +"""主会话 Steer / Guided Middleware。""" from langchain.agents.middleware import AgentMiddleware, hook_config class SteerMiddleware(AgentMiddleware): - """在安全生命周期边界结束当前 Run,让队列优先执行 Steer。""" + """在模型调用前的安全边界处理队列干预。 + + - guided:把等待注入的补充消息追加进当前 Run 的 messages,模型下一轮即看到, + 当前 Run 不终止(Claude Code 式中途修正)。 + - steer:结束当前 Run,让位给高优先级 Steer 请求。 + """ @hook_config(can_jump_to=["end"]) async def abefore_model(self, state, runtime): # noqa: ARG002 - return await self._jump_if_steer_requested(runtime) + guided_update = await self._collect_guided_update(runtime) + steer_jump = await self._jump_if_steer_requested(runtime) + if steer_jump is not None: + # 让位时 guided 消息已标记注入并随 checkpoint 留给下一个 Run。 + return steer_jump + return guided_update @hook_config(can_jump_to=["end"]) async def aafter_model(self, state, runtime): @@ -17,6 +27,20 @@ async def aafter_model(self, state, runtime): return None return await self._jump_if_steer_requested(runtime) + async def _collect_guided_update(self, runtime): + from yuxi.services.agent_request_queue_service import take_pending_guided_messages + + run_id = getattr(runtime.context, "run_id", None) + if not run_id: + return None + try: + messages = await take_pending_guided_messages(run_id) + except Exception: # noqa: BLE001 + return None + if not messages: + return None + return {"messages": messages} + async def _jump_if_steer_requested(self, runtime): from yuxi.services.agent_request_queue_service import should_end_run_for_steer diff --git a/backend/package/yuxi/repositories/agent_run_request_repository.py b/backend/package/yuxi/repositories/agent_run_request_repository.py index edbd80d9d..76b0b5e85 100644 --- a/backend/package/yuxi/repositories/agent_run_request_repository.py +++ b/backend/package/yuxi/repositories/agent_run_request_repository.py @@ -11,7 +11,7 @@ from __future__ import annotations -from sqlalchemy import and_, func, or_, select +from sqlalchemy import and_, func, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession from yuxi.storage.postgres.models_business import AgentRunRequest @@ -109,6 +109,42 @@ async def get_pending_steer( ) return result.scalar_one_or_none() + async def list_pending_guided( + self, + *, + uid: str, + agent_slug: str, + conversation_thread_id: str, + ) -> list[AgentRunRequest]: + """读取线程内待注入的 guided 请求(FIFO,允许多条同时等待)。""" + result = await self.db.execute( + select(AgentRunRequest) + .where( + AgentRunRequest.uid == str(uid), + AgentRunRequest.agent_slug == agent_slug, + AgentRunRequest.conversation_thread_id == conversation_thread_id, + AgentRunRequest.queue_policy == "guided", + AgentRunRequest.status == "queued", + ) + .order_by(AgentRunRequest.created_at.asc(), AgentRunRequest.id.asc()) + ) + return list(result.scalars().all()) + + async def mark_guided_injected(self, request_ids: list[str]) -> int: + """把已注入当前 Run 的 guided 请求收敛为 injected 终态;返回实际更新数。""" + if not request_ids: + return 0 + result = await self.db.execute( + update(AgentRunRequest) + .where( + AgentRunRequest.request_id.in_(request_ids), + AgentRunRequest.queue_policy == "guided", + AgentRunRequest.status == "queued", + ) + .values(status="injected", updated_at=utc_now_naive()) + ) + return int(result.rowcount or 0) + async def get_queue_head( self, *, diff --git a/backend/package/yuxi/services/agent_request_queue_service.py b/backend/package/yuxi/services/agent_request_queue_service.py index 73dc8506d..91e6c86e8 100644 --- a/backend/package/yuxi/services/agent_request_queue_service.py +++ b/backend/package/yuxi/services/agent_request_queue_service.py @@ -26,6 +26,7 @@ resolve_agent_run_config, ) from yuxi.services.input_message_service import AgentRunInputMessage +from yuxi.services.run_queue_service import append_run_stream_event from yuxi.services.workdir_service import ( WorkdirBinding, resolve_conversation_workdir_binding, @@ -43,8 +44,8 @@ ) from yuxi.workspace.paths import ensure_bound_user_workdir -SUPPORTED_QUEUE_POLICIES = ("enqueue", "reject", "steer") -NOT_IMPLEMENTED_QUEUE_POLICIES = ("guided", "bridge") +SUPPORTED_QUEUE_POLICIES = ("enqueue", "reject", "steer", "guided") +NOT_IMPLEMENTED_QUEUE_POLICIES = ("bridge",) # Request lifecycle states. REQUEST_STATUS_QUEUED = "queued" @@ -52,7 +53,11 @@ REQUEST_STATUS_CANCELLED = "cancelled" REQUEST_STATUS_REJECTED = "rejected" REQUEST_STATUS_FAILED = "failed" -REQUEST_TERMINAL_STATUSES = frozenset({REQUEST_STATUS_CANCELLED, REQUEST_STATUS_REJECTED, REQUEST_STATUS_FAILED}) +# guided 请求已注入当前活跃 Run;不再参与队列派发。 +REQUEST_STATUS_INJECTED = "injected" +REQUEST_TERMINAL_STATUSES = frozenset( + {REQUEST_STATUS_CANCELLED, REQUEST_STATUS_REJECTED, REQUEST_STATUS_FAILED, REQUEST_STATUS_INJECTED} +) # Message delivery states aligned with messages.delivery_status. DELIVERY_STATUS_QUEUED = "queued" @@ -122,8 +127,8 @@ async def intake_request( 返回 IntakeResult:dispatched 时含 run_id(调用方需 commit 后 enqueue ARQ)。 """ policy = validate_queue_policy(queue_policy) - if policy == "steer" and source not in {"chat", "channel"}: - raise HTTPException(status_code=422, detail="queue_policy 'steer' 仅支持主会话 Chat/Channel") + if policy in {"steer", "guided"} and source not in {"chat", "channel"}: + raise HTTPException(status_code=422, detail=f"queue_policy '{policy}' 仅支持主会话 Chat/Channel") meta = meta or {} uid_str = str(uid) repo = AgentRunRequestRepository(db) @@ -191,8 +196,13 @@ async def existing_intake_result(binding: WorkdirBinding | None = None) -> Intak ) if latest_run is not None and latest_run.status == "interrupted": raise _queue_conflict("run_interrupted", "线程正在等待用户回答或审批") - if policy == "steer" and active_run is not None and not await _is_steerable_message_run(db=db, run=active_run): + if ( + policy in {"steer", "guided"} + and active_run is not None + and not await _is_steerable_message_run(db=db, run=active_run) + ): raise _queue_conflict("run_not_steerable", "当前运行不支持引导") + # guided 与 steer 的差别:guided 允许多条同时等待(逐条注入);steer 只允许一条。 if policy == "steer" and await repo.get_pending_steer( uid=uid_str, agent_slug=agent_slug, @@ -407,6 +417,127 @@ async def should_end_run_for_steer(run_id: str) -> bool: return request is not None +async def take_pending_guided_messages(run_id: str) -> list: + """收割并标记属于当前 Run 的 guided 消息,返回可注入 state 的 HumanMessage 列表。 + + 仅在模型调用前(middleware before_model)调用;run 由 worker lease 保证单写者, + 读取-标记在同一事务内完成。消息从 Message.extra_metadata.raw_message 恢复, + 保留原始 LangChain id,使 run 结束保存时与既有 Message 行按 id 去重。 + """ + from langchain.messages import HumanMessage + + injected: list = [] + injected_summaries: list[dict] = [] + async with pg_manager.get_async_session_context() as db: + run = await AgentRunRepository(db).get_run(run_id) + if run is None or not await _is_steerable_message_run(db=db, run=run): + return [] + repo = AgentRunRequestRepository(db) + requests = await repo.list_pending_guided( + uid=run.uid, + agent_slug=run.agent_slug, + conversation_thread_id=run.conversation_thread_id, + ) + for request in requests: + message = ( + await db.execute(select(Message).where(Message.id == request.input_message_id)) + ).scalar_one_or_none() + if message is None: + logger.warning(f"guided 请求 {request.request_id} 缺少输入消息,跳过") + continue + raw = ( + (message.extra_metadata or {}).get("raw_message") if isinstance(message.extra_metadata, dict) else None + ) + if isinstance(raw, dict): + try: + injected.append(HumanMessage(**{k: v for k, v in raw.items() if k in ("content", "id", "name")})) + except Exception: + injected.append(HumanMessage(content=message.content)) + else: + injected.append(HumanMessage(content=message.content)) + injected_summaries.append( + { + "request_id": request.request_id, + "content": message.content, + "thread_id": request.conversation_thread_id, + } + ) + if injected_summaries: + await repo.mark_guided_injected([item["request_id"] for item in injected_summaries]) + await db.commit() + + # 注入事件 best-effort:前端据此把排队消息转入对话流。 + for item in injected_summaries: + try: + await append_run_stream_event( + run_id, + "guided_injected", + {"chunk": {"status": "guided_injected", **item}}, + thread_id=item["thread_id"], + ) + except Exception: + logger.warning(f"Failed to publish guided_injected event for run {run_id}", exc_info=True) + return injected + + +async def guided_queued_request( + *, + request_id: str, + current_uid: str, + db: AsyncSession, +) -> IntakeResult: + """把普通 Chat 排队请求升级为 guided(注入当前 Run 而非等待队列)。""" + repo = AgentRunRequestRepository(db) + existing = await repo.get_by_request_id(request_id) + if existing is None or existing.uid != str(current_uid): + raise HTTPException(status_code=404, detail={"code": "request_not_found", "message": "请求不存在"}) + + await _get_thread_conversation( + db=db, + uid=existing.uid, + agent_slug=existing.agent_slug, + thread_id=existing.conversation_thread_id, + lock=True, + ) + request = await repo.lock_by_request_id(request_id) + if request is None or request.uid != str(current_uid): + raise HTTPException(status_code=404, detail={"code": "request_not_found", "message": "请求不存在"}) + if request.queue_policy == "guided" and request.status == REQUEST_STATUS_QUEUED: + return await _build_existing_intake_result( + repo=repo, + request=request, + uid=request.uid, + agent_slug=request.agent_slug, + thread_id=request.conversation_thread_id, + source=request.source, + channel=request.channel, + external_id=request.external_id, + queue_policy="guided", + ) + if request.status != REQUEST_STATUS_QUEUED or request.queue_policy != "enqueue" or request.source != "chat": + raise _queue_conflict("request_not_queued", "只有普通 Chat 排队请求可以升级为注入") + + active_run = await AgentRunRepository(db).get_active_run_by_thread_for_user( + uid=request.uid, + agent_slug=request.agent_slug, + conversation_thread_id=request.conversation_thread_id, + ) + if active_run is None or not await _is_steerable_message_run(db=db, run=active_run): + raise _queue_conflict("run_not_steerable", "当前没有可注入的运行") + + request.queue_policy = "guided" + request.updated_at = utc_now_naive() + await db.flush() + return IntakeResult( + request_id=request.request_id, + status=request.status, + queue_policy=request.queue_policy, + message_id=request.input_message_id, + thread_id=request.conversation_thread_id, + queue_position=await repo.get_queue_position(request_id), + ) + + async def finalize_intake( *, db: AsyncSession, diff --git a/backend/package/yuxi/services/agent_run_service.py b/backend/package/yuxi/services/agent_run_service.py index 1b0105f6e..ffc4609bd 100644 --- a/backend/package/yuxi/services/agent_run_service.py +++ b/backend/package/yuxi/services/agent_run_service.py @@ -19,7 +19,6 @@ import uuid from collections.abc import AsyncIterator from dataclasses import dataclass -from random import uniform from time import monotonic from typing import Any, Literal @@ -41,11 +40,11 @@ ) from yuxi.services.langfuse_service import get_trace_url_by_id_async from yuxi.services.run_queue_service import ( + blocking_read_run_stream_events, build_run_event_envelope, get_arq_pool, get_last_run_stream_seq, list_recent_run_stream_events, - list_run_stream_events, normalize_after_seq, publish_cancel_signals, ) @@ -57,6 +56,7 @@ from yuxi.utils.sse_utils import ( SSE_HEARTBEAT_SECONDS, SSE_MAX_CONNECTION_MINUTES, + SSE_STREAM_BLOCK_MS, format_heartbeat, format_sse, ) @@ -64,12 +64,7 @@ RUN_PROGRESS_RECENT_EVENT_SCAN_LIMIT = 100 RUN_PROGRESS_MESSAGE_LIMIT = 3 RUN_PROGRESS_CONTENT_MAX_CHARS = 800 -RUN_SSE_ACTIVE_POLL_SECONDS = 0.1 -RUN_SSE_SHORT_IDLE_MAX_POLL_SECONDS = 1.0 -RUN_SSE_LONG_IDLE_AFTER_SECONDS = 120.0 -RUN_SSE_LONG_IDLE_MAX_POLL_SECONDS = 4.0 RUN_SSE_STATUS_POLL_SECONDS = 5.0 -RUN_SSE_POLL_JITTER_RATIO = 0.2 def _resolve_agent_run_request_id( @@ -949,22 +944,6 @@ async def _load_stream_run(run_id: str): return await AgentRunRepository(db).get_run(run_id) -def _next_run_sse_poll_interval(current_interval: float, idle_seconds: float) -> float: - """按空闲时长扩大 Run 事件轮询间隔。""" - max_interval = ( - RUN_SSE_LONG_IDLE_MAX_POLL_SECONDS - if idle_seconds >= RUN_SSE_LONG_IDLE_AFTER_SECONDS - else RUN_SSE_SHORT_IDLE_MAX_POLL_SECONDS - ) - return min(max(current_interval * 2, RUN_SSE_ACTIVE_POLL_SECONDS), max_interval) - - -def _jitter_run_sse_poll_interval(interval: float) -> float: - """为轮询间隔增加有限抖动,分散并发连接尖峰。""" - multiplier = uniform(1 - RUN_SSE_POLL_JITTER_RATIO, 1 + RUN_SSE_POLL_JITTER_RATIO) - return interval * multiplier - - async def stream_agent_run_events( *, run_id: str, @@ -977,9 +956,7 @@ async def stream_agent_run_events( last_heartbeat_ts = started_at last_seq = normalize_after_seq(after_seq) started_monotonic = monotonic() - last_event_at = started_monotonic next_status_check_at = started_monotonic + RUN_SSE_STATUS_POLL_SECONDS - poll_interval = RUN_SSE_ACTIVE_POLL_SECONDS try: try: @@ -1003,7 +980,12 @@ async def stream_agent_run_events( while True: try: - events = await list_run_stream_events(run_id, after_seq=last_seq, limit=200) + events = await blocking_read_run_stream_events( + run_id, + after_seq=last_seq, + block_ms=SSE_STREAM_BLOCK_MS, + limit=200, + ) except Exception as e: logger.warning(f"Run SSE redis error for run {run_id}: {e}") yield format_sse( @@ -1016,10 +998,6 @@ async def stream_agent_run_events( ) return - if events: - last_event_at = monotonic() - poll_interval = RUN_SSE_ACTIVE_POLL_SECONDS - emitted_terminal = False for event in events: seq = str(event.get("seq") or "0-0") @@ -1037,8 +1015,9 @@ async def stream_agent_run_events( if emitted_terminal: return + # 阻塞读超时(无新事件)且到达状态检查节流点时才查库,确认 run 状态并检测终止态。 now_monotonic = monotonic() - if now_monotonic >= next_status_check_at: + if not events and now_monotonic >= next_status_check_at: try: run = await _load_stream_run(run_id) if not run: @@ -1059,31 +1038,27 @@ async def stream_agent_run_events( return next_status_check_at = monotonic() + RUN_SSE_STATUS_POLL_SECONDS - if ( - run.status in TERMINAL_RUN_STATUSES - and not bool(getattr(run, "runtime_cleanup_pending", False)) - and not events - ): - terminal_seq = last_seq - if terminal_seq in {"", "0-0"}: - terminal_seq = await get_last_run_stream_seq(run_id) - if terminal_seq in {"", "0-0"}: - terminal_seq = None - terminal_envelope = build_run_event_envelope( - run_id=run_id, - thread_id=run.conversation_thread_id, - event_type="end", - payload={"status": run.status, "request_id": run.request_id}, - created_at=utc_now_naive().isoformat(), - ) - if not verbose: - terminal_envelope = _compact_run_event_envelope(terminal_envelope) - yield format_sse( - terminal_envelope, - event="end", - event_id=terminal_seq, - ) - return + if run.status in TERMINAL_RUN_STATUSES and not bool(getattr(run, "runtime_cleanup_pending", False)): + terminal_seq = last_seq + if terminal_seq in {"", "0-0"}: + terminal_seq = await get_last_run_stream_seq(run_id) + if terminal_seq in {"", "0-0"}: + terminal_seq = None + terminal_envelope = build_run_event_envelope( + run_id=run_id, + thread_id=run.conversation_thread_id, + event_type="end", + payload={"status": run.status, "request_id": run.request_id}, + created_at=utc_now_naive().isoformat(), + ) + if not verbose: + terminal_envelope = _compact_run_event_envelope(terminal_envelope) + yield format_sse( + terminal_envelope, + event="end", + event_id=terminal_seq, + ) + return now = utc_now_naive() elapsed_seconds = (now - started_at).total_seconds() @@ -1094,13 +1069,6 @@ async def stream_agent_run_events( if elapsed_seconds >= SSE_MAX_CONNECTION_MINUTES * 60: return - - status_check_delay = max(0.0, next_status_check_at - monotonic()) - sleep_seconds = min(_jitter_run_sse_poll_interval(poll_interval), status_check_delay) - await asyncio.sleep(sleep_seconds) - if not events: - idle_seconds = monotonic() - last_event_at - poll_interval = _next_run_sse_poll_interval(poll_interval, idle_seconds) except asyncio.CancelledError: return diff --git a/backend/package/yuxi/services/run_queue_service.py b/backend/package/yuxi/services/run_queue_service.py index 1bd7a034a..4a5474a3d 100644 --- a/backend/package/yuxi/services/run_queue_service.py +++ b/backend/package/yuxi/services/run_queue_service.py @@ -224,6 +224,31 @@ async def list_run_stream_events( return events +async def blocking_read_run_stream_events( + run_id: str, + *, + after_seq: str = "0-0", + block_ms: int = 1000, + limit: int = 200, +) -> list[dict]: + """阻塞读取 run 事件流:有事件立即返回,无事件挂起最多 block_ms。 + + 用于 SSE 实时流:替代「非阻塞读 + sleep 轮询」,把流式延迟从轮询间隔 + 降到毫秒级,同时保留 Stream 的游标续读能力(after_seq 语义同 xrange)。 + """ + redis = await get_redis_client() + key = _event_stream_key(run_id) + start = "0-0" if after_seq in {"0-0", ""} else after_seq + result = await redis.xread(streams={key: start}, block=block_ms, count=limit) + if not result: + return [] + events: list[dict] = [] + for _stream_key, rows in result: + for event_id, fields in rows: + events.append(_decode_run_stream_row(run_id, str(event_id), fields)) + return events + + async def list_recent_run_stream_events(run_id: str, *, limit: int = 100) -> list[dict]: """从 Redis Stream 反向读取最近的 run events,返回顺序为新到旧。""" redis = await get_redis_client() diff --git a/backend/package/yuxi/utils/sse_utils.py b/backend/package/yuxi/utils/sse_utils.py index 0c1aa7f62..5c80c3726 100644 --- a/backend/package/yuxi/utils/sse_utils.py +++ b/backend/package/yuxi/utils/sse_utils.py @@ -9,7 +9,10 @@ # Compose limits development-server graceful shutdown separately, so a live # stream cannot block hot reload for this full connection lifetime. SSE_MAX_CONNECTION_MINUTES = int(os.getenv("RUN_SSE_MAX_CONNECTION_MINUTES", "30")) -SSE_POLL_INTERVAL_SECONDS = float(os.getenv("RUN_SSE_POLL_INTERVAL_SECONDS", "1.0")) +SSE_POLL_INTERVAL_SECONDS = float(os.getenv("RUN_SSE_POLL_INTERVAL_SECONDS", "0.5")) +# SSE 阻塞读取 Redis Stream 的单次挂起上限(毫秒)。有事件立即返回,无事件挂起 +# 最多这么久再检查 heartbeat/终止态,把流式延迟从轮询间隔降到毫秒级。 +SSE_STREAM_BLOCK_MS = int(os.getenv("RUN_SSE_STREAM_BLOCK_MS", "1000")) def format_sse(data: dict, event: str, event_id: str | None = None) -> str: diff --git a/backend/test/unit/agents/test_steer_middleware.py b/backend/test/unit/agents/test_steer_middleware.py new file mode 100644 index 000000000..e6e391596 --- /dev/null +++ b/backend/test/unit/agents/test_steer_middleware.py @@ -0,0 +1,218 @@ +"""Steer/Guided Middleware 单元测试。""" + +from __future__ import annotations + +from types import SimpleNamespace + +from langchain_core.messages import AIMessage +from yuxi.services import agent_request_queue_service + +import pytest + +import yuxi.agents.middlewares.steer as steer_module +from yuxi.agents.middlewares.steer import SteerMiddleware + +pytestmark = [pytest.mark.unit] + + +def _runtime(run_id="run-1"): + return SimpleNamespace(context=SimpleNamespace(run_id=run_id)) + + +@pytest.mark.asyncio +async def test_before_model_injects_pending_guided_messages(monkeypatch): + from langchain.messages import HumanMessage + + middleware = SteerMiddleware() + messages = [HumanMessage(content="补充:只看 2024 年后的数据", id="m-1")] + + async def fake_take(run_id): + assert run_id == "run-1" + return messages + + async def fake_steer(run_id): + return False + + monkeypatch.setattr(steer_module, "take_pending_guided_messages", fake_take, raising=False) + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.take_pending_guided_messages", fake_take + ) + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.should_end_run_for_steer", fake_steer + ) + + result = await middleware.abefore_model({}, _runtime()) + + assert result == {"messages": messages} + + +@pytest.mark.asyncio +async def test_before_model_returns_none_without_interventions(monkeypatch): + middleware = SteerMiddleware() + + async def fake_take(run_id): + return [] + + async def fake_steer(run_id): + return False + + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.take_pending_guided_messages", fake_take + ) + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.should_end_run_for_steer", fake_steer + ) + + assert await middleware.abefore_model({}, _runtime()) is None + + +@pytest.mark.asyncio +async def test_before_model_prefers_steer_jump_when_both_pending(monkeypatch): + from langchain.messages import HumanMessage + + middleware = SteerMiddleware() + + async def fake_take(run_id): + return [HumanMessage(content="guided")] + + async def fake_steer(run_id): + return True + + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.take_pending_guided_messages", fake_take + ) + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.should_end_run_for_steer", fake_steer + ) + + result = await middleware.abefore_model({}, _runtime()) + + assert result == {"jump_to": "end"} + + +@pytest.mark.asyncio +async def test_before_model_guided_failure_does_not_break_run(monkeypatch): + middleware = SteerMiddleware() + + async def fake_take(run_id): + raise RuntimeError("db down") + + async def fake_steer(run_id): + return False + + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.take_pending_guided_messages", fake_take + ) + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.should_end_run_for_steer", fake_steer + ) + + assert await middleware.abefore_model({}, _runtime()) is None + + +@pytest.mark.asyncio +async def test_before_model_without_run_id_skips_queries(monkeypatch): + middleware = SteerMiddleware() + + async def fail_take(run_id): + raise AssertionError("无 run_id 不应查询 guided") + + async def fail_steer(run_id): + raise AssertionError("无 run_id 不应查询 steer") + + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.take_pending_guided_messages", fail_take + ) + monkeypatch.setattr( + "yuxi.services.agent_request_queue_service.should_end_run_for_steer", fail_steer + ) + + assert await middleware.abefore_model({}, SimpleNamespace(context=SimpleNamespace())) is None + + +@pytest.mark.asyncio +async def test_before_model_ends_run_when_steer_is_waiting(monkeypatch: pytest.MonkeyPatch): + """存在待处理 Steer 时,在下一次模型调用前结束当前 Graph。""" + + async def should_end(run_id: str) -> bool: + return run_id == "run-1" + + monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) + runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) + + result = await SteerMiddleware().abefore_model({}, runtime) + + assert result == {"jump_to": "end"} + +@pytest.mark.asyncio +async def test_before_model_continues_without_steer(monkeypatch: pytest.MonkeyPatch): + """没有 Steer 时继续正常模型调用。""" + + async def should_end(run_id: str) -> bool: + return False + + monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) + runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) + + assert await SteerMiddleware().abefore_model({}, runtime) is None + +@pytest.mark.asyncio +async def test_before_model_ignores_context_without_run_id(monkeypatch: pytest.MonkeyPatch): + """缺少 Run 上下文时不查询队列。""" + called = False + + async def should_end(run_id: str) -> bool: + nonlocal called + called = True + return True + + monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) + runtime = SimpleNamespace(context=SimpleNamespace()) + + assert await SteerMiddleware().abefore_model({}, runtime) is None + assert called is False + +@pytest.mark.asyncio +async def test_after_model_ends_tool_free_turn_when_steer_arrives(monkeypatch: pytest.MonkeyPatch): + """模型轮次结束后才到达的 Steer 仍会让旧 Run 让位。""" + + async def should_end(run_id: str) -> bool: + return run_id == "run-1" + + monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) + runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) + + result = await SteerMiddleware().aafter_model( + {"messages": [AIMessage(content="已完成当前回答")]}, + runtime, + ) + + assert result == {"jump_to": "end"} + +@pytest.mark.asyncio +async def test_after_model_does_not_skip_tool_batch(monkeypatch: pytest.MonkeyPatch): + """模型生成工具调用时,Steer 不能跳过尚未执行的工具批次。""" + called = False + + async def should_end(run_id: str) -> bool: + nonlocal called + called = True + return True + + monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) + runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) + + result = await SteerMiddleware().aafter_model( + { + "messages": [ + AIMessage( + content="", + tool_calls=[{"id": "call-1", "name": "tool", "args": {}}], + ) + ] + }, + runtime, + ) + + assert result is None + assert called is False diff --git a/backend/test/unit/middlewares/test_steer_middleware.py b/backend/test/unit/middlewares/test_steer_middleware.py deleted file mode 100644 index de065416e..000000000 --- a/backend/test/unit/middlewares/test_steer_middleware.py +++ /dev/null @@ -1,97 +0,0 @@ -"""Steer Middleware 单元测试。""" - -from types import SimpleNamespace - -import pytest -from langchain_core.messages import AIMessage -from yuxi.agents.middlewares.steer import SteerMiddleware -from yuxi.services import agent_request_queue_service - -pytestmark = [pytest.mark.unit, pytest.mark.asyncio] - - -async def test_before_model_ends_run_when_steer_is_waiting(monkeypatch: pytest.MonkeyPatch): - """存在待处理 Steer 时,在下一次模型调用前结束当前 Graph。""" - - async def should_end(run_id: str) -> bool: - return run_id == "run-1" - - monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) - runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) - - result = await SteerMiddleware().abefore_model({}, runtime) - - assert result == {"jump_to": "end"} - - -async def test_before_model_continues_without_steer(monkeypatch: pytest.MonkeyPatch): - """没有 Steer 时继续正常模型调用。""" - - async def should_end(run_id: str) -> bool: - return False - - monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) - runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) - - assert await SteerMiddleware().abefore_model({}, runtime) is None - - -async def test_before_model_ignores_context_without_run_id(monkeypatch: pytest.MonkeyPatch): - """缺少 Run 上下文时不查询队列。""" - called = False - - async def should_end(run_id: str) -> bool: - nonlocal called - called = True - return True - - monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) - runtime = SimpleNamespace(context=SimpleNamespace()) - - assert await SteerMiddleware().abefore_model({}, runtime) is None - assert called is False - - -async def test_after_model_ends_tool_free_turn_when_steer_arrives(monkeypatch: pytest.MonkeyPatch): - """模型轮次结束后才到达的 Steer 仍会让旧 Run 让位。""" - - async def should_end(run_id: str) -> bool: - return run_id == "run-1" - - monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) - runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) - - result = await SteerMiddleware().aafter_model( - {"messages": [AIMessage(content="已完成当前回答")]}, - runtime, - ) - - assert result == {"jump_to": "end"} - - -async def test_after_model_does_not_skip_tool_batch(monkeypatch: pytest.MonkeyPatch): - """模型生成工具调用时,Steer 不能跳过尚未执行的工具批次。""" - called = False - - async def should_end(run_id: str) -> bool: - nonlocal called - called = True - return True - - monkeypatch.setattr(agent_request_queue_service, "should_end_run_for_steer", should_end) - runtime = SimpleNamespace(context=SimpleNamespace(run_id="run-1")) - - result = await SteerMiddleware().aafter_model( - { - "messages": [ - AIMessage( - content="", - tool_calls=[{"id": "call-1", "name": "tool", "args": {}}], - ) - ] - }, - runtime, - ) - - assert result is None - assert called is False diff --git a/backend/test/unit/services/test_agent_request_queue_service.py b/backend/test/unit/services/test_agent_request_queue_service.py index 57223d8a8..2a1a7fd32 100644 --- a/backend/test/unit/services/test_agent_request_queue_service.py +++ b/backend/test/unit/services/test_agent_request_queue_service.py @@ -19,6 +19,7 @@ cancel_queued_request, finalize_dispatch, finalize_intake, + guided_queued_request, intake_request, steer_queued_request, validate_queue_policy, @@ -1498,3 +1499,156 @@ async def resolve_config(*_args): assert result.status == "dispatched" assert result.run_id is not None + + +# ── guided: 注入当前 Run 而非排队等待 ── + + +@pytest.mark.asyncio +async def test_queued_request_can_be_upgraded_to_guided(session): + await _seed_thread(session) + await _seed_active_run(session) + await _create_request(session, request_id="request-guide") + + result = await guided_queued_request(request_id="request-guide", current_uid="user-1", db=session) + + request = await session.scalar( + select(AgentRunRequest).where(AgentRunRequest.request_id == "request-guide") + ) + assert result.status == "queued" + assert result.queue_policy == "guided" + assert request.queue_policy == "guided" + assert request.status == "queued" + + +@pytest.mark.asyncio +async def test_guided_upgrade_requires_running_main_chat(session): + from fastapi import HTTPException + + await _seed_thread(session) + await _create_request(session, request_id="request-guide") + + with pytest.raises(HTTPException) as exc_info: + await guided_queued_request(request_id="request-guide", current_uid="user-1", db=session) + + assert exc_info.value.status_code == 409 + assert exc_info.value.detail["code"] == "run_not_steerable" + + +@pytest.mark.asyncio +async def test_multiple_pending_guided_requests_are_allowed(session): + await _seed_thread(session) + await _seed_active_run(session) + session.add(Message(id=200, conversation_id=10, role="user", content="first")) + session.add(Message(id=201, conversation_id=10, role="user", content="second")) + await session.commit() + await _create_request(session, request_id="guided-1", msg_id=200, queue_policy="guided") + await _create_request(session, request_id="guided-2", msg_id=201, queue_policy="guided") + + from yuxi.repositories.agent_run_request_repository import AgentRunRequestRepository + + pending = await AgentRunRequestRepository(session).list_pending_guided( + uid="user-1", agent_slug="main", conversation_thread_id="t1" + ) + + # guided 与 steer 不同:允许多条同时等待,逐条注入。 + assert [item.request_id for item in pending] == ["guided-1", "guided-2"] + + +@pytest.mark.asyncio +async def test_guided_is_supported_queue_policy(): + assert validate_queue_policy("guided") == "guided" + assert "guided" not in NOT_IMPLEMENTED_QUEUE_POLICIES + + +@pytest.mark.asyncio +async def test_take_pending_guided_messages_restores_langchain_messages(session, monkeypatch): + import yuxi.services.agent_request_queue_service as queue_service + from langchain.messages import HumanMessage + + await _seed_thread(session) + await _seed_active_run(session) + session.add( + Message( + id=300, + conversation_id=10, + role="user", + content="补充:只关注 2024 年之后的数据", + extra_metadata={ + "request_id": "guided-1", + "raw_message": HumanMessage( + content="补充:只关注 2024 年之后的数据", id="langchain-msg-300" + ).model_dump(), + }, + ) + ) + await session.commit() + await _create_request(session, request_id="guided-1", msg_id=300, queue_policy="guided") + + published: list[tuple[str, str]] = [] + + @asynccontextmanager + async def fake_session_ctx(): + yield session + + async def fake_append_event(run_id, event_type, payload, *, thread_id=None): + published.append((run_id, event_type)) + + monkeypatch.setattr(queue_service.pg_manager, "get_async_session_context", fake_session_ctx) + monkeypatch.setattr(queue_service, "append_run_stream_event", fake_append_event) + + injected = await queue_service.take_pending_guided_messages("active-run") + + assert len(injected) == 1 + assert injected[0].content == "补充:只关注 2024 年之后的数据" + assert injected[0].id == "langchain-msg-300" + request = await session.scalar( + select(AgentRunRequest).where(AgentRunRequest.request_id == "guided-1") + ) + assert request.status == "injected" + assert published == [("active-run", "guided_injected")] + + +@pytest.mark.asyncio +async def test_take_pending_guided_messages_returns_empty_without_active_run(session, monkeypatch): + import yuxi.services.agent_request_queue_service as queue_service + + await _seed_thread(session) + await _create_request(session, request_id="guided-1", queue_policy="guided") + + @asynccontextmanager + async def fake_session_ctx(): + yield session + + monkeypatch.setattr(queue_service.pg_manager, "get_async_session_context", fake_session_ctx) + + assert await queue_service.take_pending_guided_messages("missing-run") == [] + + +@pytest.mark.asyncio +async def test_injected_guided_request_leaves_queue(session, monkeypatch): + import yuxi.services.agent_request_queue_service as queue_service + from yuxi.repositories.agent_run_request_repository import AgentRunRequestRepository + + await _seed_thread(session) + await _seed_active_run(session) + session.add(Message(id=310, conversation_id=10, role="user", content="补充")) + await session.commit() + await _create_request(session, request_id="guided-1", msg_id=310, queue_policy="guided") + + @asynccontextmanager + async def fake_session_ctx(): + yield session + + async def fake_append_event(run_id, event_type, payload, *, thread_id=None): + del run_id, event_type, payload, thread_id + + monkeypatch.setattr(queue_service.pg_manager, "get_async_session_context", fake_session_ctx) + monkeypatch.setattr(queue_service, "append_run_stream_event", fake_append_event) + + await queue_service.take_pending_guided_messages("active-run") + + queued = await AgentRunRequestRepository(session).list_queued( + uid="user-1", agent_slug="main", conversation_thread_id="t1" + ) + assert all(item.request_id != "guided-1" for item in queued) diff --git a/backend/test/unit/services/test_agent_run_service.py b/backend/test/unit/services/test_agent_run_service.py index 999933591..68efedebf 100644 --- a/backend/test/unit/services/test_agent_run_service.py +++ b/backend/test/unit/services/test_agent_run_service.py @@ -50,27 +50,6 @@ def _run_stream_event(seq: str, event_type: str, payload: dict) -> dict: } -def test_run_sse_poll_interval_caps_short_and_long_idle_periods(): - interval = agent_run_service.RUN_SSE_ACTIVE_POLL_SECONDS - short_idle_intervals = [] - for _ in range(6): - interval = agent_run_service._next_run_sse_poll_interval(interval, idle_seconds=30) - short_idle_intervals.append(interval) - - assert short_idle_intervals == [0.2, 0.4, 0.8, 1.0, 1.0, 1.0] - assert agent_run_service._next_run_sse_poll_interval(1.0, idle_seconds=120) == 2.0 - assert agent_run_service._next_run_sse_poll_interval(2.0, idle_seconds=120) == 4.0 - assert agent_run_service._next_run_sse_poll_interval(4.0, idle_seconds=120) == 4.0 - - -def test_run_sse_poll_jitter_stays_within_twenty_percent(monkeypatch: pytest.MonkeyPatch): - multipliers = iter([0.8, 1.2]) - monkeypatch.setattr(agent_run_service, "uniform", lambda _low, _high: next(multipliers)) - - assert agent_run_service._jitter_run_sse_poll_interval(1.0) == 0.8 - assert agent_run_service._jitter_run_sse_poll_interval(1.0) == 1.2 - - def test_openai_content_parts_build_and_restore_multimodal_message(): input_message = build_chat_input_message_from_openai_content( [ @@ -414,8 +393,13 @@ async def get_run_for_user(self, run_id: str, uid: str): del run_id, uid raise RuntimeError("db down") + async def fake_blocking_read(run_id: str, *, after_seq: str, block_ms: int, limit: int): + del run_id, after_seq, block_ms, limit + return [] + monkeypatch.setattr(agent_run_service.pg_manager, "get_async_session_context", fake_session_ctx) monkeypatch.setattr(agent_run_service, "AgentRunRepository", BrokenRepo) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", fake_blocking_read) chunks = [] async for chunk in agent_run_service.stream_agent_run_events( @@ -440,7 +424,7 @@ async def unexpected_list_events(*_args, **_kwargs): pytest.fail("未授权连接不得读取 Redis Run 事件") monkeypatch.setattr(agent_run_service, "_load_stream_run_for_user", fake_load_run) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", unexpected_list_events) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", unexpected_list_events) chunks = [ chunk @@ -471,7 +455,7 @@ async def fake_list_events(*_args, **_kwargs): monkeypatch.setattr(agent_run_service, "_load_stream_run_for_user", fake_load_run) monkeypatch.setattr(agent_run_service, "_load_stream_run", broken_refresh) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", fake_list_events) monkeypatch.setattr(agent_run_service, "RUN_SSE_STATUS_POLL_SECONDS", 0.0) chunks = [ @@ -504,8 +488,8 @@ async def get_run_for_user(self, run_id: str, uid: str): calls = {"count": 0} - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - del run_id, after_seq, limit + async def fake_blocking_read(run_id: str, *, after_seq: str, block_ms: int, limit: int): + del run_id, after_seq, block_ms, limit calls["count"] += 1 if calls["count"] == 1: return [ @@ -540,7 +524,7 @@ async def fake_list_events(run_id: str, *, after_seq: str, limit: int): monkeypatch.setattr(agent_run_service.pg_manager, "get_async_session_context", fake_session_ctx) monkeypatch.setattr(agent_run_service, "AgentRunRepository", Repo) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", fake_blocking_read) chunks = [] async for chunk in agent_run_service.stream_agent_run_events( @@ -556,174 +540,6 @@ async def fake_list_events(run_id: str, *, after_seq: str, limit: int): assert "id: 1700000000001-0" in chunks[-1] -@pytest.mark.asyncio -async def test_stream_agent_run_events_decouples_pg_checks_from_redis_polling( - monkeypatch: pytest.MonkeyPatch, -): - """高频 Redis 空轮询不得同步放大 PostgreSQL 可见性查询。""" - pg_reads = 0 - - async def fake_load_run(run_id: str, uid: str): - nonlocal pg_reads - del run_id, uid - pg_reads += 1 - return _run_state() - - async def fake_refresh_run(run_id: str): - nonlocal pg_reads - del run_id - pg_reads += 1 - return _run_state() - - redis_reads = 0 - - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - nonlocal redis_reads - del run_id, after_seq, limit - redis_reads += 1 - if redis_reads < 4: - return [] - return [_run_stream_event("1700000000004-0", "end", {"status": "completed"})] - - sleep_intervals = [] - - async def fake_sleep(seconds: float): - sleep_intervals.append(seconds) - - monkeypatch.setattr(agent_run_service, "_load_stream_run_for_user", fake_load_run) - monkeypatch.setattr(agent_run_service, "_load_stream_run", fake_refresh_run) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) - monkeypatch.setattr(agent_run_service.asyncio, "sleep", fake_sleep) - monkeypatch.setattr(agent_run_service, "monotonic", lambda: 0.0) - monkeypatch.setattr(agent_run_service, "uniform", lambda _low, _high: 1.0) - - chunks = [] - async for chunk in agent_run_service.stream_agent_run_events( - run_id="run-1", - after_seq="0", - current_uid="user-1", - ): - chunks.append(chunk) - - assert pg_reads == 1 - assert redis_reads == 4 - assert sleep_intervals == [0.1, 0.2, 0.4] - assert chunks[-1].startswith("event: end") - - -@pytest.mark.asyncio -async def test_stream_agent_run_events_resets_adaptive_poll_after_event( - monkeypatch: pytest.MonkeyPatch, -): - """任一 Run 事件都必须把退避后的轮询恢复到低延迟档。""" - - async def fake_load_run(run_id: str, uid: str): - del run_id, uid - return _run_state() - - async def fake_refresh_run(run_id: str): - del run_id - return _run_state() - - redis_results = iter( - [ - [], - [], - [_run_stream_event("1700000000001-0", "messages", {"items": []})], - [], - [_run_stream_event("1700000000002-0", "end", {"status": "completed"})], - ] - ) - - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - del run_id, after_seq, limit - return next(redis_results) - - sleep_intervals = [] - - async def fake_sleep(seconds: float): - sleep_intervals.append(seconds) - - monkeypatch.setattr(agent_run_service, "_load_stream_run_for_user", fake_load_run) - monkeypatch.setattr(agent_run_service, "_load_stream_run", fake_refresh_run) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) - monkeypatch.setattr(agent_run_service.asyncio, "sleep", fake_sleep) - monkeypatch.setattr(agent_run_service, "monotonic", lambda: 0.0) - monkeypatch.setattr(agent_run_service, "uniform", lambda _low, _high: 1.0) - - chunks = [] - async for chunk in agent_run_service.stream_agent_run_events( - run_id="run-1", - after_seq="0", - current_uid="user-1", - ): - chunks.append(chunk) - - assert sleep_intervals == [0.1, 0.2, 0.1, 0.1] - assert [chunk.splitlines()[0] for chunk in chunks] == ["event: messages", "event: end"] - - -@pytest.mark.asyncio -async def test_stream_agent_run_events_refreshes_pg_before_cleanup_fallback( - monkeypatch: pytest.MonkeyPatch, -): - """Redis 缺少 end 时,低频 PG 探测仍须等待 cleanup fence 后补发终态。""" - visibility_reads = 0 - status_reads = 0 - - async def fake_load_run(run_id: str, uid: str): - nonlocal visibility_reads - del run_id, uid - visibility_reads += 1 - return _run_state("completed", cleanup_pending=True) - - async def fake_refresh_run(run_id: str): - nonlocal status_reads - del run_id - status_reads += 1 - return _run_state("completed") - - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - del run_id, after_seq, limit - return [] - - async def fake_last_stream_seq(run_id: str): - del run_id - return "0-0" - - clock = 0.0 - sleep_intervals = [] - - async def fake_sleep(seconds: float): - nonlocal clock - sleep_intervals.append(seconds) - clock += seconds - - monkeypatch.setattr(agent_run_service, "_load_stream_run_for_user", fake_load_run) - monkeypatch.setattr(agent_run_service, "_load_stream_run", fake_refresh_run) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) - monkeypatch.setattr(agent_run_service, "get_last_run_stream_seq", fake_last_stream_seq) - monkeypatch.setattr(agent_run_service.asyncio, "sleep", fake_sleep) - monkeypatch.setattr(agent_run_service, "monotonic", lambda: clock) - monkeypatch.setattr(agent_run_service, "uniform", lambda _low, _high: 1.2) - monkeypatch.setattr(agent_run_service, "RUN_SSE_LONG_IDLE_AFTER_SECONDS", 0.0) - - chunks = [] - async for chunk in agent_run_service.stream_agent_run_events( - run_id="run-1", - after_seq="0", - current_uid="user-1", - verbose=False, - ): - chunks.append(chunk) - - assert visibility_reads == 1 - assert status_reads == 1 - assert clock == agent_run_service.RUN_SSE_STATUS_POLL_SECONDS - assert sleep_intervals[-1] < 3.2 * 1.2 - assert len(chunks) == 1 - assert chunks[0].startswith("event: end") - assert _sse_data(chunks[0])["payload"] == {"status": "completed"} @pytest.mark.asyncio @@ -740,8 +556,8 @@ async def get_run_for_user(self, run_id: str, uid: str): del run_id, uid return SimpleNamespace(status="completed", conversation_thread_id="thread-1") - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - del run_id, after_seq, limit + async def fake_blocking_read(run_id: str, *, after_seq: str, block_ms: int, limit: int): + del run_id, after_seq, block_ms, limit return [ { "seq": "1700000000000-0", @@ -885,7 +701,7 @@ async def fake_list_events(run_id: str, *, after_seq: str, limit: int): monkeypatch.setattr(agent_run_service.pg_manager, "get_async_session_context", fake_session_ctx) monkeypatch.setattr(agent_run_service, "AgentRunRepository", Repo) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", fake_blocking_read) chunks = [] async for chunk in agent_run_service.stream_agent_run_events( @@ -949,8 +765,17 @@ async def get_run_for_user(self, run_id: str, uid: str): runtime_cleanup_pending=False, ) - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - del run_id, after_seq, limit + async def get_run(self, run_id: str): + del run_id + return SimpleNamespace( + status="completed", + conversation_thread_id="thread-1", + request_id="req-1", + runtime_cleanup_pending=False, + ) + + async def fake_blocking_read(run_id: str, *, after_seq: str, block_ms: int, limit: int): + del run_id, after_seq, block_ms, limit return [] async def fake_get_last_run_stream_seq(run_id: str): @@ -959,7 +784,7 @@ async def fake_get_last_run_stream_seq(run_id: str): monkeypatch.setattr(agent_run_service.pg_manager, "get_async_session_context", fake_session_ctx) monkeypatch.setattr(agent_run_service, "AgentRunRepository", Repo) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", fake_blocking_read) monkeypatch.setattr(agent_run_service, "get_last_run_stream_seq", fake_get_last_run_stream_seq) chunks = [] @@ -1002,21 +827,18 @@ async def get_run_for_user(self, run_id: str, uid: str): runtime_cleanup_pending=True, ) - async def fake_list_events(run_id: str, *, after_seq: str, limit: int): - del run_id, after_seq, limit - return [] - - sleep_calls = 0 + calls = {"count": 0} - async def stop_after_one_poll(_seconds: float): - nonlocal sleep_calls - sleep_calls += 1 + async def fake_blocking_read(run_id: str, *, after_seq: str, block_ms: int, limit: int): + del run_id, after_seq, block_ms, limit + calls["count"] += 1 + if calls["count"] == 1: + return [] raise agent_run_service.asyncio.CancelledError monkeypatch.setattr(agent_run_service.pg_manager, "get_async_session_context", fake_session_ctx) monkeypatch.setattr(agent_run_service, "AgentRunRepository", Repo) - monkeypatch.setattr(agent_run_service, "list_run_stream_events", fake_list_events) - monkeypatch.setattr(agent_run_service.asyncio, "sleep", stop_after_one_poll) + monkeypatch.setattr(agent_run_service, "blocking_read_run_stream_events", fake_blocking_read) chunks = [] async for chunk in agent_run_service.stream_agent_run_events( @@ -1027,7 +849,7 @@ async def stop_after_one_poll(_seconds: float): ): chunks.append(chunk) - assert sleep_calls == 1 + assert calls["count"] == 2 assert not any(chunk.startswith("event: end") for chunk in chunks) diff --git a/backend/test/unit/services/test_run_queue_service.py b/backend/test/unit/services/test_run_queue_service.py index 8c3255144..3f4f60474 100644 --- a/backend/test/unit/services/test_run_queue_service.py +++ b/backend/test/unit/services/test_run_queue_service.py @@ -67,6 +67,18 @@ async def xrevrange(self, key: str, max: str, min: str, count: int): rows = list(reversed(self.streams.get(key, []))) return rows[:count] + async def xread(self, streams: dict, block: int, count: int): + del block + result = [] + for key, start in streams.items(): + rows = list(self.streams.get(key, [])) + if start not in ("0-0", ""): + rows = [(event_id, fields) for event_id, fields in rows if event_id > start] + rows = rows[:count] + if rows: + result.append([key, rows]) + return result + @pytest.mark.asyncio async def test_run_stream_event_roundtrip(monkeypatch: pytest.MonkeyPatch): @@ -142,6 +154,42 @@ async def fake_get_async_redis_client(): ] +@pytest.mark.asyncio +async def test_blocking_read_run_stream_events_reads_after_cursor(monkeypatch: pytest.MonkeyPatch): + fake_redis = _FakeStreamRedis() + + async def fake_get_async_redis_client(): + return fake_redis + + monkeypatch.setattr(run_queue_service, "get_async_redis_client", fake_get_async_redis_client) + + run_id = "run-1" + seq1 = await run_queue_service.append_run_stream_event(run_id, "loading", {"items": [1]}) + seq2 = await run_queue_service.append_run_stream_event(run_id, "finished", {"chunk": {}}) + + events = await run_queue_service.blocking_read_run_stream_events(run_id, after_seq="0-0", block_ms=100) + assert [item["event_type"] for item in events] == ["loading", "finished"] + + tail = await run_queue_service.blocking_read_run_stream_events(run_id, after_seq=seq1, block_ms=100) + assert [item["seq"] for item in tail] == [seq2] + + +@pytest.mark.asyncio +async def test_blocking_read_run_stream_events_empty_when_no_new_events(monkeypatch: pytest.MonkeyPatch): + fake_redis = _FakeStreamRedis() + + async def fake_get_async_redis_client(): + return fake_redis + + monkeypatch.setattr(run_queue_service, "get_async_redis_client", fake_get_async_redis_client) + + run_id = "run-1" + seq = await run_queue_service.append_run_stream_event(run_id, "loading", {"items": [1]}) + + events = await run_queue_service.blocking_read_run_stream_events(run_id, after_seq=seq, block_ms=100) + assert events == [] + + def test_normalize_after_seq_stream_id_only(): assert run_queue_service.normalize_after_seq(None) == "0-0" assert run_queue_service.normalize_after_seq("1700000000000-3") == "1700000000000-3"