Auto-run queued session inputs
This commit is contained in:
@@ -74,6 +74,7 @@ from src.session_store import (
|
||||
append_agent_display_message,
|
||||
agent_session_delta,
|
||||
cancel_session_input_queue,
|
||||
consume_next_session_input,
|
||||
consume_session_guidance,
|
||||
delete_agent_session,
|
||||
deserialize_runtime_config,
|
||||
@@ -2538,6 +2539,41 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
'cancelled_input_queue_items': queue_cancelled,
|
||||
}
|
||||
|
||||
def _schedule_next_queued_turn(
|
||||
*,
|
||||
account_id: str | None,
|
||||
session_id: str,
|
||||
directory: Path,
|
||||
) -> dict[str, Any] | None:
|
||||
item = consume_next_session_input(session_id, directory=directory)
|
||||
if item is None:
|
||||
return None
|
||||
content = str(item.get('content') or '').strip()
|
||||
if not content:
|
||||
return None
|
||||
try:
|
||||
return _start_chat_run(
|
||||
ChatRequest(
|
||||
prompt=content,
|
||||
account_id=account_id,
|
||||
session_id=session_id,
|
||||
resume_session_id=session_id,
|
||||
)
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
update_session_input_queue_item(
|
||||
session_id,
|
||||
int(item.get('id') or 0),
|
||||
directory=directory,
|
||||
status='pending',
|
||||
)
|
||||
print(
|
||||
f'[input-queue] failed to schedule queued turn '
|
||||
f'session={session_id} item={item.get("id")}: {exc}',
|
||||
flush=True,
|
||||
)
|
||||
return None
|
||||
|
||||
def _run_chat_payload(
|
||||
request: ChatRequest,
|
||||
event_sink: Any | None = None,
|
||||
@@ -2845,6 +2881,12 @@ def create_app(state: AgentState) -> FastAPI:
|
||||
result.session_id,
|
||||
status='cancelled',
|
||||
)
|
||||
if not run_record.cancel_event.is_set() and result.session_id:
|
||||
_schedule_next_queued_turn(
|
||||
account_id=request.account_id,
|
||||
session_id=result.session_id,
|
||||
directory=session_directory,
|
||||
)
|
||||
return payload
|
||||
|
||||
def _start_chat_run(request: ChatRequest) -> dict[str, Any]:
|
||||
|
||||
@@ -471,6 +471,46 @@ def consume_session_guidance(
|
||||
return [_queue_row_to_dict(row) for row in rows]
|
||||
|
||||
|
||||
def consume_next_session_input(
|
||||
session_id: str,
|
||||
directory: Path | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
target_dir = directory or DEFAULT_AGENT_SESSION_DIR
|
||||
now = time.time()
|
||||
with _connect_agent_session_db(target_dir) as conn:
|
||||
row = conn.execute(
|
||||
"""
|
||||
select *
|
||||
from agent_input_queue
|
||||
where session_id = ?
|
||||
and status = 'pending'
|
||||
and kind = 'next_turn'
|
||||
order by id asc
|
||||
limit 1
|
||||
""",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
cursor = conn.execute(
|
||||
"""
|
||||
update agent_input_queue
|
||||
set status = 'consumed',
|
||||
updated_at = ?,
|
||||
consumed_at = ?
|
||||
where id = ? and status = 'pending'
|
||||
""",
|
||||
(now, now, int(row['id'])),
|
||||
)
|
||||
if int(cursor.rowcount or 0) <= 0:
|
||||
return None
|
||||
item = _queue_row_to_dict(row)
|
||||
item['status'] = 'consumed'
|
||||
item['updated_at'] = now
|
||||
item['consumed_at'] = now
|
||||
return item
|
||||
|
||||
|
||||
def delete_agent_session(session_id: str, directory: Path | None = None) -> bool:
|
||||
target_dir = directory or DEFAULT_AGENT_SESSION_DIR
|
||||
deleted = False
|
||||
|
||||
@@ -20,8 +20,11 @@ from src.session_store import (
|
||||
_deserialize_output_schema,
|
||||
_optional_float,
|
||||
_optional_int,
|
||||
consume_next_session_input,
|
||||
deserialize_model_config,
|
||||
deserialize_runtime_config,
|
||||
enqueue_session_input,
|
||||
list_session_input_queue,
|
||||
load_agent_session,
|
||||
load_session,
|
||||
save_agent_session,
|
||||
@@ -234,6 +237,58 @@ class TestStoredAgentSessionRoundTrip(unittest.TestCase):
|
||||
self.assertEqual(loaded.plugin_state, {})
|
||||
|
||||
|
||||
class TestSessionInputQueue(unittest.TestCase):
|
||||
def test_consume_next_turn_uses_fifo_and_preserves_guidance(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
directory = Path(td) / 'sessions'
|
||||
directory.mkdir()
|
||||
first = enqueue_session_input(
|
||||
'session-1',
|
||||
directory=directory,
|
||||
content='first queued turn',
|
||||
kind='next_turn',
|
||||
)
|
||||
guidance = enqueue_session_input(
|
||||
'session-1',
|
||||
directory=directory,
|
||||
content='guide current run',
|
||||
kind='guidance',
|
||||
)
|
||||
second = enqueue_session_input(
|
||||
'session-1',
|
||||
directory=directory,
|
||||
content='second queued turn',
|
||||
kind='next_turn',
|
||||
)
|
||||
|
||||
consumed_first = consume_next_session_input(
|
||||
'session-1',
|
||||
directory=directory,
|
||||
)
|
||||
consumed_second = consume_next_session_input(
|
||||
'session-1',
|
||||
directory=directory,
|
||||
)
|
||||
consumed_none = consume_next_session_input(
|
||||
'session-1',
|
||||
directory=directory,
|
||||
)
|
||||
pending = list_session_input_queue('session-1', directory=directory)
|
||||
|
||||
self.assertIsNotNone(consumed_first)
|
||||
self.assertIsNotNone(consumed_second)
|
||||
assert consumed_first is not None
|
||||
assert consumed_second is not None
|
||||
self.assertEqual(consumed_first['id'], first['id'])
|
||||
self.assertEqual(consumed_first['content'], 'first queued turn')
|
||||
self.assertEqual(consumed_first['status'], 'consumed')
|
||||
self.assertEqual(consumed_second['id'], second['id'])
|
||||
self.assertEqual(consumed_second['content'], 'second queued turn')
|
||||
self.assertIsNone(consumed_none)
|
||||
self.assertEqual([item['id'] for item in pending], [guidance['id']])
|
||||
self.assertEqual(pending[0]['kind'], 'guidance')
|
||||
|
||||
|
||||
class TestModelConfigSerialization(unittest.TestCase):
|
||||
"""serialize_model_config + deserialize_model_config round-trip."""
|
||||
|
||||
@@ -389,8 +444,8 @@ class TestRuntimeConfigSerialization(unittest.TestCase):
|
||||
config = deserialize_runtime_config(payload)
|
||||
|
||||
self.assertEqual(config.max_turns, 50)
|
||||
self.assertAlmostEqual(config.command_timeout_seconds, 30.0)
|
||||
self.assertEqual(config.max_output_chars, 12000)
|
||||
self.assertAlmostEqual(config.command_timeout_seconds, 300.0)
|
||||
self.assertEqual(config.max_output_chars, 50000)
|
||||
self.assertFalse(config.stream_model_responses)
|
||||
self.assertIsNone(config.auto_snip_threshold_tokens)
|
||||
self.assertIsNone(config.auto_compact_threshold_tokens)
|
||||
|
||||
Reference in New Issue
Block a user