Add runtime guidance input queue
This commit is contained in:
+171
-1
@@ -71,13 +71,19 @@ from src.run_state_store import ACTIVE_RUN_STATUSES, RunStateStore
|
||||
from src.session_store import (
|
||||
DEFAULT_AGENT_SESSION_DIR,
|
||||
StoredAgentSession,
|
||||
agent_session_delta,
|
||||
cancel_session_input_queue,
|
||||
consume_session_guidance,
|
||||
delete_agent_session,
|
||||
deserialize_runtime_config,
|
||||
enqueue_session_input,
|
||||
list_agent_sessions,
|
||||
list_session_input_queue,
|
||||
load_agent_session,
|
||||
save_agent_session,
|
||||
serialize_model_config,
|
||||
serialize_runtime_config,
|
||||
update_session_input_queue_item,
|
||||
)
|
||||
from src.token_budget import calculate_token_budget
|
||||
|
||||
@@ -1061,6 +1067,22 @@ class RunCancelRequest(BaseModel):
|
||||
run_id: str | None = None
|
||||
|
||||
|
||||
class SessionInputQueueCreate(BaseModel):
|
||||
content: str = Field(min_length=1, max_length=20000)
|
||||
kind: str = 'next_turn'
|
||||
run_id: str | None = None
|
||||
account_id: str | None = None
|
||||
metadata: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class SessionInputQueueUpdate(BaseModel):
|
||||
content: str | None = Field(default=None, max_length=20000)
|
||||
kind: str | None = None
|
||||
status: str | None = None
|
||||
run_id: str | None = None
|
||||
account_id: str | None = None
|
||||
|
||||
|
||||
class StateUpdate(BaseModel):
|
||||
model: str | None = None
|
||||
base_url: str | None = None
|
||||
@@ -1737,6 +1759,126 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
raise HTTPException(status_code=404, detail='Session not found')
|
||||
return _serialize_stored_session(stored)
|
||||
|
||||
@app.get('/api/sessions/{session_id}/state')
|
||||
async def get_session_state_delta(
|
||||
session_id: str,
|
||||
account_id: str | None = None,
|
||||
after_message_seq: int = 0,
|
||||
after_event_seq: int = 0,
|
||||
) -> dict[str, Any]:
|
||||
directory = state.account_paths(account_id)['sessions']
|
||||
safe_id = _safe_session_id(session_id)
|
||||
if safe_id is None:
|
||||
raise HTTPException(status_code=400, detail='Invalid session id')
|
||||
account_key = state._account_key(account_id)
|
||||
message_delta = agent_session_delta(
|
||||
safe_id,
|
||||
directory=directory,
|
||||
after_message_seq=after_message_seq,
|
||||
)
|
||||
event_delta = state.run_state_store.events_since(
|
||||
account_key,
|
||||
safe_id,
|
||||
after_event_seq=after_event_seq,
|
||||
)
|
||||
run_status = await latest_run(safe_id, account_id=account_id)
|
||||
queue_items = list_session_input_queue(safe_id, directory=directory)
|
||||
return {
|
||||
'session_id': safe_id,
|
||||
'session': message_delta.get('session', {'session_id': safe_id}),
|
||||
'messages': message_delta.get('messages', []),
|
||||
'latest_message_seq': message_delta.get('latest_message_seq', 0),
|
||||
'activity_events': event_delta.get('events', []),
|
||||
'latest_event_seq': event_delta.get('latest_event_seq', 0),
|
||||
'run': run_status,
|
||||
'input_queue': queue_items,
|
||||
}
|
||||
|
||||
@app.get('/api/sessions/{session_id}/input-queue')
|
||||
async def get_session_input_queue(
|
||||
session_id: str,
|
||||
account_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
directory = state.account_paths(account_id)['sessions']
|
||||
safe_id = _safe_session_id(session_id)
|
||||
if safe_id is None:
|
||||
raise HTTPException(status_code=400, detail='Invalid session id')
|
||||
return {
|
||||
'session_id': safe_id,
|
||||
'items': list_session_input_queue(safe_id, directory=directory),
|
||||
}
|
||||
|
||||
@app.post('/api/sessions/{session_id}/input-queue')
|
||||
async def create_session_input_queue_item(
|
||||
session_id: str,
|
||||
payload: SessionInputQueueCreate,
|
||||
) -> dict[str, Any]:
|
||||
directory = state.account_paths(payload.account_id)['sessions']
|
||||
safe_id = _safe_session_id(session_id)
|
||||
if safe_id is None:
|
||||
raise HTTPException(status_code=400, detail='Invalid session id')
|
||||
content = payload.content.strip()
|
||||
if not content:
|
||||
raise HTTPException(status_code=400, detail='Empty input content')
|
||||
run_id = payload.run_id
|
||||
if payload.kind == 'guidance' and not run_id:
|
||||
account_key = state._account_key(payload.account_id)
|
||||
latest = state.run_state_store.snapshot_latest(account_key, safe_id)
|
||||
if latest and latest.get('status') in ACTIVE_RUN_STATUSES:
|
||||
run_id = str(latest.get('run_id') or '')
|
||||
item = enqueue_session_input(
|
||||
safe_id,
|
||||
directory=directory,
|
||||
run_id=run_id,
|
||||
kind=payload.kind,
|
||||
content=content,
|
||||
metadata=payload.metadata,
|
||||
)
|
||||
return {'session_id': safe_id, 'item': item}
|
||||
|
||||
@app.patch('/api/sessions/{session_id}/input-queue/{item_id}')
|
||||
async def update_session_input_queue(
|
||||
session_id: str,
|
||||
item_id: int,
|
||||
payload: SessionInputQueueUpdate,
|
||||
) -> dict[str, Any]:
|
||||
directory = state.account_paths(payload.account_id)['sessions']
|
||||
safe_id = _safe_session_id(session_id)
|
||||
if safe_id is None:
|
||||
raise HTTPException(status_code=400, detail='Invalid session id')
|
||||
item = update_session_input_queue_item(
|
||||
safe_id,
|
||||
item_id,
|
||||
directory=directory,
|
||||
content=payload.content.strip() if payload.content is not None else None,
|
||||
kind=payload.kind,
|
||||
status=payload.status,
|
||||
run_id=payload.run_id,
|
||||
)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail='Queue item not found')
|
||||
return {'session_id': safe_id, 'item': item}
|
||||
|
||||
@app.delete('/api/sessions/{session_id}/input-queue/{item_id}')
|
||||
async def delete_session_input_queue(
|
||||
session_id: str,
|
||||
item_id: int,
|
||||
account_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
directory = state.account_paths(account_id)['sessions']
|
||||
safe_id = _safe_session_id(session_id)
|
||||
if safe_id is None:
|
||||
raise HTTPException(status_code=400, detail='Invalid session id')
|
||||
item = update_session_input_queue_item(
|
||||
safe_id,
|
||||
item_id,
|
||||
directory=directory,
|
||||
status='cancelled',
|
||||
)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail='Queue item not found')
|
||||
return {'session_id': safe_id, 'item': item}
|
||||
|
||||
@app.delete('/api/sessions/{session_id}')
|
||||
async def delete_session(session_id: str, account_id: str | None = None) -> dict[str, Any]:
|
||||
directory = state.account_paths(account_id)['sessions']
|
||||
@@ -2356,8 +2498,18 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
status='cancelled',
|
||||
stage='用户已取消',
|
||||
)
|
||||
queue_cancelled = cancel_session_input_queue(
|
||||
safe_id,
|
||||
state.account_paths(payload.account_id)['sessions'],
|
||||
run_id=payload.run_id,
|
||||
)
|
||||
cancelled_ids.update(stored_cancelled)
|
||||
cancelled = bool(cancelled_ids or cancelled_bg_task_ids or abandoned_runtime)
|
||||
cancelled = bool(
|
||||
cancelled_ids
|
||||
or cancelled_bg_task_ids
|
||||
or abandoned_runtime
|
||||
or queue_cancelled
|
||||
)
|
||||
if cancelled:
|
||||
_mark_session_interrupted(
|
||||
state.account_paths(payload.account_id)['sessions'],
|
||||
@@ -2370,6 +2522,7 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
'cancelled_run_ids': sorted(cancelled_ids),
|
||||
'cancelled_bg_task_ids': sorted(cancelled_bg_task_ids),
|
||||
'abandoned_runtime': abandoned_runtime,
|
||||
'cancelled_input_queue_items': queue_cancelled,
|
||||
}
|
||||
|
||||
def _run_chat_payload(
|
||||
@@ -2527,6 +2680,22 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
except Exception: # noqa: BLE001
|
||||
return False
|
||||
|
||||
def _drain_runtime_guidance() -> tuple[dict[str, object], ...]:
|
||||
items = consume_session_guidance(
|
||||
requested_session_id,
|
||||
run_record.run_id,
|
||||
directory=session_directory,
|
||||
)
|
||||
return tuple(
|
||||
{
|
||||
'id': int(item.get('id') or 0),
|
||||
'content': str(item.get('content') or ''),
|
||||
}
|
||||
for item in items
|
||||
)
|
||||
|
||||
previous_guidance_provider = agent.runtime_guidance_provider
|
||||
agent.runtime_guidance_provider = _drain_runtime_guidance
|
||||
agent.tool_context = replace(
|
||||
agent.tool_context,
|
||||
cancel_event=run_record.cancel_event,
|
||||
@@ -2602,6 +2771,7 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
agent.runtime_guidance_provider = previous_guidance_provider
|
||||
agent.tool_context = replace(
|
||||
agent.tool_context,
|
||||
cancel_event=None,
|
||||
|
||||
Reference in New Issue
Block a user