graphharbor 0.13.0.post33__tar.gz → 0.13.0.post35__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (22) hide show
  1. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/PKG-INFO +2 -2
  2. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/pyproject.toml +2 -2
  3. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/protocol_api.py +98 -29
  4. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/server.py +3 -291
  5. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/store_api.py +14 -11
  6. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/streaming.py +93 -23
  7. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_application_authorization.py +32 -0
  8. graphharbor-0.13.0.post35/tests/test_stream_heartbeat_resilience.py +344 -0
  9. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/.gitignore +0 -0
  10. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/LICENSE +0 -0
  11. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/README.md +0 -0
  12. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/__init__.py +0 -0
  13. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/__main__.py +0 -0
  14. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/cli.py +0 -0
  15. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/core_api.py +0 -0
  16. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/mcp_transport.py +0 -0
  17. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_cli.py +0 -0
  18. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_graph_discovery.py +0 -0
  19. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_mcp_transport.py +0 -0
  20. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_official_protocol_compare.py +0 -0
  21. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_server_paths.py +0 -0
  22. {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_thread_state_projection.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: graphharbor
3
- Version: 0.13.0.post33
3
+ Version: 0.13.0.post35
4
4
  Summary: GraphHarbor CLI for a self-hosted LangGraph Agent Server on PostgreSQL and Redis
5
5
  Project-URL: Homepage, https://github.com/ljxpython/graphharbor
6
6
  Project-URL: Repository, https://github.com/ljxpython/graphharbor
@@ -26,7 +26,7 @@ Classifier: Topic :: Software Development :: Libraries :: Python Modules
26
26
  Classifier: Typing :: Typed
27
27
  Requires-Python: >=3.11
28
28
  Requires-Dist: click>=8.1.7
29
- Requires-Dist: graphharbor-runtime==0.13.0.post33
29
+ Requires-Dist: graphharbor-runtime==0.13.0.post35
30
30
  Requires-Dist: langgraph-cli<0.5,>=0.4.0
31
31
  Requires-Dist: mcp<2,>=1.23
32
32
  Requires-Dist: pyfiglet>=1.0.0
@@ -1,7 +1,7 @@
1
1
  [project]
2
2
  name = "graphharbor"
3
3
  # Lockstep with graphharbor-runtime; independent of langgraph-api releases.
4
- version = "0.13.0.post33"
4
+ version = "0.13.0.post35"
5
5
  description = "GraphHarbor CLI for a self-hosted LangGraph Agent Server on PostgreSQL and Redis"
6
6
  readme = "README.md"
7
7
  license = "MIT"
@@ -50,7 +50,7 @@ dependencies = [
50
50
  "mcp>=1.23,<2",
51
51
  "uvicorn[standard]>=0.51.0",
52
52
  # Exact same version as graphharbor-runtime (lockstep releases).
53
- "graphharbor-runtime==0.13.0.post33",
53
+ "graphharbor-runtime==0.13.0.post35",
54
54
  ]
55
55
 
56
56
  [project.urls]
@@ -189,6 +189,12 @@ async def protocol_commands(request: Request) -> JSONResponse:
189
189
  key: value for key, value in persisted.interrupts.items() if key != interrupt_id
190
190
  }
191
191
  persisted.status = "busy"
192
+ old_run = await conn.session.get(RunRow, latest.run_id)
193
+ if old_run is not None:
194
+ old_kwargs = dict(old_run.kwargs or {})
195
+ if old_kwargs.get("stream_resumable") is True:
196
+ old_kwargs["stream_resumable"] = False
197
+ old_run.kwargs = old_kwargs
192
198
  return JSONResponse(
193
199
  {
194
200
  "id": command_id,
@@ -271,6 +277,15 @@ def _frame(wire: dict[str, Any]) -> str:
271
277
  return f"id: {seq}\nevent: event\ndata: {data}\n\n"
272
278
 
273
279
 
280
+ def _safe_uuid(value: Any) -> UUID | None:
281
+ if not value:
282
+ return None
283
+ try:
284
+ return UUID(str(value))
285
+ except (ValueError, TypeError):
286
+ return None
287
+
288
+
274
289
  async def protocol_event_stream(request: Request) -> JSONResponse | StreamingResponse:
275
290
  from langhost.streaming import _resumable_run_ids
276
291
 
@@ -300,67 +315,117 @@ async def protocol_event_stream(request: Request) -> JSONResponse | StreamingRes
300
315
  queue = await manager.add_thread_stream(thread_id)
301
316
  try:
302
317
  watermark, replay = await _load_protocol_events(thread_id, since)
318
+ if since and since < watermark:
319
+ await manager.remove_thread_stream(thread_id, queue)
320
+ return JSONResponse(
321
+ {"code": "cursor_expired", "detail": "cursor_expired", "recovery": "thread_snapshot"},
322
+ status_code=410,
323
+ )
324
+ parsed_run_ids = {
325
+ uid
326
+ for wire in replay
327
+ if (uid := _safe_uuid(wire.get("params", {}).get("run_id"))) is not None
328
+ }
329
+ resumable = await _resumable_run_ids(parsed_run_ids)
330
+ active_interrupt_ids = (
331
+ set(thread.interrupts.keys()) if isinstance(thread.interrupts, dict) else set()
332
+ )
333
+
334
+ def _is_active_protocol_event(w: dict[str, Any]) -> bool:
335
+ method = w.get("method")
336
+ params = w.get("params") or {}
337
+ data = params.get("data") or {}
338
+ is_input_req = (
339
+ method == "input.requested"
340
+ or (method == "input" and data.get("event") == "requested")
341
+ or data.get("event") == "input.requested"
342
+ )
343
+ if is_input_req:
344
+ iid = str(data.get("interrupt_id") or params.get("interrupt_id") or "")
345
+ if iid and iid not in active_interrupt_ids:
346
+ return False
347
+ return True
348
+
349
+ replay = [
350
+ wire
351
+ for wire in replay
352
+ if ((uid := _safe_uuid(wire.get("params", {}).get("run_id"))) is None or uid in resumable)
353
+ and _is_active_protocol_event(wire)
354
+ ]
303
355
  except Exception:
304
356
  await manager.remove_thread_stream(thread_id, queue)
305
357
  raise
306
- if since and since < watermark:
307
- await manager.remove_thread_stream(thread_id, queue)
308
- return JSONResponse(
309
- {"code": "cursor_expired", "detail": "cursor_expired", "recovery": "thread_snapshot"},
310
- status_code=410,
311
- )
312
- resumable = await _resumable_run_ids(
313
- {UUID(wire["params"]["run_id"]) for wire in replay if wire.get("params", {}).get("run_id")}
314
- )
315
- replay = [
316
- wire
317
- for wire in replay
318
- if not wire.get("params", {}).get("run_id") or UUID(wire["params"]["run_id"]) in resumable
319
- ]
320
358
 
321
359
  async def stream() -> AsyncIterator[str]:
322
360
  metric_inc("graphharbor_protocol_connections_opened_total")
323
361
  if since:
324
362
  metric_inc("graphharbor_protocol_replays_total")
325
363
  seen: set[int] = set()
326
- try:
364
+ loop = asyncio.get_running_loop()
365
+ last_sent_at = loop.time()
366
+ last_auth_at = loop.time()
367
+
368
+ async def _check_authorized() -> bool:
369
+ nonlocal last_auth_at
370
+ now_ = loop.time()
371
+ if now_ - last_auth_at < 10.0:
372
+ return True
327
373
  if await _thread(request, thread_id) is None:
374
+ return False
375
+ last_auth_at = now_
376
+ return True
377
+
378
+ try:
379
+ if not await _check_authorized():
328
380
  return
329
381
  for wire in replay:
330
- if await _thread(request, thread_id) is None:
331
- return
332
382
  seq = wire.get("seq")
333
383
  if isinstance(seq, int) and seq not in seen and _wire_matches(wire, body):
334
384
  seen.add(seq)
335
385
  metric_inc("graphharbor_protocol_events_total")
336
386
  yield _frame(wire)
337
- started = asyncio.get_running_loop().time()
338
- while asyncio.get_running_loop().time() - started < timeout:
387
+ last_sent_at = loop.time()
388
+ started = loop.time()
389
+ while loop.time() - started < timeout:
390
+ now = loop.time()
391
+ remaining = max(0.1, heartbeat - (now - last_sent_at))
339
392
  try:
340
- message = await asyncio.wait_for(queue.get(), timeout=heartbeat)
393
+ message = await asyncio.wait_for(queue.get(), timeout=remaining)
341
394
  except TimeoutError:
342
395
  if await request.is_disconnected():
343
396
  return
344
- if await _thread(request, thread_id) is None:
397
+ if not await _check_authorized():
345
398
  return
346
399
  yield ": heartbeat\n\n"
400
+ last_sent_at = loop.time()
347
401
  continue
348
- if await _thread(request, thread_id) is None:
402
+ if not await _check_authorized():
349
403
  return
350
404
  try:
351
- wire = json.loads(message.data)
405
+ payload = json.loads(message.data)
352
406
  except (TypeError, ValueError, json.JSONDecodeError):
407
+ payload = None
408
+ if not isinstance(payload, dict):
409
+ if loop.time() - last_sent_at >= heartbeat:
410
+ yield ": heartbeat\n\n"
411
+ last_sent_at = loop.time()
353
412
  continue
354
- if not isinstance(wire, dict):
355
- continue
413
+ wire = payload
356
414
  seq = wire.get("seq")
357
- if not isinstance(seq, int) or seq <= since or seq in seen:
358
- continue
359
- if not _wire_matches(wire, body):
415
+ if (
416
+ not isinstance(seq, int)
417
+ or seq <= since
418
+ or seq in seen
419
+ or not _wire_matches(wire, body)
420
+ ):
421
+ if loop.time() - last_sent_at >= heartbeat:
422
+ yield ": heartbeat\n\n"
423
+ last_sent_at = loop.time()
360
424
  continue
361
425
  seen.add(seq)
362
426
  metric_inc("graphharbor_protocol_events_total")
363
427
  yield _frame(wire)
428
+ last_sent_at = loop.time()
364
429
  yield ": stream timeout\n\n"
365
430
  finally:
366
431
  metric_inc("graphharbor_protocol_connections_closed_total")
@@ -369,7 +434,11 @@ async def protocol_event_stream(request: Request) -> JSONResponse | StreamingRes
369
434
  return StreamingResponse(
370
435
  stream(),
371
436
  media_type="text/event-stream",
372
- headers={"Cache-Control": "no-store", "X-Accel-Buffering": "no"},
437
+ headers={
438
+ "Cache-Control": "no-cache, no-transform",
439
+ "Connection": "keep-alive",
440
+ "X-Accel-Buffering": "no",
441
+ },
373
442
  )
374
443
 
375
444
 
@@ -8,40 +8,29 @@ import os
8
8
  import pathlib
9
9
  import sys
10
10
  from contextlib import asynccontextmanager
11
- from datetime import UTC, datetime
11
+ from datetime import datetime
12
12
  from typing import Any
13
13
  from uuid import UUID
14
14
 
15
15
  import uvicorn
16
16
  from langgraph_cli.config import validate_config_file
17
- from sqlalchemy import func, select
18
17
  from starlette.applications import Starlette
19
18
  from starlette.middleware import Middleware
20
19
  from starlette.middleware.cors import CORSMiddleware
21
20
  from starlette.requests import Request
22
- from starlette.responses import HTMLResponse, JSONResponse, Response
21
+ from starlette.responses import HTMLResponse, JSONResponse
23
22
  from starlette.routing import Mount, Route
24
23
 
25
24
  from langgraph_runtime_pg.auth import (
26
25
  PrincipalMiddleware,
27
- in_principal_scope,
28
26
  principal_from_scope,
29
- scoped_idempotency_key,
30
27
  )
31
28
  from langgraph_runtime_pg.checkpoint import get_checkpointer
32
- from langgraph_runtime_pg.database import connect, pool_stats
29
+ from langgraph_runtime_pg.database import pool_stats
33
30
  from langgraph_runtime_pg.graph_registry import GraphRegistry, resolve_within_base_dir
34
31
  from langgraph_runtime_pg.metrics import prometheus_text, set_gauge
35
- from langgraph_runtime_pg.models import (
36
- AssistantRow,
37
- AssistantVersionRow,
38
- RunRow,
39
- ThreadRow,
40
- )
41
32
  from langgraph_runtime_pg.production import RuntimeReadiness, lifespan as runtime_lifespan
42
33
  from langgraph_runtime_pg.protocol import official_info_document
43
- from langgraph_runtime_pg.redis_stream import wake_run_queue
44
- from langgraph_runtime_pg.run_store import RunRepository
45
34
  from langhost.core_api import (
46
35
  assistants_count,
47
36
  assistants_create,
@@ -360,137 +349,6 @@ async def _metrics(_: Request):
360
349
  return PlainTextResponse(prometheus_text(), media_type="text/plain; version=0.0.4")
361
350
 
362
351
 
363
- def _no_content() -> Response:
364
- return Response(status_code=204)
365
-
366
-
367
- async def _capability_unavailable(request: Request) -> JSONResponse:
368
- capability = request.path_params.get("capability", "stream_v2")
369
- return JSONResponse(
370
- {
371
- "detail": f"capability {capability!r} is not enabled in the foundation profile",
372
- "capability": capability,
373
- "status": 501,
374
- },
375
- status_code=501,
376
- )
377
-
378
-
379
- def _scope_query(query: Any, model: Any, principal: Any) -> Any:
380
- return query
381
-
382
-
383
- def _metadata_query(query: Any, model: Any, metadata: Any) -> Any:
384
- if isinstance(metadata, dict) and metadata:
385
- query = query.where(model.metadata_.contains(metadata))
386
- return query
387
-
388
-
389
- def _request_limit_offset(request: Request) -> tuple[int, int]:
390
- try:
391
- limit = max(1, min(int(request.query_params.get("limit", "10")), 1000))
392
- offset = max(0, int(request.query_params.get("offset", "0")))
393
- except ValueError as exc:
394
- raise ValueError("limit and offset must be integers") from exc
395
- return limit, offset
396
-
397
-
398
- async def _assistant_search(request: Request) -> JSONResponse:
399
- principal = _principal(request)
400
- payload = await request.json()
401
- try:
402
- limit, offset = _request_limit_offset(request)
403
- except ValueError as exc:
404
- return JSONResponse({"detail": str(exc)}, status_code=422)
405
- query = select(AssistantRow).order_by(AssistantRow.created_at.desc())
406
- query = _scope_query(query, AssistantRow, principal)
407
- query = _metadata_query(query, AssistantRow, payload.get("metadata"))
408
- if payload.get("graph_id"):
409
- query = query.where(AssistantRow.graph_id == str(payload["graph_id"]))
410
- if payload.get("name"):
411
- query = query.where(AssistantRow.name.ilike(f"%{payload['name']}%"))
412
- async with connect() as conn:
413
- rows = (await conn.session.execute(query.limit(limit).offset(offset))).scalars().all()
414
- values = [_assistant_payload(row) for row in rows]
415
- if payload.get("response_format") == "object":
416
- return JSONResponse({"assistants": values, "next": None})
417
- return JSONResponse(values)
418
-
419
-
420
- async def _assistant_count(request: Request) -> JSONResponse:
421
- principal = _principal(request)
422
- payload = await request.json()
423
- query = select(func.count()).select_from(AssistantRow)
424
- query = _scope_query(query, AssistantRow, principal)
425
- query = _metadata_query(query, AssistantRow, payload.get("metadata"))
426
- if payload.get("graph_id"):
427
- query = query.where(AssistantRow.graph_id == str(payload["graph_id"]))
428
- if payload.get("name"):
429
- query = query.where(AssistantRow.name.ilike(f"%{payload['name']}%"))
430
- async with connect() as conn:
431
- count = int(await conn.session.scalar(query) or 0)
432
- return JSONResponse(count)
433
-
434
-
435
- async def _assistant_update(request: Request) -> JSONResponse:
436
- principal = _principal(request)
437
- try:
438
- assistant_id = UUID(request.path_params["assistant_id"])
439
- except ValueError:
440
- return JSONResponse({"detail": "assistant not found"}, status_code=404)
441
- payload = await request.json()
442
- async with connect() as conn:
443
- query = _scope_query(
444
- select(AssistantRow).where(AssistantRow.assistant_id == assistant_id),
445
- AssistantRow,
446
- principal,
447
- )
448
- row = (await conn.session.execute(query)).scalar_one_or_none()
449
- if row is None:
450
- return JSONResponse({"detail": "assistant not found"}, status_code=404)
451
- for field in ("graph_id", "name", "description", "config", "context"):
452
- if field in payload:
453
- setattr(row, field, payload[field])
454
- if isinstance(payload.get("metadata"), dict):
455
- row.metadata_ = {**row.metadata_, **payload["metadata"]}
456
- row.version += 1
457
- row.updated_at = datetime.now(UTC)
458
- conn.session.add(
459
- AssistantVersionRow(
460
- assistant_id=row.assistant_id,
461
- version=row.version,
462
- graph_id=row.graph_id,
463
- config=row.config,
464
- context=row.context,
465
- metadata_=row.metadata_,
466
- name=row.name,
467
- description=row.description,
468
- )
469
- )
470
- await conn.session.flush()
471
- return JSONResponse(_assistant_payload(row))
472
-
473
-
474
- async def _assistant_delete(request: Request) -> JSONResponse | Response:
475
- principal = _principal(request)
476
- try:
477
- assistant_id = UUID(request.path_params["assistant_id"])
478
- except ValueError:
479
- return _no_content()
480
- async with connect() as conn:
481
- query = _scope_query(
482
- select(AssistantRow).where(AssistantRow.assistant_id == assistant_id),
483
- AssistantRow,
484
- principal,
485
- )
486
- row = (await conn.session.execute(query)).scalar_one_or_none()
487
- if row is None:
488
- return _no_content()
489
- await conn.session.delete(row)
490
- await conn.session.flush()
491
- return _no_content()
492
-
493
-
494
352
  async def _assistants(request: Request) -> JSONResponse:
495
353
  request._json = dict(request.query_params)
496
354
  return await assistants_search(request)
@@ -501,152 +359,6 @@ async def _threads(request: Request) -> JSONResponse:
501
359
  return await threads_search(request)
502
360
 
503
361
 
504
- async def _assistant_get(request: Request) -> JSONResponse:
505
- try:
506
- assistant_id = UUID(request.path_params["assistant_id"])
507
- except ValueError:
508
- return JSONResponse({"detail": "assistant not found"}, status_code=404)
509
- async with connect() as conn:
510
- query = select(AssistantRow).where(AssistantRow.assistant_id == assistant_id)
511
- row = (await conn.session.execute(query)).scalar_one_or_none()
512
- if row is None:
513
- return JSONResponse({"detail": "assistant not found"}, status_code=404)
514
- return JSONResponse(_assistant_payload(row))
515
-
516
-
517
- async def _thread_get(request: Request) -> JSONResponse:
518
- principal = _principal(request)
519
- try:
520
- thread_id = UUID(request.path_params["thread_id"])
521
- except ValueError:
522
- return JSONResponse({"detail": "thread not found"}, status_code=404)
523
- async with connect() as conn:
524
- row = await conn.session.get(ThreadRow, thread_id)
525
- if row is None or not in_principal_scope(row, principal):
526
- return JSONResponse({"detail": "thread not found"}, status_code=404)
527
- return JSONResponse(_thread_payload(row))
528
-
529
-
530
- async def _resolve_assistant(
531
- session: Any, assistant_value: str, principal: Any
532
- ) -> AssistantRow | None:
533
- try:
534
- assistant_id = UUID(assistant_value)
535
- query = select(AssistantRow).where(AssistantRow.assistant_id == assistant_id)
536
- except ValueError:
537
- query = select(AssistantRow).where(AssistantRow.graph_id == assistant_value)
538
- return (await session.execute(query.limit(1))).scalar_one_or_none()
539
-
540
-
541
- async def _run_create(request: Request) -> JSONResponse:
542
- principal = _principal(request)
543
- payload = await request.json()
544
- assistant_value = str(payload.get("assistant_id", ""))
545
- thread_value = request.path_params.get("thread_id")
546
- if not assistant_value:
547
- return JSONResponse({"detail": "assistant_id is required"}, status_code=422)
548
- async with connect() as conn:
549
- assistant = await _resolve_assistant(conn.session, assistant_value, principal)
550
- if assistant is None:
551
- return JSONResponse({"detail": "assistant not found"}, status_code=404)
552
- thread = None
553
- thread_id = UUID(str(thread_value)) if thread_value else None
554
- if thread_id is not None:
555
- thread = await conn.session.get(ThreadRow, thread_id)
556
- if thread is None or not in_principal_scope(thread, principal):
557
- return JSONResponse({"detail": "thread not found"}, status_code=404)
558
- raw_idempotency_key = request.headers.get("idempotency-key") or payload.get(
559
- "idempotency_key"
560
- )
561
- idempotency_key = scoped_idempotency_key(principal, raw_idempotency_key)
562
- run = await RunRepository().create(
563
- conn.session,
564
- assistant_id=assistant.assistant_id,
565
- thread_id=thread_id,
566
- kwargs=payload,
567
- metadata=payload.get("metadata") or {},
568
- idempotency_key=idempotency_key,
569
- )
570
- await conn.session.refresh(run)
571
- conn.schedule_after_commit(wake_run_queue)
572
- return JSONResponse(_run_payload(run), status_code=201)
573
-
574
-
575
- async def _run_get(request: Request) -> JSONResponse:
576
- principal = _principal(request)
577
- run_id = UUID(request.path_params["run_id"])
578
- thread_id = UUID(request.path_params["thread_id"])
579
- async with connect() as conn:
580
- run = await conn.session.get(RunRow, run_id)
581
- if run is None or run.thread_id != thread_id or not in_principal_scope(run, principal):
582
- return JSONResponse({"detail": "run not found"}, status_code=404)
583
- return JSONResponse(_run_payload(run))
584
-
585
-
586
- async def _run_list(request: Request) -> JSONResponse:
587
- thread_id = UUID(request.path_params["thread_id"])
588
- async with connect() as conn:
589
- query = (
590
- select(RunRow).where(RunRow.thread_id == thread_id).order_by(RunRow.created_at.desc())
591
- )
592
- rows = (await conn.session.execute(query)).scalars().all()
593
- return JSONResponse([_run_payload(row) for row in rows])
594
-
595
-
596
- async def _run_cancel(request: Request) -> JSONResponse:
597
- return await runs_cancel(request)
598
-
599
-
600
- def _assistant_payload(row: AssistantRow) -> dict[str, Any]:
601
- return _plain(
602
- {
603
- "assistant_id": row.assistant_id,
604
- "graph_id": row.graph_id,
605
- "name": row.name,
606
- "description": row.description,
607
- "config": row.config,
608
- "context": row.context,
609
- "metadata": row.metadata_,
610
- "version": row.version,
611
- "created_at": row.created_at,
612
- "updated_at": row.updated_at,
613
- }
614
- )
615
-
616
-
617
- def _thread_payload(row: ThreadRow) -> dict[str, Any]:
618
- return _plain(
619
- {
620
- "thread_id": row.thread_id,
621
- "status": row.status,
622
- "metadata": row.metadata_,
623
- "config": row.config,
624
- "values": row.values_,
625
- "interrupts": row.interrupts,
626
- "error": row.error,
627
- "created_at": row.created_at,
628
- "updated_at": row.updated_at,
629
- "state_updated_at": row.state_updated_at,
630
- }
631
- )
632
-
633
-
634
- def _run_payload(row: RunRow) -> dict[str, Any]:
635
- return _plain(
636
- {
637
- "run_id": row.run_id,
638
- "thread_id": row.thread_id,
639
- "assistant_id": row.assistant_id,
640
- "status": row.status,
641
- "metadata": row.metadata_,
642
- "kwargs": row.kwargs,
643
- "multitask_strategy": row.multitask_strategy,
644
- "created_at": row.created_at,
645
- "updated_at": row.updated_at,
646
- }
647
- )
648
-
649
-
650
362
  def create_app(
651
363
  config: dict[str, Any] | Any,
652
364
  *,
@@ -31,13 +31,13 @@ def _namespace_error(namespace: Sequence[str]) -> Response | None:
31
31
  return None
32
32
 
33
33
 
34
- def _namespace(request: Request, value: Any) -> tuple[str, ...] | None:
34
+ def _namespace(value: Any, request: Request | None = None) -> tuple[str, ...] | None:
35
35
  if not isinstance(value, list) or not all(isinstance(label, str) for label in value):
36
36
  return None
37
37
  return tuple(value)
38
38
 
39
39
 
40
- def _public_namespace(request: Request, value: Sequence[str]) -> list[str]:
40
+ def _public_namespace(value: Sequence[str], request: Request | None = None) -> list[str]:
41
41
  return list(value)
42
42
 
43
43
 
@@ -61,7 +61,7 @@ async def _authorize_store(request: Request, action: str, value: dict[str, Any])
61
61
 
62
62
  def _item(request: Request, value: Any) -> dict[str, Any]:
63
63
  data = value.dict()
64
- data["namespace"] = _public_namespace(request, data["namespace"])
64
+ data["namespace"] = _public_namespace(data["namespace"])
65
65
  return data
66
66
 
67
67
 
@@ -77,10 +77,10 @@ async def store_put(request: Request) -> Response:
77
77
  payload = await _body(request)
78
78
  if payload is None or "key" not in payload or "value" not in payload:
79
79
  return JSONResponse({"detail": "namespace, key and value are required"}, status_code=422)
80
- namespace = _namespace(request, payload.get("namespace"))
80
+ namespace = _namespace(payload.get("namespace"))
81
81
  if namespace is None:
82
82
  return JSONResponse({"detail": "namespace must be an array of strings"}, status_code=422)
83
- if error := _namespace_error(namespace[-len(payload["namespace"]) :]):
83
+ if error := _namespace_error(namespace):
84
84
  return error
85
85
  if not isinstance(payload["key"], str) or not isinstance(payload["value"], dict):
86
86
  return JSONResponse(
@@ -110,8 +110,11 @@ async def store_put(request: Request) -> Response:
110
110
 
111
111
 
112
112
  async def store_get(request: Request) -> JSONResponse | Response:
113
- labels = request.query_params.get("namespace", "").split(".")
114
- namespace = _namespace(request, labels)
113
+ raw_namespace = request.query_params.get("namespace")
114
+ if not raw_namespace:
115
+ return JSONResponse({"detail": "namespace is required"}, status_code=422)
116
+ labels = raw_namespace.split(".")
117
+ namespace = _namespace(labels)
115
118
  if namespace is None:
116
119
  return JSONResponse({"detail": "namespace must be an array of strings"}, status_code=422)
117
120
  if error := _namespace_error(labels):
@@ -133,10 +136,10 @@ async def store_delete(request: Request) -> JSONResponse | Response:
133
136
  payload = await _body(request)
134
137
  if payload is None or "key" not in payload:
135
138
  return JSONResponse({"detail": "namespace and key are required"}, status_code=422)
136
- namespace = _namespace(request, payload.get("namespace"))
139
+ namespace = _namespace(payload.get("namespace"))
137
140
  if namespace is None:
138
141
  return JSONResponse({"detail": "namespace must be an array of strings"}, status_code=422)
139
- if error := _namespace_error(namespace[-len(payload["namespace"]) :]):
142
+ if error := _namespace_error(namespace):
140
143
  return error
141
144
  if not isinstance(payload["key"], str):
142
145
  return JSONResponse({"detail": "key must be a string"}, status_code=422)
@@ -150,7 +153,7 @@ async def store_search(request: Request) -> JSONResponse | Response:
150
153
  if payload is None:
151
154
  return JSONResponse({"detail": "request body must be an object"}, status_code=422)
152
155
  labels = payload.get("namespace_prefix")
153
- namespace = _namespace(request, labels)
156
+ namespace = _namespace(labels)
154
157
  if namespace is None:
155
158
  return JSONResponse(
156
159
  {"detail": "namespace_prefix must be an array of strings"}, status_code=422
@@ -222,7 +225,7 @@ async def store_list_namespaces(request: Request) -> JSONResponse | Response:
222
225
  limit=authorized.get("limit", 100),
223
226
  offset=authorized.get("offset", 0),
224
227
  )
225
- return JSONResponse({"namespaces": [_public_namespace(request, item) for item in namespaces]})
228
+ return JSONResponse({"namespaces": [_public_namespace(item) for item in namespaces]})
226
229
 
227
230
 
228
231
  __all__ = ["store_delete", "store_get", "store_list_namespaces", "store_put", "store_search"]
@@ -158,7 +158,9 @@ async def _resumable_run_ids(run_ids: set[UUID]) -> set[UUID]:
158
158
  )
159
159
 
160
160
 
161
- async def _thread_frame(row: RuntimeEventRow, modes: set[str]) -> tuple[str, Any, str] | None:
161
+ async def _thread_frame(
162
+ row: RuntimeEventRow, modes: set[str], *, attempts: dict[UUID, int] | None = None
163
+ ) -> tuple[str, Any, str] | None:
162
164
  event = row.payload
163
165
  name = str(event.get("event") or event.get("method") or "custom")
164
166
  if name == "lifecycle":
@@ -168,10 +170,13 @@ async def _thread_frame(row: RuntimeEventRow, modes: set[str]) -> tuple[str, Any
168
170
  return None
169
171
  attempt = 1
170
172
  if row.run_id is not None:
171
- async with connect() as conn:
172
- run = await conn.session.get(RunRow, row.run_id)
173
- if run is not None:
174
- attempt = max(run.retry_count, 1)
173
+ if attempts is not None and row.run_id in attempts:
174
+ attempt = attempts[row.run_id]
175
+ else:
176
+ async with connect() as conn:
177
+ run = await conn.session.get(RunRow, row.run_id)
178
+ if run is not None:
179
+ attempt = max(run.retry_count, 1)
175
180
  return "metadata", {"run_id": str(row.run_id), "attempt": attempt}, f"{row.sequence}-0"
176
181
  if status in _TERMINAL:
177
182
  if "lifecycle" not in modes and "run_modes" not in modes:
@@ -221,12 +226,26 @@ async def thread_stream(request: Request) -> JSONResponse | StreamingResponse:
221
226
  async def body() -> AsyncIterator[str]:
222
227
  nonlocal cursor_value
223
228
  queue = await manager.add_thread_stream(thread_id)
229
+ loop = asyncio.get_running_loop()
230
+ last_sent_at = loop.time()
231
+ last_auth_at = loop.time()
232
+
233
+ async def _check_authorized() -> bool:
234
+ nonlocal last_auth_at
235
+ now_ = loop.time()
236
+ if now_ - last_auth_at < 10.0:
237
+ return True
238
+ if (await _get_thread(request))[0] is None:
239
+ return False
240
+ last_auth_at = now_
241
+ return True
242
+
224
243
  try:
225
244
  if cursor_value < 0:
226
245
  cursor_value = await _thread_event_sequence(thread_id)
227
246
  initial_replay = True
228
247
  while True:
229
- if (await _get_thread(request))[0] is None:
248
+ if not await _check_authorized():
230
249
  return
231
250
  watermark, rows = await _thread_events(thread_id, cursor_value)
232
251
  if (
@@ -240,23 +259,50 @@ async def thread_stream(request: Request) -> JSONResponse | StreamingResponse:
240
259
  if initial_replay
241
260
  else set()
242
261
  )
262
+ running_run_ids = {
263
+ row.run_id
264
+ for row in rows
265
+ if row.run_id is not None
266
+ and (row.payload.get("event") or row.payload.get("method")) == "lifecycle"
267
+ and row.payload.get("status") == RunStatus.RUNNING.value
268
+ }
269
+ attempts: dict[UUID, int] = {}
270
+ if running_run_ids:
271
+ async with connect() as conn:
272
+ run_rows = (
273
+ await conn.session.scalars(
274
+ select(RunRow).where(RunRow.run_id.in_(running_run_ids))
275
+ )
276
+ ).all()
277
+ attempts = {r.run_id: max(r.retry_count, 1) for r in run_rows}
278
+ emitted = False
243
279
  for row in rows:
244
- if (await _get_thread(request))[0] is None:
245
- return
246
280
  cursor_value = row.sequence
247
281
  if initial_replay and row.run_id is not None and row.run_id not in resumable:
248
282
  continue
249
- frame = await _thread_frame(row, modes)
283
+ frame = await _thread_frame(row, modes, attempts=attempts)
250
284
  if frame is not None:
251
285
  name, data, event_id = frame
252
286
  yield _sse(name, data, event_id=event_id, event_id_last=True)
287
+ last_sent_at = loop.time()
288
+ emitted = True
253
289
  initial_replay = False
290
+
291
+ if not emitted and (loop.time() - last_sent_at >= heartbeat):
292
+ yield ": heartbeat\n\n"
293
+ last_sent_at = loop.time()
294
+
295
+ now = loop.time()
296
+ remaining = max(0.1, heartbeat - (now - last_sent_at))
254
297
  try:
255
- await asyncio.wait_for(queue.get(), timeout=heartbeat)
298
+ await asyncio.wait_for(queue.get(), timeout=remaining)
256
299
  except TimeoutError:
257
300
  if await request.is_disconnected():
258
301
  return
302
+ if not await _check_authorized():
303
+ return
259
304
  yield ": heartbeat\n\n"
305
+ last_sent_at = loop.time()
260
306
  finally:
261
307
  await manager.remove_thread_stream(thread_id, queue)
262
308
 
@@ -264,7 +310,7 @@ async def thread_stream(request: Request) -> JSONResponse | StreamingResponse:
264
310
  body(),
265
311
  media_type="text/event-stream",
266
312
  headers={
267
- "Cache-Control": "no-store",
313
+ "Cache-Control": "no-cache, no-transform",
268
314
  "Connection": "keep-alive",
269
315
  "X-Accel-Buffering": "no",
270
316
  },
@@ -458,39 +504,63 @@ async def _run_sse(
458
504
  metric_inc("graphharbor_sse_events_total", labels={"version": version})
459
505
  yield _sse(name, data, event_id=sequence if resumable else None)
460
506
 
461
- for envelope in replay:
507
+ loop = asyncio.get_running_loop()
508
+ last_sent_at = loop.time()
509
+ last_auth_at = loop.time()
510
+
511
+ async def _check_authorized() -> bool:
512
+ nonlocal last_auth_at
513
+ now_ = loop.time()
514
+ if now_ - last_auth_at < 10.0:
515
+ return True
462
516
  if await _run_snapshot(request, run_id) is None:
463
- return
517
+ return False
518
+ last_auth_at = now_
519
+ return True
520
+
521
+ for envelope in replay:
464
522
  async for frame in emit_envelope(envelope):
465
523
  yield frame
524
+ last_sent_at = loop.time()
466
525
  snapshot = await _run_snapshot(request, run_id)
467
526
  if snapshot is None or snapshot.status in _TERMINAL:
468
527
  return
469
528
 
470
- started = asyncio.get_running_loop().time()
471
- while asyncio.get_running_loop().time() - started < timeout:
529
+ started = loop.time()
530
+ while loop.time() - started < timeout:
531
+ now = loop.time()
532
+ remaining = max(0.1, heartbeat - (now - last_sent_at))
472
533
  try:
473
- message = await asyncio.wait_for(queue.get(), timeout=heartbeat)
534
+ message = await asyncio.wait_for(queue.get(), timeout=remaining)
474
535
  except TimeoutError:
475
- snapshot = await _run_snapshot(request, run_id)
476
- if snapshot is None:
536
+ if not await _check_authorized():
477
537
  return
478
538
  yield ": heartbeat\n\n"
479
- if snapshot.status in _TERMINAL:
539
+ last_sent_at = loop.time()
540
+ snapshot = await _run_snapshot(request, run_id)
541
+ if snapshot is not None and snapshot.status in _TERMINAL:
480
542
  for envelope in await _load_events(run_id, after=max(seen or {cursor})):
481
- if await _run_snapshot(request, run_id) is None:
482
- return
483
543
  async for frame in emit_envelope(envelope):
484
544
  yield frame
545
+ last_sent_at = loop.time()
485
546
  return
486
547
  continue
487
548
  live_envelope = _message_envelope(message)
488
549
  if live_envelope is None:
550
+ if loop.time() - last_sent_at >= heartbeat:
551
+ yield ": heartbeat\n\n"
552
+ last_sent_at = loop.time()
489
553
  continue
490
- if await _run_snapshot(request, run_id) is None:
554
+ if not await _check_authorized():
491
555
  return
556
+ emitted = False
492
557
  async for frame in emit_envelope(live_envelope):
493
558
  yield frame
559
+ last_sent_at = loop.time()
560
+ emitted = True
561
+ if not emitted and (loop.time() - last_sent_at >= heartbeat):
562
+ yield ": heartbeat\n\n"
563
+ last_sent_at = loop.time()
494
564
  event = live_envelope.get("event")
495
565
  if isinstance(event, dict) and event.get("event") == "lifecycle":
496
566
  status = str(event.get("status", ""))
@@ -502,7 +572,7 @@ async def _run_sse(
502
572
  await manager.remove_queue(run_id, thread_id, queue)
503
573
 
504
574
  headers = {
505
- "Cache-Control": "no-cache",
575
+ "Cache-Control": "no-cache, no-transform",
506
576
  "Connection": "keep-alive",
507
577
  "X-Accel-Buffering": "no",
508
578
  }
@@ -582,3 +582,35 @@ def test_filters_compile_as_bound_json_before_pagination():
582
582
  assert "jsonb_typeof" in sql and "@>" in sql and "::JSONB" in sql
583
583
  assert "alice" not in sql
584
584
  assert sql.index("WHERE") < sql.index("LIMIT")
585
+
586
+
587
+ @pytest.mark.asyncio
588
+ async def test_synchronous_auth_handler_is_supported():
589
+ auth = Auth()
590
+ auth._handlers[("threads", "read")] = [lambda ctx, value: {"owner": ctx.user.identity}]
591
+
592
+ user = {"identity": "alice", "permissions": []}
593
+ result = await authorize(auth, user, "threads", "read", {"thread_id": "test"})
594
+ assert result == {"owner": "alice"}
595
+
596
+
597
+ @pytest.mark.asyncio
598
+ async def test_store_get_requires_namespace():
599
+ import json
600
+
601
+ from starlette.requests import Request
602
+
603
+ from langhost.store_api import store_get
604
+
605
+ scope = {
606
+ "type": "http",
607
+ "method": "GET",
608
+ "path": "/store/items",
609
+ "query_string": b"key=test",
610
+ "headers": [],
611
+ }
612
+ request = Request(scope)
613
+ res = await store_get(request)
614
+ assert res.status_code == 422
615
+ assert json.loads(res.body.decode()) == {"detail": "namespace is required"}
616
+
@@ -0,0 +1,344 @@
1
+ """Tests for deterministic stream heartbeats and resilience under filtered traffic."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import contextlib
7
+ import json
8
+ from uuid import UUID, uuid4
9
+
10
+ import pytest
11
+ from starlette.requests import Request
12
+ from starlette.responses import StreamingResponse
13
+
14
+ from langgraph_runtime_pg.models import ThreadRow
15
+ from langgraph_runtime_pg.protocol import protocol_event
16
+ from langgraph_runtime_pg.redis_stream import Message
17
+ from langhost.protocol_api import protocol_event_stream
18
+ from langhost.streaming import thread_stream
19
+
20
+
21
+ class InMemoryStreamManager:
22
+ def __init__(self) -> None:
23
+ self.queues: dict[UUID, list[asyncio.Queue]] = {}
24
+ self.removed_count = 0
25
+
26
+ async def add_thread_stream(self, thread_id: UUID) -> asyncio.Queue:
27
+ q: asyncio.Queue = asyncio.Queue()
28
+ self.queues.setdefault(thread_id, []).append(q)
29
+ return q
30
+
31
+ async def remove_thread_stream(self, thread_id: UUID, queue: asyncio.Queue) -> None:
32
+ if thread_id in self.queues:
33
+ self.queues[thread_id] = [q for q in self.queues[thread_id] if q is not queue]
34
+ self.removed_count += 1
35
+
36
+ async def publish_thread_event(self, thread_id: UUID, message: Message) -> None:
37
+ for q in self.queues.get(thread_id, []):
38
+ await q.put(message)
39
+
40
+
41
+ def _make_scope_and_receive(
42
+ method: str,
43
+ path: str,
44
+ path_params: dict[str, str],
45
+ body_dict: dict | None = None,
46
+ headers: dict[str, str] | None = None,
47
+ ) -> tuple[dict, callable]:
48
+ raw_headers = []
49
+ if headers:
50
+ for k, v in headers.items():
51
+ raw_headers.append((k.lower().encode("latin1"), v.encode("latin1")))
52
+ scope = {
53
+ "type": "http",
54
+ "method": method,
55
+ "path": path,
56
+ "headers": raw_headers,
57
+ "path_params": path_params,
58
+ "query_string": b"",
59
+ }
60
+ body_bytes = json.dumps(body_dict).encode("utf-8") if body_dict is not None else b""
61
+
62
+ async def receive():
63
+ return {"type": "http.request", "body": body_bytes, "more_body": False}
64
+
65
+ return scope, receive
66
+
67
+
68
+ @pytest.mark.asyncio
69
+ async def test_protocol_event_stream_heartbeat_under_filtered_traffic(monkeypatch):
70
+ """Ensure protocol stream emits heartbeats even when inundated with filtered events."""
71
+ monkeypatch.setenv("GRAPHHARBOR_PROTOCOL_HEARTBEAT_SECONDS", "0.15")
72
+ monkeypatch.setenv("GRAPHHARBOR_PROTOCOL_TIMEOUT_SECONDS", "2.0")
73
+
74
+ thread_id = uuid4()
75
+ mock_thread = ThreadRow(
76
+ thread_id=thread_id,
77
+ status="idle",
78
+ metadata_={},
79
+ config={},
80
+ values_={},
81
+ interrupts={},
82
+ error=None,
83
+ )
84
+
85
+ manager = InMemoryStreamManager()
86
+
87
+ monkeypatch.setattr("langhost.protocol_api.get_stream_manager", lambda: manager)
88
+ monkeypatch.setattr(
89
+ "langhost.protocol_api._thread",
90
+ lambda *a, **k: asyncio.sleep(0, result=mock_thread),
91
+ )
92
+ monkeypatch.setattr(
93
+ "langhost.protocol_api._load_protocol_events",
94
+ lambda *a, **k: asyncio.sleep(0, result=(0, [])),
95
+ )
96
+
97
+ scope, receive = _make_scope_and_receive(
98
+ "POST",
99
+ f"/threads/{thread_id}/stream/events",
100
+ {"thread_id": str(thread_id)},
101
+ body_dict={"channels": ["input"], "since": 0},
102
+ )
103
+ req = Request(scope, receive)
104
+ res = await protocol_event_stream(req)
105
+ assert isinstance(res, StreamingResponse)
106
+
107
+ async def pump_filtered_events():
108
+ for seq in range(1, 20):
109
+ await asyncio.sleep(0.03)
110
+ wire = protocol_event(
111
+ event_id=f"evt-{seq}",
112
+ sequence=seq,
113
+ run_id=str(uuid4()),
114
+ thread_id=str(thread_id),
115
+ event={"event": "filtered_internal", "data": "ignore_me"},
116
+ )
117
+ wire["method"] = "debug:internal"
118
+ await manager.publish_thread_event(
119
+ thread_id,
120
+ Message(
121
+ topic=b"thread",
122
+ id=f"{seq}-0".encode("ascii"),
123
+ data=json.dumps(wire).encode("utf-8"),
124
+ ),
125
+ )
126
+
127
+ pump_task = asyncio.create_task(pump_filtered_events())
128
+ heartbeats_received = 0
129
+ try:
130
+ async for chunk in res.body_iterator:
131
+ if chunk.startswith(": heartbeat"):
132
+ heartbeats_received += 1
133
+ if heartbeats_received >= 2:
134
+ break
135
+ finally:
136
+ pump_task.cancel()
137
+ with contextlib.suppress(asyncio.CancelledError):
138
+ await pump_task
139
+ await res.body_iterator.aclose()
140
+
141
+ assert heartbeats_received >= 2, f"Expected >= 2 heartbeats, got {heartbeats_received}"
142
+ assert manager.removed_count >= 1, "Queue must be cleaned up in finally block"
143
+
144
+
145
+ @pytest.mark.asyncio
146
+ async def test_thread_stream_heartbeat_under_empty_queue(monkeypatch):
147
+ """Ensure thread_stream emits heartbeat when idle."""
148
+ monkeypatch.setenv("GRAPHHARBOR_THREAD_STREAM_HEARTBEAT_SECONDS", "0.15")
149
+
150
+ thread_id = uuid4()
151
+ mock_thread = ThreadRow(
152
+ thread_id=thread_id,
153
+ status="idle",
154
+ metadata_={},
155
+ config={},
156
+ values_={},
157
+ interrupts={},
158
+ error=None,
159
+ )
160
+
161
+ manager = InMemoryStreamManager()
162
+
163
+ monkeypatch.setattr("langhost.streaming.get_stream_manager", lambda: manager)
164
+ monkeypatch.setattr(
165
+ "langhost.core_api._get_thread",
166
+ lambda *a, **k: asyncio.sleep(0, result=(mock_thread, None, thread_id)),
167
+ )
168
+ monkeypatch.setattr(
169
+ "langhost.streaming._thread_event_sequence",
170
+ lambda *a, **k: asyncio.sleep(0, result=0),
171
+ )
172
+ monkeypatch.setattr(
173
+ "langhost.streaming._thread_events",
174
+ lambda *a, **k: asyncio.sleep(0, result=(0, [])),
175
+ )
176
+
177
+ scope, receive = _make_scope_and_receive(
178
+ "GET",
179
+ f"/threads/{thread_id}/stream",
180
+ {"thread_id": str(thread_id)},
181
+ )
182
+ req = Request(scope, receive)
183
+ res = await thread_stream(req)
184
+ assert isinstance(res, StreamingResponse)
185
+
186
+ heartbeats_received = 0
187
+ try:
188
+ async for chunk in res.body_iterator:
189
+ if chunk.startswith(": heartbeat"):
190
+ heartbeats_received += 1
191
+ if heartbeats_received >= 2:
192
+ break
193
+ finally:
194
+ await res.body_iterator.aclose()
195
+
196
+ assert heartbeats_received >= 2
197
+ assert manager.removed_count >= 1
198
+
199
+
200
+ @pytest.mark.asyncio
201
+ async def test_stream_response_headers_compliance(monkeypatch):
202
+ """Verify that SSE stream endpoints provide standardized headers."""
203
+ thread_id = uuid4()
204
+ mock_thread = ThreadRow(
205
+ thread_id=thread_id,
206
+ status="idle",
207
+ metadata_={},
208
+ config={},
209
+ values_={},
210
+ interrupts={},
211
+ error=None,
212
+ )
213
+
214
+ manager = InMemoryStreamManager()
215
+ monkeypatch.setattr("langhost.protocol_api.get_stream_manager", lambda: manager)
216
+ monkeypatch.setattr("langhost.streaming.get_stream_manager", lambda: manager)
217
+ monkeypatch.setattr(
218
+ "langhost.protocol_api._thread",
219
+ lambda *a, **k: asyncio.sleep(0, result=mock_thread),
220
+ )
221
+ monkeypatch.setattr(
222
+ "langhost.protocol_api._load_protocol_events",
223
+ lambda *a, **k: asyncio.sleep(0, result=(0, [])),
224
+ )
225
+ monkeypatch.setattr(
226
+ "langhost.core_api._get_thread",
227
+ lambda *a, **k: asyncio.sleep(0, result=(mock_thread, None, thread_id)),
228
+ )
229
+ monkeypatch.setattr(
230
+ "langhost.streaming._thread_event_sequence",
231
+ lambda *a, **k: asyncio.sleep(0, result=0),
232
+ )
233
+ monkeypatch.setattr(
234
+ "langhost.streaming._thread_events",
235
+ lambda *a, **k: asyncio.sleep(0, result=(0, [])),
236
+ )
237
+
238
+ # 1. Test protocol stream
239
+ scope, receive = _make_scope_and_receive(
240
+ "POST",
241
+ f"/threads/{thread_id}/stream/events",
242
+ {"thread_id": str(thread_id)},
243
+ body_dict={"channels": ["input"], "since": 0},
244
+ )
245
+ res_proto = await protocol_event_stream(Request(scope, receive))
246
+ assert res_proto.status_code == 200
247
+ assert res_proto.headers["connection"] == "keep-alive"
248
+ assert "no-cache" in res_proto.headers["cache-control"]
249
+ assert res_proto.headers["x-accel-buffering"] == "no"
250
+ await res_proto.body_iterator.aclose()
251
+
252
+ # 2. Test thread stream
253
+ scope_t, receive_t = _make_scope_and_receive(
254
+ "GET",
255
+ f"/threads/{thread_id}/stream",
256
+ {"thread_id": str(thread_id)},
257
+ )
258
+ res_thread = await thread_stream(Request(scope_t, receive_t))
259
+ assert res_thread.status_code == 200
260
+ assert res_thread.headers["connection"] == "keep-alive"
261
+ assert "no-cache" in res_thread.headers["cache-control"]
262
+ assert res_thread.headers["x-accel-buffering"] == "no"
263
+ await res_thread.body_iterator.aclose()
264
+
265
+
266
+ @pytest.mark.asyncio
267
+ async def test_zombie_interrupt_replay_filtered(monkeypatch):
268
+ """Ensure resolved/historical input.requested events are filtered from replay."""
269
+ thread_id = uuid4()
270
+ run_id = uuid4()
271
+
272
+ # Active interrupts on thread only contains 'int-active'
273
+ mock_thread = ThreadRow(
274
+ thread_id=thread_id,
275
+ status="interrupted",
276
+ metadata_={},
277
+ config={},
278
+ values_={},
279
+ interrupts={"int-active": {"id": "int-active", "value": "Approve step 2"}},
280
+ error=None,
281
+ )
282
+
283
+ manager = InMemoryStreamManager()
284
+ monkeypatch.setattr("langhost.protocol_api.get_stream_manager", lambda: manager)
285
+ monkeypatch.setattr(
286
+ "langhost.protocol_api._thread",
287
+ lambda *a, **k: asyncio.sleep(0, result=mock_thread),
288
+ )
289
+ monkeypatch.setattr(
290
+ "langhost.streaming._resumable_run_ids",
291
+ lambda run_ids: asyncio.sleep(0, result=set(run_ids)),
292
+ )
293
+
294
+ # 2 historical events: seq 1 is zombie (int-old), seq 2 is active (int-active)
295
+ wire_old = protocol_event(
296
+ event_id="evt-1",
297
+ sequence=1,
298
+ run_id=str(run_id),
299
+ thread_id=str(thread_id),
300
+ event={
301
+ "event": "input.requested",
302
+ "data": {"interrupt_id": "int-old", "value": "Old resolved request"},
303
+ },
304
+ )
305
+ wire_active = protocol_event(
306
+ event_id="evt-2",
307
+ sequence=2,
308
+ run_id=str(run_id),
309
+ thread_id=str(thread_id),
310
+ event={
311
+ "event": "input.requested",
312
+ "data": {"interrupt_id": "int-active", "value": "Active request"},
313
+ },
314
+ )
315
+
316
+ monkeypatch.setattr(
317
+ "langhost.protocol_api._load_protocol_events",
318
+ lambda *a, **k: asyncio.sleep(0, result=(0, [wire_old, wire_active])),
319
+ )
320
+
321
+ scope, receive = _make_scope_and_receive(
322
+ "POST",
323
+ f"/threads/{thread_id}/stream/events",
324
+ {"thread_id": str(thread_id)},
325
+ body_dict={"channels": ["input"], "since": 0},
326
+ )
327
+ req = Request(scope, receive)
328
+ res = await protocol_event_stream(req)
329
+ assert isinstance(res, StreamingResponse)
330
+
331
+ frames: list[str] = []
332
+ try:
333
+ async for chunk in res.body_iterator:
334
+ frames.append(chunk)
335
+ # Replay emits immediately, break after replay
336
+ if "int-active" in chunk or len(frames) >= 2:
337
+ break
338
+ finally:
339
+ await res.body_iterator.aclose()
340
+
341
+ all_output = "".join(frames)
342
+ assert "int-old" not in all_output, "Zombie interrupt 'int-old' must NOT be replayed to client!"
343
+ assert "int-active" in all_output, "Active interrupt 'int-active' must be replayed to client!"
344
+