mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-16 17:46:20 +00:00
* fix(tools): run tool assembly off-loop at async entry points get_available_tools() may block on MCP cache initialization while it is called on async agent-assembly paths (task_tool, durable batch execution), stalling the calling event loop for the full discovery duration. Dispatch the (unchanged, synchronous) assembly call to a worker thread via asyncio.to_thread at the two async entry points so the loop keeps processing requests, SSE frames, cancellations, and timers. Fixes #5172 * fix(tools): offload lead-agent assembly off-loop and pin with blocking-io anchors Review follow-up for #5224: - run_agent now dispatches agent_factory(...) through asyncio.to_thread, so lead-agent assembly (including both get_available_tools call sites in _assemble_lead_agent) runs off the event loop — the Gateway headline scenario from issue #5172. - _ensure_sync_invocable_tool takes a double-checked threading.Lock, making the in-place tool.func wrap on the shared tool singletons explicitly single-shot now that assembly can run concurrently on worker threads. - Add backend/tests/blocking_io/test_tool_assembly_offloop.py: blocking-probe anchors for task_tool and SubagentBatchService._execute_item under the strict Blockbuster gate, plus a meta-check proving the gate trips on the exact syscall class (ExtensionsConfig.from_file on the loop). Verified the anchor goes red when the offload is flattened back to a plain call. * fix(gateway): build checkpoint state accessor off-loop; anchor run_agent offload Review follow-up for #5224: - Add abuild_checkpoint_state_accessor (asyncio.to_thread around the unchanged sync builder) and switch every async call site to it: the stateless_wait route, thread_runs, both threads call sites, and the build_thread_checkpoint_state_accessor boundary. The agent-factory assembly re-enters get_available_tools() and may block on MCP cache initialization; repeat calls hit _state_accessor_graph_cache and only pay the thread hop. - Add a third blocking-io anchor driving the real run_agent with minimal RunManager/bridge stubs; the factory performs a real production blocking read (ExtensionsConfig.from_file()) and the test asserts assembly never runs on the main thread. Verified the anchor goes red when the run_agent offload is flattened back to a plain call. - Adapt the test_threads_router checkpoint-builder patch sites to the new async name. * refactor(tools): carry assembly offloads on a dedicated bounded pool Review follow-up for #5224: - Add utils/assembly_io.py: a dedicated ThreadPoolExecutor (default 8 workers, DEER_FLOW_ASSEMBLY_WORKERS-overridable, mirroring utils/file_io.py and tools/sync.py) with run_assembly(), which copies contextvars explicitly. A hung stdio MCP server parks its worker for the full MCP timeout; carrying assembly hops on the loop's default executor would let a few parked assemblies queue every other to_thread/run_in_executor(None, ...) caller behind them. - Switch all four offloads (run_agent, task_tool, batch _execute_item, abuild_checkpoint_state_accessor) to run_assembly(). - State the cold-path behavior in the accessor docstring: the graph cache validates factory identity, so non-identity-stable factories may duplicate lead-agent assembly across concurrent readers (MCP discovery stays process-wide single-flight); the pool bounds the duplicates. - Add a fourth blocking-io anchor driving build_thread_checkpoint_state_ accessor with a per-resolution fresh factory (always a cache miss) and the real production blocking read; enumerate all four offloads in the gate's module docstring. Verified the anchor goes red when abuild_checkpoint_state_accessor is flattened back to a plain call. * fix(subagents): revalidate batch item before launch; make assembly pool observable Review follow-up for #5224: - _execute_item() revalidates the durable state right after assembly and before executor.execute_async(): renew_item_lease() returns valid=False when cancel_batch() terminalized the item or the lease was lost while assembly was parked, and the launch is skipped (the canceller already finalized the item). Previously the launch was unconditional and the poll loop's cancellation checks only started after execution began. - Regression test driving the real SQLite repository: a blocking assembly probe parks _execute_item, cancel_batch() lands, and the launch is skipped with the item staying cancelled. Verified the test goes red when the revalidation is removed. - run_assembly() tracks pending assemblies and logs a throttled WARNING once the pending count exceeds the worker count, so assembly starvation (workers parked on a hung MCP server) is distinguishable from idle. - The run_agent blocking-io anchor now binds a sentinel extension snapshot via ctx.extensions and asserts the factory observed it through get_agent_build_extensions(), pinning run_assembly()'s ContextVar propagation. Verified red when ctx.run is dropped. - Document the assembly pool in backend/AGENTS.md. * fix(utils): decrement the assembly pending count on the pool thread The pending-assembly counter behind the starvation warning decremented from the asyncio future's done callback, which never fires once the submitting loop is closed while its worker is still running: the count ratcheted up permanently and eventually fired the starvation warning with no starvation behind it (reproduced at 97dc9bec by review). Decrement instead from the dispatched work item: run_assembly() wraps func so a finally drops the count under the pending lock on the pool thread, and the done callback is gone. Pin the counter with tests/test_assembly_io.py: a healthy call returns the count to zero, and an abandoned loop (stopped while the worker is parked) does not wedge it — the abandoned case goes red against the old done-callback decrement. * docs(utils): fix the pending-counter comment after the decrement move The comment still described the removed done-callback decrement, contradicting _work()'s own comment; state the actual mechanism (increment on the loop before dispatch, decrement from the dispatched work item's finally on a pool thread). * test(gateway): retarget checkpoint-accessor stubs to the services seam thread_runs and runs now call abuild_checkpoint_state_accessor, so the upstream wait-reader, regenerate-prepare, and idempotency tests must stub the sync builder where abuild resolves it (app.gateway.services); stubbing the removed router re-exports fails with AttributeError at setup. The async seam semantics are unchanged: run_assembly invokes the stubbed sync builder off-loop and propagates its return values and exceptions. Move the agent/tool assembly off-load note from backend/AGENTS.md to deerflow/utils/AGENTS.md (next to assembly_io.py) so the effective instruction chain for agents/middlewares no longer grows past the AG002 hard limit. * fix(runtime): serialize same-key accessor assembly and release queued-cancel slots Address the three review follow-ups on the assembly off-load: - assembly_io: a job cancelled while still queued never runs its work item, so the dispatched finally never fired and _pending_assemblies stayed elevated until a false starvation warning. Exactly-once cleanup now rides the concurrent future's cancelled() state — cancel() only succeeds before the executor starts the item, so cancelled() is true precisely when the finally will never run — plus a submit-failure release; the one-worker queued-cancellation case is pinned red/green. - services: overlapping cold readers sharing one cache key could both run full agent assembly. _state_accessor_graph now serializes per key through a thread-side KeyedLockTable (pool threads, no running loop) and re-validates factory/app-config identity under the lock, so the factory runs exactly once while identity changes still rebuild. Cache dict access is lock-guarded now that construction runs off-loop. - guidance inventory: register deerflow/utils/AGENTS.md in EXPECTED_GUIDANCE_PATHS so test_repository_has_the_approved_scoped_ guidance_shape matches the relocated assembly note (CI shard 4). * test(keyed-lock): pin KeyedLockTable reclamation and waiter bypass directly Thread-side counterparts of the async table's own tests: overlapping hold() calls serialize (a late arrival joins the live entry instead of creating a second lock that bypasses a queued waiter), the last check-in pops the entry, and many unique keys leave the registry empty. Both regressions verified red — popping unconditionally trips the late-arrival test, never reclaiming trips the many-keys test. --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
1878 lines
76 KiB
Python
1878 lines
76 KiB
Python
"""Runs endpoints — create, stream, wait, cancel.
|
|
|
|
Implements the LangGraph Platform runs API on top of
|
|
:class:`deerflow.agents.runs.RunManager` and
|
|
:class:`deerflow.agents.stream_bridge.StreamBridge`.
|
|
|
|
SSE format is aligned with the LangGraph Platform protocol so that
|
|
the ``useStream`` React hook from ``@langchain/langgraph-sdk/react``
|
|
works without modification.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import logging
|
|
import re
|
|
import uuid
|
|
from copy import deepcopy
|
|
from datetime import UTC, datetime
|
|
from typing import Annotated, Any, Literal
|
|
|
|
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request
|
|
from fastapi.responses import Response, StreamingResponse
|
|
from langchain_core.messages import BaseMessage
|
|
from pydantic import BaseModel, Field
|
|
from starlette.background import BackgroundTask
|
|
|
|
from app.gateway.artifact_archive import ArtifactArchiveError, ArtifactArchiveResult, build_artifact_archive
|
|
from app.gateway.authz import require_cancel_permission_if, require_permission
|
|
from app.gateway.checkpoint_lineage import (
|
|
CheckpointLineageError,
|
|
CheckpointParentMissingError,
|
|
checkpoint_configurable,
|
|
checkpoint_messages,
|
|
find_checkpoint_before_message,
|
|
find_checkpoint_before_message_chronologically,
|
|
is_duration_only_checkpoint,
|
|
)
|
|
from app.gateway.context_usage import build_context_usage
|
|
from app.gateway.deps import get_current_user, get_feedback_repo, get_run_event_store, get_run_manager, get_run_store, get_stream_bridge
|
|
from app.gateway.internal_auth import get_trusted_internal_owner_user_id
|
|
from app.gateway.pagination import trim_run_message_page
|
|
from app.gateway.run_models import RunCreateRequest
|
|
from app.gateway.services import abuild_checkpoint_state_accessor, build_thread_checkpoint_state_accessor, sse_consumer, start_run, wait_for_run_completion
|
|
from app.gateway.utils import sanitize_log_param
|
|
from deerflow.agents.middlewares.dynamic_context_middleware import strip_injected_user_message_id_suffix
|
|
from deerflow.authz.sandbox_authz import safe_app_config_async
|
|
from deerflow.config.paths import get_paths, make_safe_user_id
|
|
from deerflow.runtime import CancelOutcome, ConflictError, RunRecord, RunStatus, ThreadOperationKind, serialize_channel_values_for_api
|
|
from deerflow.runtime.runs.store.base import format_run_cursor_created_at, normalize_run_created_at_iso
|
|
from deerflow.runtime.secret_context import redact_config_secrets, redact_metadata_secrets
|
|
from deerflow.runtime.user_context import get_effective_user_id
|
|
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY, get_original_user_content_text, message_to_text
|
|
from deerflow.utils.thread_id import ThreadId
|
|
from deerflow.workspace_changes import get_workspace_changes_response
|
|
|
|
logger = logging.getLogger(__name__)
|
|
router = APIRouter(prefix="/api/threads", tags=["runs"])
|
|
REGENERATE_HISTORY_SCAN_LIMIT = 200
|
|
# Doubled to keep ~200 effective checkpoints when duration-only checkpoints
|
|
# (one per successful run in steady state) consume roughly half of history.
|
|
REGENERATE_HISTORY_RAW_SCAN_LIMIT = REGENERATE_HISTORY_SCAN_LIMIT * 2
|
|
THREAD_MESSAGE_PAGE_SCAN_BATCH = 201
|
|
_artifact_archive_slots = asyncio.Semaphore(4)
|
|
_MISSING_REGENERATE_BASE_DETAIL = "Could not find an addressable checkpoint before the target user message"
|
|
_UNSAFE_REGENERATE_LINEAGE_DETAIL = "Could not safely resolve the checkpoint before the target user message"
|
|
THREAD_MESSAGE_LEGACY_SCAN_BATCH = 201
|
|
|
|
|
|
IdempotencyKeyHeader = Annotated[
|
|
str | None,
|
|
Header(
|
|
alias="Idempotency-Key",
|
|
max_length=255,
|
|
description="Retry key for idempotent run admission within this thread",
|
|
),
|
|
]
|
|
|
|
|
|
def _scope_http_run_idempotency_key(request: Request, thread_id: str, key: str | None) -> str | None:
|
|
"""Namespace a caller key for the process-wide run idempotency index."""
|
|
if not isinstance(key, str):
|
|
return None
|
|
key = key.strip()
|
|
if not key:
|
|
raise HTTPException(status_code=422, detail="Idempotency-Key must not be blank")
|
|
owner_id = get_trusted_internal_owner_user_id(request)
|
|
if owner_id is None:
|
|
user = getattr(request.state, "user", None)
|
|
user_id = getattr(user, "id", None)
|
|
owner_id = str(user_id) if user_id is not None else get_effective_user_id()
|
|
digest = hashlib.sha256(f"{owner_id}\0{thread_id}\0{key}".encode()).hexdigest()
|
|
return f"http-run:{digest}"
|
|
|
|
|
|
async def _refresh_store_backed_run(run_mgr: Any, record: Any) -> Any:
|
|
"""Overlay durable status/error onto a hydrated store-only record."""
|
|
if not getattr(record, "store_only", False):
|
|
return record
|
|
store = getattr(run_mgr, "_store", None)
|
|
get = getattr(store, "get", None)
|
|
if get is None:
|
|
return record
|
|
try:
|
|
row = get(record.run_id)
|
|
if hasattr(row, "__await__"):
|
|
row = await row
|
|
except Exception:
|
|
logger.exception("Failed to refresh store-backed run %s", getattr(record, "run_id", None))
|
|
return record
|
|
if not isinstance(row, dict):
|
|
return record
|
|
raw_status = row.get("status")
|
|
if raw_status:
|
|
try:
|
|
record.status = RunStatus(raw_status)
|
|
except ValueError:
|
|
pass
|
|
if "error" in row:
|
|
record.error = row.get("error")
|
|
return record
|
|
|
|
|
|
def _is_duration_only_checkpoint(checkpoint_tuple: Any) -> bool:
|
|
return is_duration_only_checkpoint(checkpoint_tuple)
|
|
|
|
|
|
def compute_run_durations(runs) -> dict[str, int]:
|
|
"""Map run_id -> duration in seconds from run timestamps."""
|
|
from datetime import datetime
|
|
|
|
durations: dict[str, int] = {}
|
|
for r in runs:
|
|
if r.created_at and r.updated_at:
|
|
try:
|
|
created = datetime.fromisoformat(r.created_at.replace("Z", "+00:00"))
|
|
updated = datetime.fromisoformat(r.updated_at.replace("Z", "+00:00"))
|
|
# Note: updated_at - created_at represents the row's total lifetime,
|
|
# which can slightly overshoot the actual AI turn end if the row is mutated later.
|
|
durations[r.run_id] = int((updated - created).total_seconds())
|
|
except Exception:
|
|
logger.warning("Failed to parse timestamps for run %s", r.run_id, exc_info=True)
|
|
return durations
|
|
|
|
|
|
def stamp_turn_duration_on_last_ai(messages, run_durations: dict[str, int]) -> None:
|
|
"""Attach each run's elapsed seconds to that run's final visible AI message only.
|
|
|
|
``turn_duration`` is the run's wall-clock lifetime (``compute_run_durations``),
|
|
not model thinking time — it belongs to the run, not to individual messages.
|
|
Stamping every AI message made the UI repeat the same number once per
|
|
message and let tool-wait time read as thinking latency (#4152).
|
|
Middleware-caller messages (e.g. title generation) are skipped so the badge
|
|
lands on the assistant's actual final answer.
|
|
|
|
Accepts both message shapes that carry ``run_id``: event-store rows, which
|
|
wrap the message payload in a ``content`` dict, and flat serialized
|
|
checkpoint messages (``/history``), where the payload is the row itself.
|
|
|
|
The middleware skip is only effective on the event-store shape: checkpoint
|
|
messages replayed on ``/history`` never carry ``metadata.caller`` (they are
|
|
plain serialized LangChain messages), so this skip is inert there. That is
|
|
not a gap in practice — middleware writes (e.g. title generation) go to
|
|
thread metadata, not the ``messages`` channel, so no middleware message
|
|
reaches a checkpoint's ``messages`` list to begin with.
|
|
"""
|
|
stamped: set[str] = set()
|
|
for msg in reversed(messages):
|
|
rid = msg.get("run_id")
|
|
if not rid or rid in stamped or rid not in run_durations:
|
|
continue
|
|
content = msg.get("content")
|
|
payload = content if isinstance(content, dict) else msg
|
|
metadata = msg.get("metadata") or {}
|
|
is_middleware = str(metadata.get("caller", "")).startswith("middleware:")
|
|
if payload.get("type") == "ai" and not is_middleware:
|
|
payload.setdefault("additional_kwargs", {})["turn_duration"] = run_durations[rid]
|
|
stamped.add(rid)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Request / response models
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class RegeneratePrepareRequest(BaseModel):
|
|
message_id: str = Field(..., min_length=1, description="Assistant message id to regenerate")
|
|
|
|
|
|
class RegeneratePrepareResponse(BaseModel):
|
|
input: dict[str, Any]
|
|
checkpoint: dict[str, Any]
|
|
metadata: dict[str, Any]
|
|
target_run_id: str
|
|
|
|
|
|
class EditRegeneratePrepareRequest(BaseModel):
|
|
human_message_id: str = Field(..., min_length=1, description="Source human message id to edit and rerun")
|
|
replacement_text: str = Field(..., min_length=1, description="Replacement user-visible text")
|
|
|
|
|
|
class EditRegeneratePrepareResponse(RegeneratePrepareResponse):
|
|
replacement_human_message_id: str
|
|
source_message_ids: list[str]
|
|
|
|
|
|
class ThreadMessagesPageResponse(BaseModel):
|
|
data: list[dict[str, Any]]
|
|
has_more: bool
|
|
next_before_seq: int | None = None
|
|
|
|
|
|
class RunResponse(BaseModel):
|
|
run_id: str
|
|
thread_id: str
|
|
assistant_id: str | None = None
|
|
status: str
|
|
metadata: dict[str, Any] = Field(default_factory=dict)
|
|
kwargs: dict[str, Any] = Field(default_factory=dict)
|
|
multitask_strategy: str = "reject"
|
|
created_at: str = ""
|
|
updated_at: str = ""
|
|
total_input_tokens: int = 0
|
|
total_output_tokens: int = 0
|
|
total_tokens: int = 0
|
|
llm_call_count: int = 0
|
|
lead_agent_tokens: int = 0
|
|
subagent_tokens: int = 0
|
|
middleware_tokens: int = 0
|
|
message_count: int = 0
|
|
stop_reason: str | None = None
|
|
|
|
|
|
class ThreadRunsPageResponse(BaseModel):
|
|
data: list[RunResponse]
|
|
has_more: bool
|
|
next_before_created_at: str | None = None
|
|
next_before_run_id: str | None = None
|
|
|
|
|
|
class ArtifactArchiveManifestResponse(BaseModel):
|
|
file_count: int
|
|
|
|
|
|
class ThreadTokenUsageModelBreakdown(BaseModel):
|
|
tokens: int = 0
|
|
runs: int = Field(
|
|
default=0,
|
|
description="Number of runs in which this model appeared; counts are non-exclusive for runs that used multiple models.",
|
|
)
|
|
|
|
|
|
class ThreadTokenUsageCallerBreakdown(BaseModel):
|
|
lead_agent: int = 0
|
|
subagent: int = 0
|
|
middleware: int = 0
|
|
|
|
|
|
class ThreadContextUsage(BaseModel):
|
|
token_count: int = 0
|
|
max_context_tokens: int | None = None
|
|
percentage: float | None = None
|
|
|
|
|
|
class ThreadTokenUsageResponse(BaseModel):
|
|
thread_id: str
|
|
total_tokens: int = 0
|
|
total_input_tokens: int = 0
|
|
total_output_tokens: int = 0
|
|
total_runs: int = 0
|
|
by_model: dict[str, ThreadTokenUsageModelBreakdown] = Field(default_factory=dict)
|
|
by_caller: ThreadTokenUsageCallerBreakdown = Field(default_factory=ThreadTokenUsageCallerBreakdown)
|
|
context_usage: ThreadContextUsage | None = None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def require_cancel_permission_when_action(request: Request, action: str | None) -> None:
|
|
"""Conditionally require ``runs:cancel`` for cancel-then-stream requests.
|
|
|
|
``stream_existing_run`` is gated at ``runs:read`` so action-less stream
|
|
joins keep working with read-only credentials, but its ``action`` branch
|
|
cancels the run — a separate permission. A read-only PAT (or any read-only
|
|
credential) must not reach the cancel path, and decorators cannot express
|
|
query-parameter-conditional permissions, so the check lives here. See
|
|
``authz.require_cancel_permission_if`` — the shared primitive for every
|
|
request dimension that carries cancel capability.
|
|
"""
|
|
require_cancel_permission_if(request, action is not None)
|
|
|
|
|
|
def _cancel_conflict_detail(run_id: str, record: RunRecord) -> str:
|
|
if record.status in (RunStatus.pending, RunStatus.running):
|
|
return f"Run {run_id} is not active on this worker and cannot be cancelled"
|
|
return f"Run {run_id} is not cancellable (status: {record.status.value})"
|
|
|
|
|
|
def _compute_retry_after(lease_expires_at: str | None, grace_seconds: int) -> int | None:
|
|
"""Return seconds until the lease expires + grace, for ``Retry-After``.
|
|
|
|
Returns ``None`` when the lease is NULL or unparseable so the caller
|
|
can decide whether to send a generic 409 without the header.
|
|
|
|
The ``max(1, ...)`` floor means a lease just about to expire yields
|
|
``Retry-After: 1``. This is a lower bound, not a recommended poll
|
|
interval — clients that honour this header should apply minimum
|
|
backoff / jitter rather than retrying every second.
|
|
"""
|
|
if lease_expires_at is None:
|
|
return None
|
|
try:
|
|
dt = datetime.fromisoformat(lease_expires_at)
|
|
if dt.tzinfo is None:
|
|
dt = dt.replace(tzinfo=UTC)
|
|
except (ValueError, TypeError):
|
|
return None
|
|
remaining = (dt - datetime.now(UTC)).total_seconds() + grace_seconds
|
|
return max(1, int(remaining))
|
|
|
|
|
|
async def _raise_lease_valid_elsewhere(
|
|
run_id: str,
|
|
run_mgr, # RunManager (avoid import for testability)
|
|
record: RunRecord,
|
|
) -> None:
|
|
"""Re-fetch the lease and raise HTTP 409 + Retry-After.
|
|
|
|
``record.lease_expires_at`` may be stale (fetched at request start while
|
|
the owner renewed between our read and the conditional UPDATE). Re-read
|
|
from the store to get the fresh value so ``Retry-After`` is accurate.
|
|
"""
|
|
fresh = await run_mgr.get(run_id)
|
|
if fresh is not None:
|
|
record = fresh
|
|
retry_after = _compute_retry_after(record.lease_expires_at, run_mgr.grace_seconds)
|
|
headers: dict[str, str] = {}
|
|
if retry_after is not None:
|
|
headers["Retry-After"] = str(retry_after)
|
|
raise HTTPException(
|
|
status_code=409,
|
|
detail=f"Run {run_id} is active on another worker; retry after lease expiry.",
|
|
headers=headers,
|
|
)
|
|
|
|
|
|
def _record_to_response(record: RunRecord) -> RunResponse:
|
|
kwargs = dict(record.kwargs or {})
|
|
if "config" in kwargs:
|
|
kwargs["config"] = redact_config_secrets(kwargs["config"])
|
|
|
|
return RunResponse(
|
|
run_id=record.run_id,
|
|
thread_id=record.thread_id,
|
|
assistant_id=record.assistant_id,
|
|
status=record.status.value,
|
|
metadata=redact_metadata_secrets(record.metadata),
|
|
kwargs=kwargs,
|
|
multitask_strategy=record.multitask_strategy,
|
|
created_at=record.created_at,
|
|
updated_at=record.updated_at,
|
|
total_input_tokens=record.total_input_tokens,
|
|
total_output_tokens=record.total_output_tokens,
|
|
total_tokens=record.total_tokens,
|
|
llm_call_count=record.llm_call_count,
|
|
lead_agent_tokens=record.lead_agent_tokens,
|
|
subagent_tokens=record.subagent_tokens,
|
|
middleware_tokens=record.middleware_tokens,
|
|
message_count=record.message_count,
|
|
stop_reason=record.stop_reason,
|
|
)
|
|
|
|
|
|
def _message_id(message: Any) -> str | None:
|
|
value = getattr(message, "id", None)
|
|
if value is None and isinstance(message, dict):
|
|
value = message.get("id")
|
|
return str(value) if value else None
|
|
|
|
|
|
def _message_type(message: Any) -> str | None:
|
|
value = getattr(message, "type", None)
|
|
if value is None and isinstance(message, dict):
|
|
value = message.get("type") or message.get("role")
|
|
if value == "assistant":
|
|
return "ai"
|
|
return str(value) if value else None
|
|
|
|
|
|
def _message_name(message: Any) -> str | None:
|
|
value = getattr(message, "name", None)
|
|
if value is None and isinstance(message, dict):
|
|
value = message.get("name")
|
|
return str(value) if value else None
|
|
|
|
|
|
def _message_content(message: Any) -> Any:
|
|
if isinstance(message, dict):
|
|
return message.get("content")
|
|
return getattr(message, "content", None)
|
|
|
|
|
|
def _message_text(message: Any) -> str:
|
|
return message_to_text(message)
|
|
|
|
|
|
def _message_additional_kwargs(message: Any) -> dict[str, Any]:
|
|
value = getattr(message, "additional_kwargs", None)
|
|
if value is None and isinstance(message, dict):
|
|
value = message.get("additional_kwargs")
|
|
return dict(value or {}) if isinstance(value, dict) else {}
|
|
|
|
|
|
def _message_tool_calls(message: Any) -> list[Any]:
|
|
value = getattr(message, "tool_calls", None)
|
|
if value is None and isinstance(message, dict):
|
|
value = message.get("tool_calls")
|
|
if value is None:
|
|
value = _message_additional_kwargs(message).get("tool_calls")
|
|
return list(value) if isinstance(value, list) else []
|
|
|
|
|
|
def _is_hidden_or_control_message(message: Any) -> bool:
|
|
message_type = _message_type(message)
|
|
additional_kwargs = _message_additional_kwargs(message)
|
|
return message_type == "remove" or _message_name(message) == "summary" or additional_kwargs.get("hide_from_ui") is True
|
|
|
|
|
|
def _is_visible_human_message(message: Any) -> bool:
|
|
return _message_type(message) == "human" and not _is_hidden_or_control_message(message)
|
|
|
|
|
|
def _is_visible_ai_message(message: Any) -> bool:
|
|
return _message_type(message) == "ai" and not _is_hidden_or_control_message(message)
|
|
|
|
|
|
def _is_thread_history_hidden_message_row(row: dict[str, Any]) -> bool:
|
|
caller = str((row.get("metadata") or {}).get("caller", ""))
|
|
return caller.startswith("middleware:") or (caller.startswith("subagent:") and _message_type(row.get("content")) == "ai")
|
|
|
|
|
|
def _checkpoint_messages(snapshot: Any) -> list[Any]:
|
|
return checkpoint_messages(snapshot)
|
|
|
|
|
|
def _checkpoint_values(snapshot: Any) -> dict[str, Any]:
|
|
values = getattr(snapshot, "values", None)
|
|
return dict(values) if isinstance(values, dict) else {}
|
|
|
|
|
|
def _checkpoint_configurable(checkpoint_tuple: Any) -> dict[str, Any]:
|
|
return checkpoint_configurable(checkpoint_tuple)
|
|
|
|
|
|
def _checkpoint_response(checkpoint_tuple: Any) -> dict[str, Any]:
|
|
configurable = _checkpoint_configurable(checkpoint_tuple)
|
|
checkpoint_id = configurable.get("checkpoint_id")
|
|
if not checkpoint_id:
|
|
raise HTTPException(status_code=409, detail="Checkpoint is missing checkpoint_id")
|
|
return {
|
|
"checkpoint_ns": str(configurable.get("checkpoint_ns") or ""),
|
|
"checkpoint_id": str(checkpoint_id),
|
|
"checkpoint_map": configurable.get("checkpoint_map"),
|
|
}
|
|
|
|
|
|
def _clean_human_message_for_regenerate(message: Any) -> dict[str, Any]:
|
|
additional_kwargs = _message_additional_kwargs(message)
|
|
content = get_original_user_content_text(_message_content(message), additional_kwargs)
|
|
additional_kwargs.pop(ORIGINAL_USER_CONTENT_KEY, None)
|
|
additional_kwargs.pop("hide_from_ui", None)
|
|
|
|
clean_message: dict[str, Any] = {
|
|
"type": "human",
|
|
"content": [{"type": "text", "text": content}],
|
|
"additional_kwargs": additional_kwargs,
|
|
}
|
|
# Replay the id the client originally sent. The dynamic-context reminder
|
|
# re-keys the first user message of a thread to `{id}__user`, and replaying
|
|
# that persisted id into a state that has no reminder yet makes the
|
|
# middleware treat the turn as already injected, silently dropping the date
|
|
# and memory block the original turn had.
|
|
message_id = strip_injected_user_message_id_suffix(_message_id(message))
|
|
if message_id:
|
|
clean_message["id"] = message_id
|
|
name = _message_name(message)
|
|
if name:
|
|
clean_message["name"] = name
|
|
return clean_message
|
|
|
|
|
|
def _clean_human_message_for_edit(message: Any, *, replacement_id: str, replacement_text: str) -> dict[str, Any]:
|
|
source_kwargs = _message_additional_kwargs(message)
|
|
additional_kwargs: dict[str, Any] = {}
|
|
for key in ("files", "referenced_message_contexts"):
|
|
if key in source_kwargs:
|
|
additional_kwargs[key] = deepcopy(source_kwargs[key])
|
|
|
|
clean_message: dict[str, Any] = {
|
|
"type": "human",
|
|
"id": replacement_id,
|
|
"content": [{"type": "text", "text": replacement_text}],
|
|
"additional_kwargs": additional_kwargs,
|
|
}
|
|
name = _message_name(message)
|
|
if name:
|
|
clean_message["name"] = name
|
|
return clean_message
|
|
|
|
|
|
def _is_terminal_assistant_text_message(message: Any) -> bool:
|
|
return _is_visible_ai_message(message) and bool(_message_text(message).strip()) and not _message_tool_calls(message)
|
|
|
|
|
|
def _has_title(values: dict[str, Any]) -> bool:
|
|
title = values.get("title")
|
|
return isinstance(title, str) and bool(title)
|
|
|
|
|
|
def _has_active_goal(snapshot: Any) -> bool:
|
|
goal = _checkpoint_values(snapshot).get("goal")
|
|
return isinstance(goal, dict) and goal.get("status") == "active"
|
|
|
|
|
|
def _latest_editable_turn(messages: list[Any], human_message_id: str) -> tuple[int, Any, int, Any, list[str]]:
|
|
latest_human_index = next((index for index in range(len(messages) - 1, -1, -1) if _is_visible_human_message(messages[index])), None)
|
|
if latest_human_index is None or _message_id(messages[latest_human_index]) != human_message_id:
|
|
raise HTTPException(status_code=409, detail="Only the latest completed user turn can be edited")
|
|
|
|
source_human = messages[latest_human_index]
|
|
last_ai_index: int | None = None
|
|
for index, message in enumerate(messages[latest_human_index + 1 :], start=latest_human_index + 1):
|
|
if _is_visible_human_message(message):
|
|
break
|
|
if _is_visible_ai_message(message):
|
|
last_ai_index = index
|
|
|
|
if last_ai_index is None or not _is_terminal_assistant_text_message(messages[last_ai_index]):
|
|
raise HTTPException(status_code=409, detail="Only completed assistant text turns can be edited")
|
|
|
|
source_message_ids = [message_id for message in messages[latest_human_index : last_ai_index + 1] if (message_id := _message_id(message))]
|
|
return latest_human_index, source_human, last_ai_index, messages[last_ai_index], source_message_ids
|
|
|
|
|
|
def _event_message_id(row: dict[str, Any]) -> str | None:
|
|
content = row.get("content")
|
|
if isinstance(content, BaseMessage):
|
|
return _message_id(content)
|
|
if isinstance(content, dict):
|
|
return _message_id(content)
|
|
return None
|
|
|
|
|
|
def _run_last_ai_matches_message(record: RunRecord, message: Any) -> bool:
|
|
last_ai_message = (record.last_ai_message or "").strip()
|
|
if not last_ai_message:
|
|
return False
|
|
target_text = _message_text(message).strip()
|
|
if not target_text:
|
|
return False
|
|
return last_ai_message == target_text[: len(last_ai_message)]
|
|
|
|
|
|
async def _find_target_run_id(
|
|
thread_id: str,
|
|
message_id: str,
|
|
target_message: Any,
|
|
source_human: Any,
|
|
request: Request,
|
|
) -> str:
|
|
event_store = get_run_event_store(request)
|
|
rows = await event_store.list_messages(thread_id, limit=REGENERATE_HISTORY_SCAN_LIMIT)
|
|
for row in reversed(rows):
|
|
if row.get("event_type") not in {"ai_message", "llm.ai.response"}:
|
|
continue
|
|
if _event_message_id(row) == message_id:
|
|
run_id = row.get("run_id")
|
|
if isinstance(run_id, str) and run_id:
|
|
return run_id
|
|
|
|
source_run_id = _message_additional_kwargs(source_human).get("run_id")
|
|
if isinstance(source_run_id, str) and source_run_id:
|
|
return source_run_id
|
|
|
|
run_mgr = get_run_manager(request)
|
|
user_id = await get_current_user(request)
|
|
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=10)
|
|
fallback_record = next(
|
|
(record for record in records if record.status == RunStatus.success and _run_last_ai_matches_message(record, target_message)),
|
|
None,
|
|
)
|
|
if fallback_record is not None:
|
|
return fallback_record.run_id
|
|
if len(rows) >= REGENERATE_HISTORY_SCAN_LIMIT:
|
|
logger.warning(
|
|
"Could not find source run for regenerate message %s in recent run events for thread %s (limit=%s)",
|
|
message_id,
|
|
thread_id,
|
|
REGENERATE_HISTORY_SCAN_LIMIT,
|
|
)
|
|
raise HTTPException(status_code=409, detail="Could not find source run for assistant message")
|
|
|
|
|
|
async def _find_base_checkpoint_before_human(
|
|
thread_id: str,
|
|
human_message_id: str,
|
|
request: Request,
|
|
*,
|
|
head_checkpoint: Any | None = None,
|
|
) -> Any:
|
|
accessor, base_config = await build_thread_checkpoint_state_accessor(request, thread_id=thread_id)
|
|
if head_checkpoint is not None:
|
|
try:
|
|
return await find_checkpoint_before_message(
|
|
accessor,
|
|
head_checkpoint,
|
|
human_message_id,
|
|
max_depth=REGENERATE_HISTORY_RAW_SCAN_LIMIT,
|
|
)
|
|
except CheckpointParentMissingError:
|
|
# Old checkpoints and imported histories may not have parent links.
|
|
# Preserve the bounded chronological fallback for those records.
|
|
logger.debug(
|
|
"Could not resolve parent lineage for regenerate thread %s; falling back to history scan",
|
|
sanitize_log_param(thread_id),
|
|
exc_info=True,
|
|
)
|
|
except CheckpointLineageError as exc:
|
|
logger.warning(
|
|
"Rejected unsafe checkpoint lineage for regenerate thread %s",
|
|
sanitize_log_param(thread_id),
|
|
exc_info=True,
|
|
)
|
|
raise HTTPException(status_code=409, detail=_UNSAFE_REGENERATE_LINEAGE_DETAIL) from exc
|
|
try:
|
|
raw_checkpoints = await accessor.ahistory(base_config, limit=REGENERATE_HISTORY_RAW_SCAN_LIMIT)
|
|
checkpoints = [item for item in raw_checkpoints if not _is_duration_only_checkpoint(item)]
|
|
except Exception as exc:
|
|
logger.exception("Failed to list checkpoints for regenerate thread %s", thread_id)
|
|
raise HTTPException(status_code=500, detail="Failed to inspect checkpoint history") from exc
|
|
|
|
previous_checkpoint, target_found = find_checkpoint_before_message_chronologically(raw_checkpoints, human_message_id)
|
|
if target_found:
|
|
if previous_checkpoint is None:
|
|
raise HTTPException(
|
|
status_code=409,
|
|
detail=_MISSING_REGENERATE_BASE_DETAIL,
|
|
)
|
|
return previous_checkpoint
|
|
|
|
if len(checkpoints) >= REGENERATE_HISTORY_SCAN_LIMIT:
|
|
logger.warning(
|
|
"Could not locate target user message %s in recent checkpoint history for thread %s (limit=%s)",
|
|
human_message_id,
|
|
thread_id,
|
|
REGENERATE_HISTORY_SCAN_LIMIT,
|
|
)
|
|
raise HTTPException(
|
|
status_code=409,
|
|
detail=(f"Could not locate target user message in recent checkpoint history (limit={REGENERATE_HISTORY_SCAN_LIMIT})"),
|
|
)
|
|
|
|
|
|
def _run_status_value(record: Any) -> str | None:
|
|
status = getattr(record, "status", None)
|
|
if isinstance(status, RunStatus):
|
|
return status.value
|
|
return str(status) if status is not None else None
|
|
|
|
|
|
async def _require_successful_source_run(thread_id: str, run_id: str, request: Request) -> RunRecord:
|
|
run_mgr = get_run_manager(request)
|
|
user_id = await get_current_user(request)
|
|
record = await run_mgr.get(run_id, user_id=user_id)
|
|
if record is None:
|
|
# The run-event journal is the authoritative lookup above. This fallback
|
|
# only covers recent in-memory/store hydration gaps for the latest turn.
|
|
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=20)
|
|
record = next((candidate for candidate in records if getattr(candidate, "run_id", None) == run_id), None)
|
|
if record is None:
|
|
raise HTTPException(status_code=409, detail="Could not find source run for assistant message")
|
|
record_thread_id = getattr(record, "thread_id", None)
|
|
if isinstance(record_thread_id, str) and record_thread_id and record_thread_id != thread_id:
|
|
raise HTTPException(status_code=409, detail="Could not find source run for assistant message")
|
|
if _run_status_value(record) != RunStatus.success.value:
|
|
raise HTTPException(status_code=409, detail="Only successful assistant runs can be edited and rerun")
|
|
return record
|
|
|
|
|
|
async def _find_interrupted_target_run_id(
|
|
thread_id: str,
|
|
source_human: Any,
|
|
request: Request,
|
|
) -> str | None:
|
|
source_run_id = _message_additional_kwargs(source_human).get("run_id")
|
|
if not isinstance(source_run_id, str) or not source_run_id:
|
|
return None
|
|
|
|
run_mgr = get_run_manager(request)
|
|
user_id = await get_current_user(request)
|
|
record = await run_mgr.get(source_run_id, user_id=user_id)
|
|
if record is None:
|
|
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=20)
|
|
record = next(
|
|
(candidate for candidate in records if getattr(candidate, "run_id", None) == source_run_id),
|
|
None,
|
|
)
|
|
if record is None:
|
|
return None
|
|
if getattr(record, "thread_id", None) != thread_id:
|
|
return None
|
|
if _run_status_value(record) != RunStatus.interrupted.value:
|
|
return None
|
|
return source_run_id
|
|
|
|
|
|
async def _prepare_regenerate_payload(thread_id: str, message_id: str, request: Request) -> RegeneratePrepareResponse:
|
|
accessor, latest_config = await build_thread_checkpoint_state_accessor(request, thread_id=thread_id)
|
|
try:
|
|
latest_checkpoint = await accessor.aget(latest_config)
|
|
except Exception as exc:
|
|
logger.exception("Failed to read latest checkpoint for regenerate thread %s", thread_id)
|
|
raise HTTPException(status_code=500, detail="Failed to read latest checkpoint") from exc
|
|
latest_checkpoint_id = _checkpoint_configurable(latest_checkpoint).get("checkpoint_id")
|
|
if not latest_checkpoint_id:
|
|
raise HTTPException(status_code=404, detail=f"Thread {thread_id} has no checkpoint")
|
|
|
|
messages = _checkpoint_messages(latest_checkpoint)
|
|
target_index = next((i for i, message in enumerate(messages) if _message_id(message) == message_id), None)
|
|
if target_index is None:
|
|
# A response interrupted during an LLM call can be visible in the live
|
|
# stream without ever reaching a checkpoint. The server-stamped run ID
|
|
# on the latest user message is the durable link to that partial turn.
|
|
previous_human = next(
|
|
(message for message in reversed(messages) if _is_visible_human_message(message)),
|
|
None,
|
|
)
|
|
target_run_id = await _find_interrupted_target_run_id(thread_id, previous_human, request) if previous_human is not None else None
|
|
if target_run_id is None:
|
|
raise HTTPException(status_code=404, detail=f"Message {message_id} not found")
|
|
else:
|
|
target_message = messages[target_index]
|
|
if not _is_visible_ai_message(target_message):
|
|
raise HTTPException(status_code=409, detail="Only visible assistant messages can be regenerated")
|
|
|
|
latest_visible_ai = next((message for message in reversed(messages) if _is_visible_ai_message(message)), None)
|
|
if _message_id(latest_visible_ai) != message_id:
|
|
raise HTTPException(status_code=409, detail="Only the latest assistant message can be regenerated")
|
|
|
|
previous_human = next((message for message in reversed(messages[:target_index]) if _is_visible_human_message(message)), None)
|
|
target_run_id = (
|
|
await _find_target_run_id(
|
|
thread_id,
|
|
message_id,
|
|
target_message,
|
|
previous_human,
|
|
request,
|
|
)
|
|
if previous_human is not None
|
|
else None
|
|
)
|
|
if previous_human is None:
|
|
raise HTTPException(status_code=409, detail="Could not find the user message for this assistant response")
|
|
if target_run_id is None:
|
|
raise HTTPException(status_code=409, detail="Could not find source run for assistant message")
|
|
previous_human_id = _message_id(previous_human)
|
|
if not previous_human_id:
|
|
raise HTTPException(status_code=409, detail="The source user message is missing an id")
|
|
|
|
base_checkpoint_tuple = await _find_base_checkpoint_before_human(
|
|
thread_id,
|
|
previous_human_id,
|
|
request,
|
|
head_checkpoint=latest_checkpoint,
|
|
)
|
|
checkpoint = _checkpoint_response(base_checkpoint_tuple)
|
|
metadata = {
|
|
"regenerate_from_message_id": message_id,
|
|
"regenerate_from_run_id": target_run_id,
|
|
"regenerate_checkpoint_id": checkpoint["checkpoint_id"],
|
|
}
|
|
regenerate_input: dict[str, Any] = {"messages": [_clean_human_message_for_regenerate(previous_human)]}
|
|
latest_values = latest_checkpoint.values if isinstance(latest_checkpoint.values, dict) else {}
|
|
latest_title = latest_values.get("title")
|
|
if isinstance(latest_title, str) and latest_title:
|
|
# Regenerate resumes from the checkpoint before the target human turn.
|
|
# That checkpoint can predate a manual rename, so replay the current
|
|
# title as graph input instead of letting checkpoint rollback restore
|
|
# the older automatically generated title (#4457).
|
|
regenerate_input["title"] = latest_title
|
|
return RegeneratePrepareResponse(
|
|
input=regenerate_input,
|
|
checkpoint=checkpoint,
|
|
metadata=metadata,
|
|
target_run_id=target_run_id,
|
|
)
|
|
|
|
|
|
async def _prepare_edit_regenerate_payload(
|
|
thread_id: str,
|
|
human_message_id: str,
|
|
replacement_text: str,
|
|
request: Request,
|
|
) -> EditRegeneratePrepareResponse:
|
|
normalized_text = replacement_text.strip()
|
|
if not normalized_text:
|
|
raise HTTPException(status_code=409, detail="Edited message cannot be empty")
|
|
|
|
accessor, latest_config = await build_thread_checkpoint_state_accessor(request, thread_id=thread_id)
|
|
try:
|
|
latest_checkpoint = await accessor.aget(latest_config)
|
|
except Exception as exc:
|
|
logger.exception("Failed to read latest checkpoint for edit replay thread %s", thread_id)
|
|
raise HTTPException(status_code=500, detail="Failed to read latest checkpoint") from exc
|
|
latest_checkpoint_id = _checkpoint_configurable(latest_checkpoint).get("checkpoint_id")
|
|
if not latest_checkpoint_id:
|
|
raise HTTPException(status_code=404, detail=f"Thread {thread_id} has no checkpoint")
|
|
|
|
messages = _checkpoint_messages(latest_checkpoint)
|
|
if _has_active_goal(latest_checkpoint):
|
|
raise HTTPException(status_code=409, detail="Cannot edit while a goal is active")
|
|
|
|
_, source_human, _, source_ai, source_message_ids = _latest_editable_turn(messages, human_message_id)
|
|
source_text = get_original_user_content_text(_message_content(source_human), _message_additional_kwargs(source_human)).strip()
|
|
if normalized_text == source_text:
|
|
raise HTTPException(status_code=409, detail="Edited message is unchanged")
|
|
|
|
source_human_id = _message_id(source_human)
|
|
source_ai_id = _message_id(source_ai)
|
|
if not source_human_id:
|
|
raise HTTPException(status_code=409, detail="The source user message is missing an id")
|
|
if not source_ai_id:
|
|
raise HTTPException(status_code=409, detail="The source assistant message is missing an id")
|
|
|
|
base_checkpoint_tuple = await _find_base_checkpoint_before_human(
|
|
thread_id,
|
|
source_human_id,
|
|
request,
|
|
head_checkpoint=latest_checkpoint,
|
|
)
|
|
target_run_id = await _find_target_run_id(thread_id, source_ai_id, source_ai, source_human, request)
|
|
source_record = await _require_successful_source_run(thread_id, target_run_id, request)
|
|
checkpoint = _checkpoint_response(base_checkpoint_tuple)
|
|
replacement_human_message_id = str(uuid.uuid4())
|
|
source_metadata = getattr(source_record, "metadata", None) or {}
|
|
existing_group_id = source_metadata.get("edit_version_group_id") if isinstance(source_metadata, dict) else None
|
|
# Reserved for future edit-chain grouping across repeated edits of the same
|
|
# original prompt; current visibility still keys off regenerate_from_run_id.
|
|
edit_version_group_id = existing_group_id if isinstance(existing_group_id, str) and existing_group_id else source_human_id
|
|
metadata = {
|
|
"replay_kind": "edit",
|
|
"regenerate_from_message_id": source_ai_id,
|
|
"regenerate_from_run_id": target_run_id,
|
|
"regenerate_checkpoint_id": checkpoint["checkpoint_id"],
|
|
"edit_from_message_id": source_human_id,
|
|
"edit_message_id": replacement_human_message_id,
|
|
"edit_version_group_id": edit_version_group_id,
|
|
}
|
|
edit_input: dict[str, Any] = {
|
|
"messages": [
|
|
_clean_human_message_for_edit(
|
|
source_human,
|
|
replacement_id=replacement_human_message_id,
|
|
replacement_text=normalized_text,
|
|
)
|
|
]
|
|
}
|
|
base_values = base_checkpoint_tuple.values if isinstance(getattr(base_checkpoint_tuple, "values", None), dict) else {}
|
|
latest_values = latest_checkpoint.values if isinstance(latest_checkpoint.values, dict) else {}
|
|
latest_title = latest_values.get("title")
|
|
if _has_title(base_values) and isinstance(latest_title, str) and latest_title:
|
|
# The replay base can predate a manual rename, so replay the current
|
|
# title rather than letting checkpoint rollback restore the older one
|
|
# (#4457, the same rollback regenerate already guards against). An
|
|
# untitled base is deliberately left alone: it belongs to a thread the
|
|
# title middleware has not named yet, and pinning the current title
|
|
# there would keep a name generated from the prompt this edit replaced.
|
|
edit_input["title"] = latest_title
|
|
return EditRegeneratePrepareResponse(
|
|
input=edit_input,
|
|
checkpoint=checkpoint,
|
|
metadata=metadata,
|
|
target_run_id=target_run_id,
|
|
replacement_human_message_id=replacement_human_message_id,
|
|
source_message_ids=source_message_ids,
|
|
)
|
|
|
|
|
|
async def _default_history_hidden_run_ids(run_mgr: Any, thread_id: str, *, user_id: str | None) -> set[str]:
|
|
superseded_run_ids = await run_mgr.list_successful_regenerate_sources(thread_id, user_id=user_id)
|
|
edit_visibility = await run_mgr.list_edit_replay_visibility(thread_id, user_id=user_id)
|
|
return set(superseded_run_ids) | set(edit_visibility.hidden_source_run_ids) | set(edit_visibility.hidden_attempt_run_ids)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Endpoints
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@router.post("/{thread_id}/runs/regenerate/prepare", response_model=RegeneratePrepareResponse)
|
|
@require_permission("runs", "create", owner_check=True, require_existing=True)
|
|
async def prepare_regenerate_run(
|
|
thread_id: ThreadId,
|
|
body: RegeneratePrepareRequest,
|
|
request: Request,
|
|
) -> RegeneratePrepareResponse:
|
|
"""Prepare input and checkpoint for regenerating the latest assistant turn."""
|
|
return await _prepare_regenerate_payload(thread_id, body.message_id, request)
|
|
|
|
|
|
@router.post("/{thread_id}/runs/edit-regenerate/prepare", response_model=EditRegeneratePrepareResponse)
|
|
@require_permission("runs", "create", owner_check=True, require_existing=True)
|
|
async def prepare_edit_regenerate_run(
|
|
thread_id: ThreadId,
|
|
body: EditRegeneratePrepareRequest,
|
|
request: Request,
|
|
) -> EditRegeneratePrepareResponse:
|
|
"""Prepare input and checkpoint for editing then rerunning the latest user turn."""
|
|
return await _prepare_edit_regenerate_payload(thread_id, body.human_message_id, body.replacement_text, request)
|
|
|
|
|
|
@router.post("/{thread_id}/runs", response_model=RunResponse)
|
|
@require_permission("runs", "create", owner_check=True, require_existing=True)
|
|
async def create_run(
|
|
thread_id: ThreadId,
|
|
body: RunCreateRequest,
|
|
request: Request,
|
|
idempotency_key: IdempotencyKeyHeader = None,
|
|
) -> RunResponse:
|
|
"""Create a background run (returns immediately)."""
|
|
record = await start_run(
|
|
body,
|
|
thread_id,
|
|
request,
|
|
idempotency_key=_scope_http_run_idempotency_key(request, thread_id, idempotency_key),
|
|
)
|
|
return _record_to_response(record)
|
|
|
|
|
|
@router.post("/{thread_id}/runs/stream")
|
|
@require_permission("runs", "create", owner_check=True, require_existing=True)
|
|
async def stream_run(
|
|
thread_id: ThreadId,
|
|
body: RunCreateRequest,
|
|
request: Request,
|
|
idempotency_key: IdempotencyKeyHeader = None,
|
|
) -> StreamingResponse:
|
|
"""Create a run and stream events via SSE.
|
|
|
|
The response includes a ``Content-Location`` header with the run's
|
|
resource URL, matching the LangGraph Platform protocol. The
|
|
``useStream`` React hook uses this to extract run metadata.
|
|
"""
|
|
bridge = get_stream_bridge(request)
|
|
run_mgr = get_run_manager(request)
|
|
record = await start_run(
|
|
body,
|
|
thread_id,
|
|
request,
|
|
idempotency_key=_scope_http_run_idempotency_key(request, thread_id, idempotency_key),
|
|
)
|
|
|
|
# Same shape join already rejects: a reused store-only handle on a
|
|
# process-local bridge has no owner stream. Subscribing would create an
|
|
# empty log and wait forever. Terminal reuse still goes through
|
|
# sse_consumer with emit_gap_on_missing_stream so a missing stream emits
|
|
# gap rather than a bare end. First-time creates keep the default `end`.
|
|
if record.store_only and not bridge.supports_cross_process and record.status in (RunStatus.pending, RunStatus.running):
|
|
raise HTTPException(
|
|
status_code=409,
|
|
detail=f"Run {record.run_id} is not active on this worker and cannot be streamed",
|
|
)
|
|
|
|
return StreamingResponse(
|
|
sse_consumer(
|
|
bridge,
|
|
record,
|
|
request,
|
|
run_mgr,
|
|
emit_gap_on_missing_stream=record.idempotency_reused,
|
|
),
|
|
media_type="text/event-stream",
|
|
headers={
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
# LangGraph Platform includes run metadata in this header.
|
|
# The SDK uses a greedy regex to extract the run id from this path,
|
|
# so it must point at the canonical run resource without extra suffixes.
|
|
"Content-Location": f"/api/threads/{thread_id}/runs/{record.run_id}",
|
|
},
|
|
)
|
|
|
|
|
|
@router.post("/{thread_id}/runs/wait", response_model=dict)
|
|
@require_permission("runs", "create", owner_check=True, require_existing=True)
|
|
async def wait_run(
|
|
thread_id: ThreadId,
|
|
body: RunCreateRequest,
|
|
request: Request,
|
|
idempotency_key: IdempotencyKeyHeader = None,
|
|
) -> dict:
|
|
"""Create a run and block until it completes, returning the final state.
|
|
|
|
A reused in-flight run that this worker cannot observe returns the durable
|
|
status without blocking. A reused completed run also returns durable
|
|
status: the latest thread checkpoint may belong to a later run.
|
|
"""
|
|
bridge = get_stream_bridge(request)
|
|
run_mgr = get_run_manager(request)
|
|
record = await start_run(
|
|
body,
|
|
thread_id,
|
|
request,
|
|
idempotency_key=_scope_http_run_idempotency_key(request, thread_id, idempotency_key),
|
|
)
|
|
# Capture before waiting: create_or_reject mutates the shared cached
|
|
# record's idempotency_reused flag, so an overlapping retry must not
|
|
# change this request's checkpoint-vs-status decision.
|
|
reused = bool(getattr(record, "idempotency_reused", False))
|
|
|
|
# Reused/hydrated records have no local task. Wait on the bridge when this
|
|
# worker can observe it; otherwise return durable status rather than
|
|
# serializing whatever checkpoint happens to exist.
|
|
if getattr(record, "store_only", False) and not getattr(bridge, "supports_cross_process", False):
|
|
record = await _refresh_store_backed_run(run_mgr, record)
|
|
return {"status": record.status.value, "error": record.error}
|
|
|
|
if record.task is not None or getattr(record, "store_only", False):
|
|
completed = await wait_for_run_completion(bridge, record, request, run_mgr)
|
|
else:
|
|
completed = True
|
|
|
|
# Idempotent reuse is not bound to a run-specific checkpoint id. The latest
|
|
# thread head may be a later run, so do not claim it as this run's result.
|
|
if completed and not reused:
|
|
try:
|
|
accessor, config = await abuild_checkpoint_state_accessor(
|
|
request,
|
|
thread_id=thread_id,
|
|
assistant_id=body.assistant_id,
|
|
)
|
|
snapshot = await accessor.aget(config)
|
|
snapshot_config = snapshot.config or {}
|
|
if snapshot_config.get("configurable", {}).get("checkpoint_id"):
|
|
return serialize_channel_values_for_api(snapshot.values)
|
|
except Exception:
|
|
logger.exception("Failed to fetch final state for run %s", record.run_id)
|
|
|
|
if completed:
|
|
record = await _refresh_store_backed_run(run_mgr, record)
|
|
return {"status": record.status.value, "error": record.error}
|
|
|
|
|
|
def _parse_run_page_created_at(value: str) -> str:
|
|
try:
|
|
normalized = normalize_run_created_at_iso(value)
|
|
datetime.fromisoformat(normalized)
|
|
except ValueError:
|
|
raise HTTPException(status_code=422, detail="before_created_at must be an ISO-8601 timestamp") from None
|
|
return normalized
|
|
|
|
|
|
@router.get("/{thread_id}/runs", response_model=list[RunResponse])
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def list_runs(thread_id: ThreadId, request: Request) -> list[RunResponse]:
|
|
"""List the newest runs for a thread (default 100, as a bare array)."""
|
|
run_mgr = get_run_manager(request)
|
|
user_id = await get_current_user(request)
|
|
records = await run_mgr.list_by_thread(thread_id, user_id=user_id)
|
|
return [_record_to_response(r) for r in records]
|
|
|
|
|
|
@router.get("/{thread_id}/runs/page", response_model=ThreadRunsPageResponse)
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def list_runs_page(
|
|
thread_id: ThreadId,
|
|
request: Request,
|
|
limit: int = Query(default=50, ge=1, le=200),
|
|
before_created_at: str | None = Query(default=None),
|
|
before_run_id: str | None = Query(default=None, min_length=1),
|
|
) -> ThreadRunsPageResponse:
|
|
"""Return a newest-first keyset page of runs for a thread.
|
|
|
|
Response: { data: [...], has_more: bool, next_before_created_at, next_before_run_id }
|
|
Pass both cursor fields from the previous page's last row to continue.
|
|
"""
|
|
if (before_created_at is None) != (before_run_id is None):
|
|
raise HTTPException(
|
|
status_code=422,
|
|
detail="before_created_at and before_run_id must be provided together",
|
|
)
|
|
if before_created_at is not None:
|
|
before_created_at = _parse_run_page_created_at(before_created_at)
|
|
|
|
run_mgr = get_run_manager(request)
|
|
user_id = await get_current_user(request)
|
|
records = await run_mgr.list_by_thread(
|
|
thread_id,
|
|
user_id=user_id,
|
|
limit=limit + 1,
|
|
before_created_at=before_created_at,
|
|
before_run_id=before_run_id,
|
|
)
|
|
has_more = len(records) > limit
|
|
page = records[:limit]
|
|
last = page[-1] if page and has_more else None
|
|
return ThreadRunsPageResponse(
|
|
data=[_record_to_response(record) for record in page],
|
|
has_more=has_more,
|
|
next_before_created_at=format_run_cursor_created_at(last.created_at) if last else None,
|
|
next_before_run_id=last.run_id if last else None,
|
|
)
|
|
|
|
|
|
@router.get("/{thread_id}/runs/{run_id}", response_model=RunResponse)
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def get_run(thread_id: ThreadId, run_id: str, request: Request) -> RunResponse:
|
|
"""Get details of a specific run."""
|
|
run_mgr = get_run_manager(request)
|
|
user_id = await get_current_user(request)
|
|
record = await run_mgr.get(run_id, user_id=user_id)
|
|
if record is None or record.thread_id != thread_id:
|
|
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
|
return _record_to_response(record)
|
|
|
|
|
|
@router.post("/{thread_id}/runs/{run_id}/cancel")
|
|
@require_permission("runs", "cancel", owner_check=True, require_existing=True)
|
|
async def cancel_run(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
wait: bool = Query(default=False, description="Block until run completes after cancel"),
|
|
action: Literal["interrupt", "rollback"] = Query(default="interrupt", description="Cancel action"),
|
|
) -> Response:
|
|
"""Cancel a running or pending run.
|
|
|
|
- action=interrupt: Stop execution, keep current checkpoint (can be resumed)
|
|
- action=rollback: Stop execution, revert to pre-run checkpoint state
|
|
- wait=true: Block until the run fully stops, return 204
|
|
- wait=false: Return immediately with 202
|
|
|
|
In multi-worker deployments, a cancel landing on a non-owning worker
|
|
durably notifies the owner when its lease is live, or takes over and
|
|
terminalizes the run when that lease has expired.
|
|
"""
|
|
run_mgr = get_run_manager(request)
|
|
record = await run_mgr.get(run_id)
|
|
if record is None or record.thread_id != thread_id:
|
|
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
|
|
|
outcome = await run_mgr.cancel(run_id, action=action)
|
|
|
|
# Success paths — the run was cancelled locally, durably requested from
|
|
# a live owner, or taken over from a dead worker.
|
|
if outcome in (
|
|
CancelOutcome.cancelled,
|
|
CancelOutcome.requested,
|
|
CancelOutcome.taken_over,
|
|
):
|
|
if wait and record.task is not None:
|
|
try:
|
|
await record.task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
return Response(status_code=204)
|
|
if wait and outcome == CancelOutcome.requested:
|
|
bridge = get_stream_bridge(request)
|
|
if record.store_only and bridge.supports_cross_process:
|
|
completed = await wait_for_run_completion(
|
|
bridge,
|
|
record,
|
|
request,
|
|
run_mgr,
|
|
)
|
|
if completed:
|
|
return Response(status_code=204)
|
|
return Response(status_code=202)
|
|
|
|
if outcome == CancelOutcome.lease_valid_elsewhere:
|
|
await _raise_lease_valid_elsewhere(run_id, run_mgr, record)
|
|
|
|
# not_cancellable, not_active_locally, unknown
|
|
raise HTTPException(status_code=409, detail=_cancel_conflict_detail(run_id, record))
|
|
|
|
|
|
@router.get("/{thread_id}/runs/{run_id}/join")
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def join_run(thread_id: ThreadId, run_id: str, request: Request) -> StreamingResponse:
|
|
"""Join an existing run's SSE stream."""
|
|
run_mgr = get_run_manager(request)
|
|
record = await run_mgr.get(run_id)
|
|
if record is None or record.thread_id != thread_id:
|
|
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
|
bridge = get_stream_bridge(request)
|
|
if record.store_only and not bridge.supports_cross_process:
|
|
raise HTTPException(status_code=409, detail=f"Run {run_id} is not active on this worker and cannot be streamed")
|
|
|
|
return StreamingResponse(
|
|
# Joins are read-only observation: the creator's cancel-on-disconnect
|
|
# policy must not fire because an observer closed their connection.
|
|
sse_consumer(bridge, record, request, run_mgr, apply_on_disconnect=False),
|
|
media_type="text/event-stream",
|
|
headers={
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
},
|
|
)
|
|
|
|
|
|
def _reject_get_stream_action(
|
|
action: Literal["interrupt", "rollback"] | None = Query(default=None, include_in_schema=False),
|
|
) -> None:
|
|
"""Keep the GET join read-only before thread ownership or run lookup."""
|
|
if action is not None:
|
|
# SameSite=Lax still sends the session cookie on a cross-site top-level
|
|
# safe navigation. Reject the state-changing action before the endpoint
|
|
# wrapper performs its thread ownership lookup.
|
|
raise HTTPException(
|
|
status_code=405,
|
|
detail="`action` is only supported on POST requests",
|
|
headers={"Allow": "POST"},
|
|
)
|
|
|
|
|
|
async def _stream_existing_run(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
*,
|
|
action: Literal["interrupt", "rollback"] | None,
|
|
wait: int,
|
|
) -> Response:
|
|
"""Join an existing run's SSE stream, optionally cancelling it first.
|
|
|
|
When ``action=interrupt`` or ``action=rollback`` is present the run is
|
|
cancelled first; the response then streams any remaining buffered events
|
|
so the client observes a clean shutdown.
|
|
"""
|
|
require_cancel_permission_when_action(request, action)
|
|
|
|
run_mgr = get_run_manager(request)
|
|
record = await run_mgr.get(run_id)
|
|
if record is None or record.thread_id != thread_id:
|
|
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
|
bridge = get_stream_bridge(request)
|
|
if record.store_only and action is None and not bridge.supports_cross_process:
|
|
raise HTTPException(status_code=409, detail=f"Run {run_id} is not active on this worker and cannot be streamed")
|
|
|
|
# Cancel if an action was requested (stop-button / interrupt flow)
|
|
if action is not None:
|
|
outcome = await run_mgr.cancel(run_id, action=action)
|
|
if outcome == CancelOutcome.taken_over:
|
|
# The run was on another worker and is now marked ``error`` in the
|
|
# store. There is no local stream to drain — return immediately so
|
|
# the client doesn't hang on an SSE subscription this worker can
|
|
# never serve.
|
|
return Response(status_code=202)
|
|
if outcome not in (CancelOutcome.cancelled, CancelOutcome.requested):
|
|
if outcome == CancelOutcome.lease_valid_elsewhere:
|
|
await _raise_lease_valid_elsewhere(run_id, run_mgr, record)
|
|
raise HTTPException(status_code=409, detail=_cancel_conflict_detail(run_id, record))
|
|
if outcome == CancelOutcome.requested and record.store_only and not bridge.supports_cross_process:
|
|
# The request is durable, but this bridge cannot observe the
|
|
# owner's stream. Returning 202 is safer than hanging forever on
|
|
# a process-local subscription.
|
|
return Response(status_code=202)
|
|
if wait and record.task is not None:
|
|
try:
|
|
await record.task
|
|
except (asyncio.CancelledError, Exception):
|
|
pass
|
|
return Response(status_code=204)
|
|
if wait and outcome == CancelOutcome.requested:
|
|
completed = await wait_for_run_completion(
|
|
bridge,
|
|
record,
|
|
request,
|
|
run_mgr,
|
|
)
|
|
return Response(status_code=204 if completed else 202)
|
|
|
|
return StreamingResponse(
|
|
# Both methods of this handler are join surfaces: a POST carrying an
|
|
# action cancels explicitly above (already gated by
|
|
# require_cancel_permission_when_action), and an action-less join is
|
|
# read-only observation — the creator's cancel-on-disconnect policy
|
|
# must not fire because a joiner closed their connection.
|
|
sse_consumer(bridge, record, request, run_mgr, apply_on_disconnect=False),
|
|
media_type="text/event-stream",
|
|
headers={
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
},
|
|
)
|
|
|
|
|
|
# Register POST before GET to preserve the historical route precedence and
|
|
# Allow header, while separate signatures keep cancel-only parameters off the
|
|
# GET schema. The shared route name keeps generated operationIds stable.
|
|
@router.post("/{thread_id}/runs/{run_id}/stream", response_model=None, name="stream_existing_run")
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def stream_existing_run(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
action: Literal["interrupt", "rollback"] | None = Query(default=None, description="Cancel action"),
|
|
wait: int = Query(default=0, description="Block until cancelled (1) or return immediately (0)"),
|
|
) -> Response:
|
|
"""Join an existing run's SSE stream, optionally cancelling it first."""
|
|
return await _stream_existing_run(thread_id, run_id, request, action=action, wait=wait)
|
|
|
|
|
|
@router.get(
|
|
"/{thread_id}/runs/{run_id}/stream",
|
|
response_model=None,
|
|
dependencies=[Depends(_reject_get_stream_action)],
|
|
name="stream_existing_run",
|
|
)
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def join_existing_run_stream(thread_id: ThreadId, run_id: str, request: Request) -> Response:
|
|
"""Join an existing run's observation-only SSE stream."""
|
|
return await _stream_existing_run(thread_id, run_id, request, action=None, wait=0)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Messages / Events / Token usage endpoints
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@router.get("/{thread_id}/messages")
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def list_thread_messages(
|
|
thread_id: ThreadId,
|
|
request: Request,
|
|
limit: int = Query(default=50, ge=1, le=200),
|
|
before_seq: int | None = Query(default=None, ge=1),
|
|
after_seq: int | None = Query(default=None, ge=1),
|
|
) -> list[dict]:
|
|
"""Return displayable messages for a thread (across all runs), with feedback attached."""
|
|
# Resolve the caller once; it is needed both to scope the feedback query
|
|
# below and to list the thread's runs for turn-duration injection.
|
|
user_id = await get_current_user(request)
|
|
run_mgr = get_run_manager(request)
|
|
hidden_run_ids = await _default_history_hidden_run_ids(run_mgr, thread_id, user_id=user_id)
|
|
messages, _ = await _scan_visible_thread_messages(
|
|
thread_id,
|
|
limit=limit,
|
|
before_seq=before_seq,
|
|
after_seq=after_seq,
|
|
request=request,
|
|
user_id=user_id,
|
|
hidden_run_ids=hidden_run_ids,
|
|
include_middleware=True,
|
|
include_extra=False,
|
|
batch_size=THREAD_MESSAGE_LEGACY_SCAN_BATCH,
|
|
)
|
|
|
|
# Find the last AI message per run_id. AI messages are persisted by
|
|
# RunJournal with event_type "llm.ai.response" (see runtime/journal.py);
|
|
# the event store returns that value verbatim, so match on it here.
|
|
last_ai_per_run: dict[str, int] = {} # run_id -> index in messages list
|
|
for i, msg in enumerate(messages):
|
|
if msg.get("event_type") == "llm.ai.response":
|
|
last_ai_per_run[msg["run_id"]] = i
|
|
|
|
# Attach feedback to the last AI message of each run. Only query when there
|
|
# is an AI message to attach it to — threads with no completed AI turn yet
|
|
# would otherwise pay for a grouped feedback lookup whose result is unused.
|
|
feedback_map: dict[str, dict] = {}
|
|
if last_ai_per_run:
|
|
feedback_repo = get_feedback_repo(request)
|
|
feedback_map = await feedback_repo.list_by_thread_grouped(thread_id, user_id=user_id)
|
|
|
|
last_ai_indices = set(last_ai_per_run.values())
|
|
for i, msg in enumerate(messages):
|
|
if i in last_ai_indices:
|
|
run_id = msg["run_id"]
|
|
fb = feedback_map.get(run_id)
|
|
msg["feedback"] = (
|
|
{
|
|
"feedback_id": fb["feedback_id"],
|
|
"rating": fb["rating"],
|
|
"comment": fb.get("comment"),
|
|
}
|
|
if fb
|
|
else None
|
|
)
|
|
else:
|
|
msg["feedback"] = None
|
|
|
|
runs = await run_mgr.list_by_thread(thread_id, user_id=user_id)
|
|
run_durations = compute_run_durations(runs)
|
|
|
|
if run_durations:
|
|
stamp_turn_duration_on_last_ai(messages, run_durations)
|
|
|
|
return messages
|
|
|
|
|
|
async def _scan_visible_thread_messages(
|
|
thread_id: str,
|
|
*,
|
|
limit: int,
|
|
before_seq: int | None,
|
|
after_seq: int | None,
|
|
request: Request,
|
|
user_id: str | None,
|
|
hidden_run_ids: set[str],
|
|
include_middleware: bool,
|
|
include_extra: bool,
|
|
batch_size: int,
|
|
) -> tuple[list[dict[str, Any]], bool]:
|
|
"""Scan raw message rows until ``limit`` visible rows survive filtering."""
|
|
event_store = get_run_event_store(request)
|
|
needed = limit + 1 if include_extra else limit
|
|
|
|
if after_seq is not None:
|
|
visible: list[dict[str, Any]] = []
|
|
scan_after = after_seq
|
|
while len(visible) < needed:
|
|
raw = await event_store.list_messages(
|
|
thread_id,
|
|
limit=batch_size,
|
|
after_seq=scan_after,
|
|
user_id=user_id,
|
|
)
|
|
if not raw:
|
|
break
|
|
_validate_message_scan_rows(raw, thread_id=thread_id, scan_before=None, scan_after=scan_after)
|
|
reached_before_bound = False
|
|
for row in raw:
|
|
if before_seq is not None and row["seq"] >= before_seq:
|
|
reached_before_bound = True
|
|
break
|
|
if (not include_middleware and _is_thread_history_hidden_message_row(row)) or row.get("run_id") in hidden_run_ids:
|
|
continue
|
|
visible.append(row)
|
|
if len(visible) == needed:
|
|
break
|
|
next_scan_after = max(row["seq"] for row in raw)
|
|
if next_scan_after <= scan_after:
|
|
_raise_non_advancing_message_scan(thread_id=thread_id, scan_before=None, scan_after=scan_after, next_cursor=next_scan_after, row_count=len(raw))
|
|
scan_after = next_scan_after
|
|
if reached_before_bound or len(raw) < batch_size:
|
|
break
|
|
has_more = len(visible) > limit
|
|
return visible[:limit], has_more
|
|
|
|
visible_desc: list[dict[str, Any]] = []
|
|
scan_before = before_seq
|
|
while len(visible_desc) < needed:
|
|
raw = await event_store.list_messages(
|
|
thread_id,
|
|
limit=batch_size,
|
|
before_seq=scan_before,
|
|
user_id=user_id,
|
|
)
|
|
if not raw:
|
|
break
|
|
_validate_message_scan_rows(raw, thread_id=thread_id, scan_before=scan_before, scan_after=None)
|
|
for row in reversed(raw):
|
|
if (not include_middleware and _is_thread_history_hidden_message_row(row)) or row.get("run_id") in hidden_run_ids:
|
|
continue
|
|
visible_desc.append(row)
|
|
if len(visible_desc) == needed:
|
|
break
|
|
next_scan_before = min(row["seq"] for row in raw)
|
|
if scan_before is not None and next_scan_before >= scan_before:
|
|
_raise_non_advancing_message_scan(thread_id=thread_id, scan_before=scan_before, scan_after=None, next_cursor=next_scan_before, row_count=len(raw))
|
|
scan_before = next_scan_before
|
|
if len(raw) < batch_size:
|
|
break
|
|
has_more = len(visible_desc) > limit
|
|
return list(reversed(visible_desc[:limit])), has_more
|
|
|
|
|
|
def _validate_message_scan_rows(
|
|
rows: list[dict[str, Any]],
|
|
*,
|
|
thread_id: str,
|
|
scan_before: int | None,
|
|
scan_after: int | None,
|
|
) -> None:
|
|
invalid_seq_rows = [row for row in rows if not isinstance(row.get("seq"), int)]
|
|
if invalid_seq_rows:
|
|
logger.error(
|
|
"Thread message scan found rows without sequence values: thread_id=%s scan_before=%s scan_after=%s row_count=%d invalid_count=%d",
|
|
thread_id,
|
|
scan_before,
|
|
scan_after,
|
|
len(rows),
|
|
len(invalid_seq_rows),
|
|
)
|
|
raise RuntimeError("Run event message rows are missing sequence values")
|
|
|
|
|
|
def _raise_non_advancing_message_scan(
|
|
*,
|
|
thread_id: str,
|
|
scan_before: int | None,
|
|
scan_after: int | None,
|
|
next_cursor: int,
|
|
row_count: int,
|
|
) -> None:
|
|
logger.error(
|
|
"Thread message scan cursor did not advance: thread_id=%s scan_before=%s scan_after=%s next_cursor=%s row_count=%d",
|
|
thread_id,
|
|
scan_before,
|
|
scan_after,
|
|
next_cursor,
|
|
row_count,
|
|
)
|
|
raise RuntimeError("Run event message scan did not advance its cursor")
|
|
|
|
|
|
async def _scan_thread_message_page(
|
|
thread_id: str,
|
|
*,
|
|
limit: int,
|
|
before_seq: int | None,
|
|
request: Request,
|
|
user_id: str | None,
|
|
) -> tuple[list[dict[str, Any]], bool]:
|
|
"""Select the newest ``limit + 1`` page-eligible rows before a cursor."""
|
|
run_mgr = get_run_manager(request)
|
|
hidden_run_ids = await _default_history_hidden_run_ids(run_mgr, thread_id, user_id=user_id)
|
|
return await _scan_visible_thread_messages(
|
|
thread_id,
|
|
limit=limit,
|
|
before_seq=before_seq,
|
|
after_seq=None,
|
|
request=request,
|
|
user_id=user_id,
|
|
hidden_run_ids=hidden_run_ids,
|
|
include_middleware=False,
|
|
include_extra=True,
|
|
batch_size=THREAD_MESSAGE_PAGE_SCAN_BATCH,
|
|
)
|
|
|
|
|
|
async def _enrich_thread_message_page(
|
|
thread_id: str,
|
|
rows: list[dict[str, Any]],
|
|
*,
|
|
request: Request,
|
|
user_id: str | None,
|
|
) -> list[dict[str, Any]]:
|
|
"""Attach run-scoped duration and feedback without mutating store rows."""
|
|
data = deepcopy(rows)
|
|
if not data:
|
|
return data
|
|
|
|
run_ids = {row["run_id"] for row in data if isinstance(row.get("run_id"), str)}
|
|
run_mgr = get_run_manager(request)
|
|
records = await run_mgr.get_many_by_thread(thread_id, run_ids, user_id=user_id)
|
|
run_durations = compute_run_durations(records.values())
|
|
|
|
event_store = get_run_event_store(request)
|
|
last_ai_seq_by_run = await event_store.get_last_visible_ai_seq_by_run(thread_id, run_ids, user_id=user_id)
|
|
feedback_map: dict[str, dict] = {}
|
|
feedback_run_ids = {run_id for row in data if isinstance((run_id := row.get("run_id")), str) and row.get("seq") == last_ai_seq_by_run.get(run_id)}
|
|
if feedback_run_ids:
|
|
feedback_repo = get_feedback_repo(request)
|
|
feedback_map = await feedback_repo.list_by_run_ids(thread_id, feedback_run_ids, user_id=user_id)
|
|
|
|
for row in data:
|
|
run_id = row.get("run_id")
|
|
row["feedback"] = None
|
|
if row.get("seq") == last_ai_seq_by_run.get(run_id):
|
|
feedback = feedback_map.get(run_id)
|
|
if feedback:
|
|
row["feedback"] = {
|
|
"feedback_id": feedback["feedback_id"],
|
|
"rating": feedback["rating"],
|
|
"comment": feedback.get("comment"),
|
|
}
|
|
|
|
# ``turn_duration`` is the run's wall-clock lifetime, not model thinking
|
|
# time — stamp it on the run's LAST visible AI message only so the UI does
|
|
# not repeat the same number on every intermediate AI message of a
|
|
# multi-step turn (#4152). The legacy ``GET /messages`` and ``/history``
|
|
# endpoints already use ``stamp_turn_duration_on_last_ai``; the page
|
|
# endpoint was inlining the equivalent loop but stamping every AI row,
|
|
# which #4163 fixed for the other paths and missed here.
|
|
stamp_turn_duration_on_last_ai(data, run_durations)
|
|
return data
|
|
|
|
|
|
@router.get("/{thread_id}/messages/page", response_model=ThreadMessagesPageResponse)
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def list_thread_messages_page(
|
|
thread_id: ThreadId,
|
|
request: Request,
|
|
limit: int = Query(default=50, ge=1, le=200),
|
|
before_seq: int | None = Query(default=None, ge=1),
|
|
) -> ThreadMessagesPageResponse:
|
|
"""Return a backward page ordered by the thread-global event sequence."""
|
|
if "after_seq" in request.query_params:
|
|
raise HTTPException(status_code=422, detail="after_seq is not supported by this backward-only endpoint")
|
|
|
|
user_id = await get_current_user(request)
|
|
rows, has_more = await _scan_thread_message_page(
|
|
thread_id,
|
|
limit=limit,
|
|
before_seq=before_seq,
|
|
request=request,
|
|
user_id=user_id,
|
|
)
|
|
data = await _enrich_thread_message_page(thread_id, rows, request=request, user_id=user_id)
|
|
return ThreadMessagesPageResponse(
|
|
data=data,
|
|
has_more=has_more,
|
|
next_before_seq=data[0]["seq"] if has_more else None,
|
|
)
|
|
|
|
|
|
@router.get("/{thread_id}/runs/{run_id}/messages")
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def list_run_messages(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
limit: int = Query(default=50, le=200, ge=1),
|
|
before_seq: int | None = Query(default=None, ge=1),
|
|
after_seq: int | None = Query(default=None, ge=1),
|
|
) -> dict:
|
|
"""Return paginated messages for a specific run.
|
|
|
|
Response: { data: [...], has_more: bool }
|
|
"""
|
|
event_store = get_run_event_store(request)
|
|
rows = await event_store.list_messages_by_run(
|
|
thread_id,
|
|
run_id,
|
|
limit=limit + 1,
|
|
before_seq=before_seq,
|
|
after_seq=after_seq,
|
|
)
|
|
data, has_more = trim_run_message_page(rows, limit=limit, after_seq=after_seq)
|
|
|
|
if data:
|
|
run_mgr = get_run_manager(request)
|
|
record = await run_mgr.get(run_id)
|
|
if record:
|
|
durations = compute_run_durations([record])
|
|
if durations:
|
|
stamp_turn_duration_on_last_ai(data, durations)
|
|
|
|
return {"data": data, "has_more": has_more}
|
|
|
|
|
|
def _archive_response_chunks(result: ArtifactArchiveResult):
|
|
try:
|
|
while chunk := result.file.read(1024 * 1024):
|
|
yield chunk
|
|
finally:
|
|
result.file.close()
|
|
|
|
|
|
async def _build_archive_without_abandoning_worker(
|
|
outputs_dir,
|
|
user_data_dir,
|
|
presented_paths: list[str],
|
|
*,
|
|
extra_reserved_dir_names: set[str],
|
|
) -> ArtifactArchiveResult:
|
|
if _artifact_archive_slots.locked():
|
|
raise ArtifactArchiveError("Too many artifact archives are being created; try again shortly", 429)
|
|
await _artifact_archive_slots.acquire()
|
|
build_task = asyncio.create_task(
|
|
asyncio.to_thread(
|
|
build_artifact_archive,
|
|
outputs_dir,
|
|
presented_paths,
|
|
user_data_dir=user_data_dir,
|
|
extra_reserved_dir_names=extra_reserved_dir_names,
|
|
)
|
|
)
|
|
try:
|
|
return await asyncio.shield(build_task)
|
|
except asyncio.CancelledError:
|
|
while not build_task.done():
|
|
try:
|
|
await asyncio.shield(build_task)
|
|
except asyncio.CancelledError:
|
|
continue
|
|
except Exception:
|
|
break
|
|
if not build_task.cancelled():
|
|
try:
|
|
build_task.result().file.close()
|
|
except Exception:
|
|
pass
|
|
raise
|
|
finally:
|
|
_artifact_archive_slots.release()
|
|
|
|
|
|
def _presented_files_from_delivery(events: list[dict]) -> list[str]:
|
|
if len(events) != 1:
|
|
raise HTTPException(status_code=409, detail="This response has no verified artifact delivery")
|
|
content = events[0].get("content")
|
|
by_tool = content.get("by_tool") if isinstance(content, dict) else None
|
|
presented = by_tool.get("present_files") if isinstance(by_tool, dict) else None
|
|
if not isinstance(presented, list) or not presented or any(not isinstance(path, str) for path in presented):
|
|
raise HTTPException(status_code=409, detail="This response has no verified artifact delivery")
|
|
return presented
|
|
|
|
|
|
async def _archive_presented_paths(thread_id: ThreadId, run_id: str, request: Request) -> list[str]:
|
|
run = await get_run_store(request).get(run_id)
|
|
if run is None or run.get("thread_id") != thread_id or run.get("operation_kind", "run") != "run":
|
|
raise HTTPException(status_code=404, detail=f"Run {run_id} not found")
|
|
if run.get("status") in {RunStatus.pending.value, RunStatus.running.value}:
|
|
raise HTTPException(status_code=409, detail="This run has not finished")
|
|
|
|
events = await get_run_event_store(request).list_events(
|
|
thread_id,
|
|
run_id,
|
|
event_types=["run.delivery"],
|
|
limit=2,
|
|
)
|
|
return _presented_files_from_delivery(events)
|
|
|
|
|
|
@router.get(
|
|
"/{thread_id}/runs/{run_id}/artifacts/archive",
|
|
response_model=ArtifactArchiveManifestResponse,
|
|
)
|
|
@require_permission("runs", "read", owner_check=True, require_existing=True)
|
|
async def get_run_artifact_archive_manifest(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
) -> ArtifactArchiveManifestResponse:
|
|
"""Return the verified terminal delivery count used by the archive."""
|
|
presented_paths = await _archive_presented_paths(thread_id, run_id, request)
|
|
return ArtifactArchiveManifestResponse(file_count=len(dict.fromkeys(presented_paths)))
|
|
|
|
|
|
@router.post("/{thread_id}/runs/{run_id}/artifacts/archive")
|
|
@require_permission("runs", "read", owner_check=True, require_existing=True)
|
|
async def create_run_artifact_archive(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
) -> StreamingResponse:
|
|
"""Download the current contents of the files presented by one terminal run."""
|
|
presented_paths = await _archive_presented_paths(thread_id, run_id, request)
|
|
|
|
raw_owner_user_id = get_trusted_internal_owner_user_id(request)
|
|
effective_user_id = make_safe_user_id(raw_owner_user_id) if raw_owner_user_id else get_effective_user_id()
|
|
app_config = await safe_app_config_async()
|
|
custom_tool_output_dir = getattr(getattr(app_config, "tool_output", None), "storage_subdir", None)
|
|
extra_reserved_dir_names = {custom_tool_output_dir} if isinstance(custom_tool_output_dir, str) else set()
|
|
paths = get_paths()
|
|
user_data_dir = paths.sandbox_user_data_dir(thread_id, user_id=effective_user_id)
|
|
outputs_dir = paths.sandbox_outputs_dir(thread_id, user_id=effective_user_id)
|
|
|
|
try:
|
|
async with get_run_manager(request).reserve_thread_operation(
|
|
thread_id,
|
|
kind=ThreadOperationKind.artifact_archive,
|
|
user_id=effective_user_id,
|
|
):
|
|
result = await _build_archive_without_abandoning_worker(
|
|
outputs_dir,
|
|
user_data_dir,
|
|
presented_paths,
|
|
extra_reserved_dir_names=extra_reserved_dir_names,
|
|
)
|
|
except ConflictError as exc:
|
|
raise HTTPException(status_code=409, detail="Artifacts are currently being modified; try again shortly") from exc
|
|
except ArtifactArchiveError as exc:
|
|
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc
|
|
|
|
safe_run_id = re.sub(r"[^A-Za-z0-9_-]", "", run_id)[:32] or "run"
|
|
logger.info(
|
|
"Created artifact archive thread_id=%s run_id=%s members=%d input_bytes=%d output_bytes=%d",
|
|
sanitize_log_param(thread_id),
|
|
sanitize_log_param(run_id),
|
|
result.member_count,
|
|
result.input_bytes,
|
|
result.size,
|
|
)
|
|
return StreamingResponse(
|
|
_archive_response_chunks(result),
|
|
media_type="application/zip",
|
|
headers={
|
|
"Content-Disposition": f'attachment; filename="artifacts-{safe_run_id}.zip"',
|
|
"Content-Length": str(result.size),
|
|
"Cache-Control": "private, no-store",
|
|
"X-Content-Type-Options": "nosniff",
|
|
},
|
|
background=BackgroundTask(result.file.close),
|
|
)
|
|
|
|
|
|
@router.get("/{thread_id}/runs/{run_id}/events")
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def list_run_events(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
event_types: str | None = Query(default=None),
|
|
task_id: str | None = Query(default=None),
|
|
limit: int = Query(default=500, ge=1, le=2000),
|
|
after_seq: int | None = Query(default=None, ge=1),
|
|
) -> list[dict]:
|
|
"""Return the full event stream for a run (debug/audit).
|
|
|
|
``task_id`` + ``after_seq`` let the subtask card page through one subagent
|
|
task's persisted steps without the run-wide ``limit`` truncating the tail (#3779).
|
|
"""
|
|
event_store = get_run_event_store(request)
|
|
types = event_types.split(",") if event_types else None
|
|
events = await event_store.list_events(
|
|
thread_id,
|
|
run_id,
|
|
event_types=types,
|
|
task_id=task_id,
|
|
limit=limit,
|
|
after_seq=after_seq,
|
|
)
|
|
return [
|
|
{
|
|
**event,
|
|
"metadata": redact_metadata_secrets(event.get("metadata")),
|
|
}
|
|
if isinstance(event, dict) and "metadata" in event
|
|
else event
|
|
for event in events
|
|
]
|
|
|
|
|
|
@router.get("/{thread_id}/runs/{run_id}/workspace-changes")
|
|
@require_permission("runs", "read", owner_check=True)
|
|
async def get_run_workspace_changes(
|
|
thread_id: ThreadId,
|
|
run_id: str,
|
|
request: Request,
|
|
include_files: bool = Query(default=True),
|
|
include_diff: bool = Query(default=True),
|
|
) -> dict:
|
|
"""Return workspace/output file changes recorded for one run."""
|
|
event_store = get_run_event_store(request)
|
|
return await get_workspace_changes_response(
|
|
event_store,
|
|
thread_id,
|
|
run_id,
|
|
include_files=include_files,
|
|
include_diff=include_diff,
|
|
)
|
|
|
|
|
|
@router.get("/{thread_id}/token-usage", response_model=ThreadTokenUsageResponse)
|
|
@require_permission("threads", "read", owner_check=True)
|
|
async def thread_token_usage(
|
|
thread_id: ThreadId,
|
|
request: Request,
|
|
include_active: bool = Query(default=False, description="Include running run progress snapshots"),
|
|
) -> ThreadTokenUsageResponse:
|
|
"""Thread-level token usage aggregation."""
|
|
run_store = get_run_store(request)
|
|
if include_active:
|
|
agg = await run_store.aggregate_tokens_by_thread(thread_id, include_active=True)
|
|
else:
|
|
agg = await run_store.aggregate_tokens_by_thread(thread_id)
|
|
context_usage = await build_context_usage(request, thread_id, run_store)
|
|
return ThreadTokenUsageResponse(thread_id=thread_id, context_usage=context_usage, **agg)
|