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.
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/PKG-INFO +2 -2
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/pyproject.toml +2 -2
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/protocol_api.py +98 -29
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/server.py +3 -291
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/store_api.py +14 -11
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/streaming.py +93 -23
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_application_authorization.py +32 -0
- graphharbor-0.13.0.post35/tests/test_stream_heartbeat_resilience.py +344 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/.gitignore +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/LICENSE +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/README.md +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/__init__.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/__main__.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/cli.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/core_api.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/src/langhost/mcp_transport.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_cli.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_graph_discovery.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_mcp_transport.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_official_protocol_compare.py +0 -0
- {graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_server_paths.py +0 -0
- {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.
|
|
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.
|
|
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.
|
|
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.
|
|
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
|
-
|
|
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
|
-
|
|
338
|
-
|
|
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=
|
|
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
|
|
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
|
|
402
|
+
if not await _check_authorized():
|
|
349
403
|
return
|
|
350
404
|
try:
|
|
351
|
-
|
|
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
|
-
|
|
355
|
-
continue
|
|
413
|
+
wire = payload
|
|
356
414
|
seq = wire.get("seq")
|
|
357
|
-
if
|
|
358
|
-
|
|
359
|
-
|
|
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={
|
|
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
|
|
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
|
|
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
|
|
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(
|
|
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(
|
|
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(
|
|
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(
|
|
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
|
|
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
|
-
|
|
114
|
-
|
|
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(
|
|
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
|
|
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(
|
|
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(
|
|
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(
|
|
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
|
-
|
|
172
|
-
|
|
173
|
-
|
|
174
|
-
|
|
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
|
|
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=
|
|
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-
|
|
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
|
-
|
|
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 =
|
|
471
|
-
while
|
|
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=
|
|
534
|
+
message = await asyncio.wait_for(queue.get(), timeout=remaining)
|
|
474
535
|
except TimeoutError:
|
|
475
|
-
|
|
476
|
-
if snapshot is None:
|
|
536
|
+
if not await _check_authorized():
|
|
477
537
|
return
|
|
478
538
|
yield ": heartbeat\n\n"
|
|
479
|
-
|
|
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
|
|
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
|
}
|
{graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_application_authorization.py
RENAMED
|
@@ -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
|
+
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_official_protocol_compare.py
RENAMED
|
File without changes
|
|
File without changes
|
{graphharbor-0.13.0.post33 → graphharbor-0.13.0.post35}/tests/test_thread_state_projection.py
RENAMED
|
File without changes
|