Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 27 additions & 3 deletions backend/package/yuxi/agents/middlewares/steer.py
Original file line number Diff line number Diff line change
@@ -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):
Expand All @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
*,
Expand Down
143 changes: 137 additions & 6 deletions backend/package/yuxi/services/agent_request_queue_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -43,16 +44,20 @@
)
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"
REQUEST_STATUS_DISPATCHED = "dispatched"
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"
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading