fix(gateway):unify thread id validation (#4589)

* fix(gateway): unify thread ID validation at the API boundary

Thread ID entry points accepted arbitrary strings while downstream
consumers (filesystem paths, Kubernetes Provisioner, JSONL event store)
each enforced different character restrictions, so invalid IDs were
persisted first and only failed later during sandbox/workspace init.

Centralize validation in deerflow.utils.thread_id (pattern
^[A-Za-z0-9_-]{1,64}$): validate at routers, RunCreateRequest,
scheduler dispatch, paths.py, JSONL store, embedded client, and align
the Provisioner pattern (pinned by a parity test). UUIDs are still
generated only when no ID is supplied; caller-supplied opaque IDs stay
supported.

Deliberate exceptions: DELETE /threads/{id} keeps str as the legacy
cleanup escape hatch (filesystem cleanup guarded), read-only
client.get_thread stays unvalidated, and scheduler rows with legacy
invalid IDs record a failed dispatch instead of raising out of the
poll loop.

* docs: document canonical thread ID contract

README: caller-supplied thread IDs need not be UUIDs; the canonical
pattern and per-endpoint behavior. AGENTS.md: the shared
deerflow.utils.thread_id contract, its enforcement boundaries, and the
legacy-ID escape hatches.

* fix(gateway): close thread ID validation gaps at remaining entry points

Follow-up to the canonical thread ID contract: a full audit found the
uniform-422 coverage only reached about half of the thread_id surfaces.

- routers: 18 routes still took a bare thread_id: str — 13 in
  thread_runs.py (including the five messages/events/workspace-changes
  reads that returned 500 on the JSONL event store vs 404/empty on the
  DB store), 4 read routes in threads.py, and the suggestions route
  flagged in review. DELETE /api/threads/{id} keeps str as the declared
  legacy-cleanup escape hatch.
- client: upload_files/delete_upload/list_uploads/get_artifact now
  validate up front, fulfilling the RFC's 'all mutating entry points'
  clause (get_thread stays unvalidated as the declared legacy read path).
- tui: the /resume literal-ref fallback validates against the canonical
  contract and reports a descriptive error instead of failing deep in
  the client.
- scripts/support_bundle.py: replace the drifted dot-allowing pattern
  with a byte-identical copy of THREAD_ID_PATTERN (kept local so the
  script still runs with a broken venv).

* test(gateway): guard the canonical thread ID contract against regressions

- test_thread_id_route_contract.py: static AST sweep asserting every
  route handler with a thread_id parameter annotates ThreadId
  (whitelist: the DELETE escape hatch), plus a runtime sweep hitting
  all 44 thread_id routes with a non-canonical ID and asserting a 422
  that names thread_id, plus a websocket upgrade-rejection case.
- test_thread_id_validation.py: client entry-point validation,
  support_bundle pattern parity, and TUI literal-ref fallback tests.
- Align two tests that encoded the old contract (dotted IDs).
This commit is contained in:
Vanzeren 2026-08-01 19:42:44 +08:00 committed by GitHub
parent 8234370a6a
commit 095092418c
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
33 changed files with 766 additions and 90 deletions

View File

@ -1048,6 +1048,18 @@ DeerFlow is model-agnostic — it works with any LLM that implements the OpenAI-
DeerFlow can be used as an embedded Python library without running the full HTTP services. The `DeerFlowClient` provides direct in-process access to all agent and Gateway capabilities, returning the same response schemas as the HTTP Gateway API. The HTTP Gateway also exposes `DELETE /api/threads/{thread_id}` to remove DeerFlow-managed local thread data after the LangGraph thread itself has been deleted:
Thread IDs may be supplied by callers and do not have to be UUIDs. Explicit
IDs must contain 164 ASCII letters, digits, hyphens, or underscores
(`^[A-Za-z0-9_-]{1,64}$`). DeerFlow generates a UUID only when `thread_id` is
omitted or `None`; an explicitly supplied empty string is invalid.
Existing route-addressable threads created under older, looser rules remain
readable and deletable, but cannot start new runs or create new filesystem or
sandbox state. Legacy deletion skips local path cleanup when the ID is not
safe under the canonical contract. For canonical legacy threads whose
conversation exists only in LangGraph checkpoints, DeerFlow seeds an empty
run-event feed from the checkpoint before the first new run so
`/messages/page` keeps both the old and new turns.
```python
from deerflow.client import DeerFlowClient

View File

@ -501,6 +501,19 @@ Localhost persistence deliberately reads the direct request `Host` and ignores `
| **GitHub Webhooks** (`/api/webhooks/github`) | `POST /` - receive GitHub App / repo webhook deliveries. Verifies `X-Hub-Signature-256` against `GITHUB_WEBHOOK_SECRET`; exempt from auth + CSRF because authenticity is enforced by HMAC. The route is fail-closed: mounted only when `GITHUB_WEBHOOK_SECRET` is set, or when explicit dev opt-in `DEER_FLOW_ALLOW_UNVERIFIED_GITHUB_WEBHOOKS=1` is set. Recognized events include `ping`, `issues`, `issue_comment`, `pull_request`, `pull_request_review`, and `pull_request_review_comment`; unknown events return 200 with `handled=false`. Fan-out runtime failures return 503, keeping the delivery recorded as failed for manual/API/scripted redelivery (GitHub does not automatically retry any failed delivery, 5xx included); permanent/non-retryable conditions such as `channels.github.enabled: false`, unknown events, malformed payloads, or unavailable channel service return 200 with a skipped/handled response. |
| **GitHub Event-Driven Agents** | Custom agents can declare a `github:` block in their `config.yaml` to bind to repos and event triggers. Webhook fan-out publishes one `InboundMessage` per matching binding to the channel bus; `GitHubChannel` routes those messages through `ChannelManager`. The response `dispatch` summarizes matched/fired/skipped agents. |
Thread identifiers use the shared `deerflow.utils.thread_id` contract
`^[A-Za-z0-9_-]{1,64}$`. Caller-provided opaque IDs remain supported; UUIDs
are generated only for `None`, while explicit empty strings fail validation.
Gateway creation and state-producing request boundaries, embedded-client
entry points, filesystem/upload/event-store consumers, scheduled launches,
and the standalone Provisioner enforce the same contract before persistence
or workspace initialization. Route-addressable legacy IDs remain accepted by
pure reads and cleanup/control endpoints. Deleting a noncanonical legacy ID
best-effort removes its metadata and checkpoints but deliberately skips local
filesystem cleanup, so the raw value is never interpolated into a host path;
new runs, workspace/sandbox operations, and other state-producing mutations
remain blocked.
**Workspace change review**: `packages/harness/deerflow/workspace_changes/`
captures a pre-run and post-run snapshot of the thread-owned `workspace` and
`outputs` directories. `runtime/runs/worker.py` performs the filesystem scan via

View File

@ -23,6 +23,7 @@ from deerflow.config.paths import make_safe_user_id
from deerflow.runtime import ConflictError, ThreadOperationKind
from deerflow.runtime.user_context import get_effective_user_id
from deerflow.sandbox.sandbox_provider import get_sandbox_provider
from deerflow.utils.thread_id import ThreadId
logger = logging.getLogger(__name__)
@ -271,7 +272,7 @@ def _read_artifact_payload(actual_path: Path, path: str, download: bool) -> tupl
description="Retrieve an artifact file generated by the AI agent. Text and binary files can be viewed inline, while active web content is always downloaded.",
)
@require_permission("threads", "read", owner_check=True)
async def get_artifact(thread_id: str, path: str, request: Request, download: bool = False) -> Response:
async def get_artifact(thread_id: ThreadId, path: str, request: Request, download: bool = False) -> Response:
"""Get an artifact file by its path.
The endpoint automatically detects file types and returns appropriate content types.
@ -385,7 +386,7 @@ async def get_artifact(thread_id: str, path: str, request: Request, download: bo
)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def update_artifact(
thread_id: str,
thread_id: ThreadId,
path: str,
body: ArtifactUpdateRequest,
request: Request,

View File

@ -10,6 +10,7 @@ from app.gateway.authz import require_permission
from app.gateway.browser_capability import browser_capability
from deerflow.config.paths import get_paths
from deerflow.runtime.user_context import get_effective_user_id, reset_current_user, set_current_user
from deerflow.utils.thread_id import ThreadId
logger = logging.getLogger(__name__)
@ -76,7 +77,7 @@ async def _browser_thread_owned_by(thread_store, thread_id: str, user_id: str) -
description="Steer the thread's live browser session to a URL from the UI and capture a screenshot.",
)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def navigate_browser(thread_id: str, body: BrowserNavigateRequest, request: Request) -> BrowserNavigateResponse:
async def navigate_browser(thread_id: ThreadId, body: BrowserNavigateRequest, request: Request) -> BrowserNavigateResponse:
user_id = str(request.state.auth.user.id)
thread_store = getattr(request.app.state, "thread_store", None)
if thread_store is None or not await _browser_thread_owned_by(thread_store, thread_id, user_id):
@ -176,7 +177,7 @@ def _ws_origin_allowed(websocket: WebSocket) -> bool:
@router.websocket("/threads/{thread_id}/browser/stream")
async def browser_stream(websocket: WebSocket, thread_id: str) -> None:
async def browser_stream(websocket: WebSocket, thread_id: ThreadId) -> None:
"""Bidirectional live browser stream.
Server client: JSON ``{"type":"frame","data":"<base64 jpeg>"}`` frames

View File

@ -14,6 +14,7 @@ from pydantic import BaseModel, Field
from app.gateway.authz import require_permission
from app.gateway.deps import get_current_user, get_feedback_repo, get_run_store
from deerflow.utils.thread_id import ThreadId
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/threads", tags=["feedback"])
@ -61,7 +62,7 @@ class FeedbackStatsResponse(BaseModel):
@router.put("/{thread_id}/runs/{run_id}/feedback", response_model=FeedbackResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def upsert_feedback(
thread_id: str,
thread_id: ThreadId,
run_id: str,
body: FeedbackUpsertRequest,
request: Request,
@ -92,7 +93,7 @@ async def upsert_feedback(
@router.delete("/{thread_id}/runs/{run_id}/feedback")
@require_permission("threads", "delete", owner_check=True, require_existing=True)
async def delete_run_feedback(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
) -> dict[str, bool]:
@ -112,7 +113,7 @@ async def delete_run_feedback(
@router.post("/{thread_id}/runs/{run_id}/feedback", response_model=FeedbackResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def create_feedback(
thread_id: str,
thread_id: ThreadId,
run_id: str,
body: FeedbackCreateRequest,
request: Request,
@ -145,7 +146,7 @@ async def create_feedback(
@router.get("/{thread_id}/runs/{run_id}/feedback", response_model=list[FeedbackResponse])
@require_permission("threads", "read", owner_check=True)
async def list_feedback(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
) -> list[dict[str, Any]]:
@ -157,7 +158,7 @@ async def list_feedback(
@router.get("/{thread_id}/runs/{run_id}/feedback/stats", response_model=FeedbackStatsResponse)
@require_permission("threads", "read", owner_check=True)
async def feedback_stats(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
) -> dict[str, Any]:
@ -169,7 +170,7 @@ async def feedback_stats(
@router.delete("/{thread_id}/runs/{run_id}/feedback/{feedback_id}")
@require_permission("threads", "delete", owner_check=True, require_existing=True)
async def delete_feedback(
thread_id: str,
thread_id: ThreadId,
run_id: str,
feedback_id: str,
request: Request,

View File

@ -8,7 +8,6 @@ is reused so that conversation history is preserved across calls.
from __future__ import annotations
import logging
import uuid
from fastapi import APIRouter, HTTPException, Query, Request
from fastapi.responses import StreamingResponse
@ -19,6 +18,7 @@ from app.gateway.pagination import trim_run_message_page
from app.gateway.run_models import RunCreateRequest
from app.gateway.services import build_checkpoint_state_accessor, sse_consumer, start_run, wait_for_run_completion
from deerflow.runtime import serialize_channel_values_for_api
from deerflow.utils.thread_id import resolve_thread_id
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/runs", tags=["runs"])
@ -27,9 +27,7 @@ router = APIRouter(prefix="/api/runs", tags=["runs"])
def _resolve_thread_id(body: RunCreateRequest) -> str:
"""Return the thread_id from the request body, or generate a new one."""
thread_id = ((body.config or {}).get("configurable") or {}).get("thread_id")
if thread_id:
return str(thread_id)
return str(uuid.uuid4())
return resolve_thread_id(thread_id)
@router.post("/stream")

View File

@ -23,6 +23,7 @@ from deerflow.scheduler.schedules import (
normalize_cron_expression,
validate_timezone,
)
from deerflow.utils.thread_id import ThreadId
router = APIRouter(prefix="/api", tags=["scheduled-tasks"])
@ -36,7 +37,7 @@ def _ensure_task_mutable(task: dict[str, Any]) -> None:
class ScheduledTaskCreateRequest(BaseModel):
thread_id: str | None = None
thread_id: ThreadId | None = None
context_mode: str = "fresh_thread_per_run"
title: str = Field(min_length=1)
prompt: str = Field(min_length=1)
@ -47,7 +48,7 @@ class ScheduledTaskCreateRequest(BaseModel):
class ScheduledTaskUpdateRequest(BaseModel):
context_mode: str | None = None
thread_id: str | None = None
thread_id: ThreadId | None = None
title: str | None = Field(default=None, min_length=1)
prompt: str | None = Field(default=None, min_length=1)
schedule_spec: dict[str, Any] | None = None
@ -302,7 +303,7 @@ async def list_scheduled_task_runs(
@router.get("/threads/{thread_id}/scheduled-tasks")
@require_permission("threads", "read", owner_check=True)
async def list_thread_scheduled_tasks(thread_id: str, request: Request):
async def list_thread_scheduled_tasks(thread_id: ThreadId, request: Request):
repo = get_scheduled_task_repo(request)
user = await get_optional_user_from_request(request)
if user is None:

View File

@ -31,6 +31,7 @@ from deerflow.skills.security_static_scanner import (
)
from deerflow.skills.storage import SkillStorage, get_or_new_user_skill_storage
from deerflow.skills.types import SKILL_MD_FILE, SkillCategory
from deerflow.utils.thread_id import ThreadId
logger = logging.getLogger(__name__)
@ -65,7 +66,7 @@ class SkillUpdateRequest(BaseModel):
class SkillInstallRequest(BaseModel):
"""Request model for installing a skill from a .skill file."""
thread_id: str = Field(..., description="The thread ID where the .skill file is located")
thread_id: ThreadId = Field(..., description="The thread ID where the .skill file is located")
path: str = Field(..., description="Virtual path to the .skill file (e.g., mnt/user-data/outputs/my-skill.skill)")

View File

@ -10,6 +10,7 @@ from app.gateway.deps import get_config
from deerflow.config.app_config import AppConfig
from deerflow.config.suggestions_config import DEFAULT_MAX_SUGGESTIONS, MAX_SUGGESTIONS_LIMIT
from deerflow.utils.oneshot_llm import run_oneshot_llm
from deerflow.utils.thread_id import ThreadId
logger = logging.getLogger(__name__)
@ -102,7 +103,7 @@ async def get_suggestions_config(
)
@require_permission("threads", "read", owner_check=True)
async def generate_suggestions(
thread_id: str,
thread_id: ThreadId,
body: SuggestionsRequest,
request: Request,
config: AppConfig = Depends(get_config),

View File

@ -43,6 +43,7 @@ from deerflow.agents.middlewares.dynamic_context_middleware import strip_injecte
from deerflow.runtime import CancelOutcome, RunRecord, RunStatus, serialize_channel_values_for_api
from deerflow.runtime.secret_context import redact_config_secrets, redact_metadata_secrets
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__)
@ -813,7 +814,7 @@ async def _default_history_hidden_run_ids(run_mgr: Any, thread_id: str, *, user_
@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: str,
thread_id: ThreadId,
body: RegeneratePrepareRequest,
request: Request,
) -> RegeneratePrepareResponse:
@ -824,7 +825,7 @@ async def prepare_regenerate_run(
@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: str,
thread_id: ThreadId,
body: EditRegeneratePrepareRequest,
request: Request,
) -> EditRegeneratePrepareResponse:
@ -834,7 +835,7 @@ async def prepare_edit_regenerate_run(
@router.post("/{thread_id}/runs", response_model=RunResponse)
@require_permission("runs", "create", owner_check=True, require_existing=True)
async def create_run(thread_id: str, body: RunCreateRequest, request: Request) -> RunResponse:
async def create_run(thread_id: ThreadId, body: RunCreateRequest, request: Request) -> RunResponse:
"""Create a background run (returns immediately)."""
record = await start_run(body, thread_id, request)
return _record_to_response(record)
@ -842,7 +843,7 @@ async def create_run(thread_id: str, body: RunCreateRequest, request: Request) -
@router.post("/{thread_id}/runs/stream")
@require_permission("runs", "create", owner_check=True, require_existing=True)
async def stream_run(thread_id: str, body: RunCreateRequest, request: Request) -> StreamingResponse:
async def stream_run(thread_id: ThreadId, body: RunCreateRequest, request: Request) -> StreamingResponse:
"""Create a run and stream events via SSE.
The response includes a ``Content-Location`` header with the run's
@ -870,7 +871,7 @@ async def stream_run(thread_id: str, body: RunCreateRequest, request: Request) -
@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: str, body: RunCreateRequest, request: Request) -> dict:
async def wait_run(thread_id: ThreadId, body: RunCreateRequest, request: Request) -> dict:
"""Create a run and block until it completes, returning the final state."""
bridge = get_stream_bridge(request)
run_mgr = get_run_manager(request)
@ -899,7 +900,7 @@ async def wait_run(thread_id: str, body: RunCreateRequest, request: Request) ->
@router.get("/{thread_id}/runs", response_model=list[RunResponse])
@require_permission("runs", "read", owner_check=True)
async def list_runs(thread_id: str, request: Request) -> list[RunResponse]:
async def list_runs(thread_id: ThreadId, request: Request) -> list[RunResponse]:
"""List all runs for a thread."""
run_mgr = get_run_manager(request)
user_id = await get_current_user(request)
@ -909,7 +910,7 @@ async def list_runs(thread_id: str, request: Request) -> list[RunResponse]:
@router.get("/{thread_id}/runs/{run_id}", response_model=RunResponse)
@require_permission("runs", "read", owner_check=True)
async def get_run(thread_id: str, run_id: str, request: Request) -> RunResponse:
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)
@ -922,7 +923,7 @@ async def get_run(thread_id: str, run_id: str, request: Request) -> RunResponse:
@router.post("/{thread_id}/runs/{run_id}/cancel")
@require_permission("runs", "cancel", owner_check=True, require_existing=True)
async def cancel_run(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
wait: bool = Query(default=False, description="Block until run completes after cancel"),
@ -981,7 +982,7 @@ async def cancel_run(
@router.get("/{thread_id}/runs/{run_id}/join")
@require_permission("runs", "read", owner_check=True)
async def join_run(thread_id: str, run_id: str, request: Request) -> StreamingResponse:
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)
@ -1010,7 +1011,7 @@ async def join_run(thread_id: str, run_id: str, request: Request) -> StreamingRe
@router.post("/{thread_id}/runs/{run_id}/stream", response_model=None)
@require_permission("runs", "read", owner_check=True)
async def stream_existing_run(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
action: Literal["interrupt", "rollback"] | None = Query(default=None, description="Cancel action"),
@ -1083,7 +1084,7 @@ async def stream_existing_run(
@router.get("/{thread_id}/messages")
@require_permission("runs", "read", owner_check=True)
async def list_thread_messages(
thread_id: str,
thread_id: ThreadId,
request: Request,
limit: int = Query(default=50, ge=1, le=200),
before_seq: int | None = Query(default=None, ge=1),
@ -1337,7 +1338,7 @@ async def _enrich_thread_message_page(
@router.get("/{thread_id}/messages/page", response_model=ThreadMessagesPageResponse)
@require_permission("runs", "read", owner_check=True)
async def list_thread_messages_page(
thread_id: str,
thread_id: ThreadId,
request: Request,
limit: int = Query(default=50, ge=1, le=200),
before_seq: int | None = Query(default=None, ge=1),
@ -1365,7 +1366,7 @@ async def list_thread_messages_page(
@router.get("/{thread_id}/runs/{run_id}/messages")
@require_permission("runs", "read", owner_check=True)
async def list_run_messages(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
limit: int = Query(default=50, le=200, ge=1),
@ -1400,7 +1401,7 @@ async def list_run_messages(
@router.get("/{thread_id}/runs/{run_id}/events")
@require_permission("runs", "read", owner_check=True)
async def list_run_events(
thread_id: str,
thread_id: ThreadId,
run_id: str,
request: Request,
event_types: str | None = Query(default=None),
@ -1437,7 +1438,7 @@ async def list_run_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: str,
thread_id: ThreadId,
run_id: str,
request: Request,
include_files: bool = Query(default=True),
@ -1457,7 +1458,7 @@ async def get_run_workspace_changes(
@router.get("/{thread_id}/token-usage", response_model=ThreadTokenUsageResponse)
@require_permission("threads", "read", owner_check=True)
async def thread_token_usage(
thread_id: str,
thread_id: ThreadId,
request: Request,
include_active: bool = Query(default=False, description="Include running run progress snapshots"),
) -> ThreadTokenUsageResponse:

View File

@ -68,6 +68,7 @@ from deerflow.runtime.runs.worker import valid_duration_entry
from deerflow.runtime.secret_context import redact_metadata_secrets
from deerflow.runtime.user_context import get_effective_user_id
from deerflow.utils.file_io import run_file_io
from deerflow.utils.thread_id import ThreadId, resolve_thread_id, validate_thread_id
from deerflow.utils.time import coerce_iso, now_iso
logger = logging.getLogger(__name__)
@ -372,7 +373,7 @@ class ThreadResponse(_MetadataRedactingResponse):
class ThreadCreateRequest(BaseModel):
"""Request body for creating a thread."""
thread_id: str | None = Field(default=None, description="Optional thread ID (auto-generated if omitted)")
thread_id: ThreadId | None = Field(default=None, description="Optional thread ID (auto-generated if omitted)")
assistant_id: str | None = Field(default=None, description="Associate thread with an assistant")
metadata: dict[str, Any] = Field(default_factory=dict, description="Initial metadata")
@ -627,8 +628,18 @@ async def delete_thread_data(thread_id: str, request: Request) -> ThreadDeleteRe
"""
from app.gateway.deps import get_thread_store
# Clean local filesystem
response = _delete_thread_data(thread_id, user_id=get_effective_user_id())
# Legacy IDs may predate the canonical filesystem-safe contract. They can
# still be removed from metadata/checkpoint stores, but must never be
# interpolated into a host path during cleanup.
try:
validate_thread_id(thread_id)
except ValueError:
response = ThreadDeleteResponse(
success=True,
message="Skipped local data cleanup for legacy thread ID",
)
else:
response = _delete_thread_data(thread_id, user_id=get_effective_user_id())
# Remove checkpoints (best-effort)
checkpointer = getattr(request.app.state, "checkpointer", None)
@ -707,7 +718,7 @@ async def create_thread(body: ThreadCreateRequest, request: Request) -> ThreadRe
checkpointer = get_checkpointer(request)
thread_store = get_thread_store(request)
thread_id = body.thread_id or str(uuid.uuid4())
thread_id = resolve_thread_id(body.thread_id)
now = now_iso()
thread_owner_user_id = get_trusted_internal_owner_user_id(request)
thread_owner_kwargs = {"user_id": thread_owner_user_id} if thread_owner_user_id else {}
@ -774,7 +785,7 @@ async def create_thread(body: ThreadCreateRequest, request: Request) -> ThreadRe
@router.post("/{thread_id}/branches", response_model=ThreadBranchResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def branch_thread(thread_id: str, body: ThreadBranchRequest, request: Request) -> ThreadBranchResponse:
async def branch_thread(thread_id: ThreadId, body: ThreadBranchRequest, request: Request) -> ThreadBranchResponse:
"""Create a new main-thread branch from a completed assistant turn."""
from app.gateway.deps import get_thread_store
@ -977,7 +988,7 @@ async def search_threads(body: ThreadSearchRequest, request: Request) -> list[Th
@router.patch("/{thread_id}", response_model=ThreadResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def patch_thread(thread_id: str, body: ThreadPatchRequest, request: Request) -> ThreadResponse:
async def patch_thread(thread_id: ThreadId, body: ThreadPatchRequest, request: Request) -> ThreadResponse:
"""Merge metadata into a thread record."""
from app.gateway.deps import get_thread_store
@ -1010,7 +1021,7 @@ async def patch_thread(thread_id: str, body: ThreadPatchRequest, request: Reques
@router.get("/{thread_id}", response_model=ThreadResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread(thread_id: str, request: Request) -> ThreadResponse:
async def get_thread(thread_id: ThreadId, request: Request) -> ThreadResponse:
"""Get thread info from metadata plus the graph's materialized state."""
from app.gateway.deps import get_thread_store
@ -1063,7 +1074,7 @@ async def get_thread(thread_id: str, request: Request) -> ThreadResponse:
@router.get("/{thread_id}/goal", response_model=ThreadGoalResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread_goal(thread_id: str, request: Request) -> ThreadGoalResponse:
async def get_thread_goal(thread_id: ThreadId, request: Request) -> ThreadGoalResponse:
"""Return the active Claude-style goal for a thread, if any."""
checkpointer = get_checkpointer(request)
try:
@ -1076,7 +1087,7 @@ async def get_thread_goal(thread_id: str, request: Request) -> ThreadGoalRespons
@router.put("/{thread_id}/goal", response_model=ThreadGoalResponse)
@require_permission("threads", "write", owner_check=True)
async def set_thread_goal(thread_id: str, body: ThreadGoalRequest, request: Request) -> ThreadGoalResponse:
async def set_thread_goal(thread_id: ThreadId, body: ThreadGoalRequest, request: Request) -> ThreadGoalResponse:
"""Set or replace the active goal for a thread.
``/chats/new`` pages already hold a generated UUID before the first run, so
@ -1102,7 +1113,7 @@ async def set_thread_goal(thread_id: str, body: ThreadGoalRequest, request: Requ
@router.delete("/{thread_id}/goal", response_model=ThreadGoalResponse)
@require_permission("threads", "write", owner_check=True)
async def clear_thread_goal(thread_id: str, request: Request) -> ThreadGoalResponse:
async def clear_thread_goal(thread_id: ThreadId, request: Request) -> ThreadGoalResponse:
"""Clear the active goal for a thread."""
checkpointer = get_checkpointer(request)
try:
@ -1133,7 +1144,7 @@ def _thread_compact_response(result: ThreadCompactionResult) -> ThreadCompactRes
@router.post("/{thread_id}/compact", response_model=ThreadCompactResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def compact_thread(thread_id: str, body: ThreadCompactRequest, request: Request) -> ThreadCompactResponse:
async def compact_thread(thread_id: ThreadId, body: ThreadCompactRequest, request: Request) -> ThreadCompactResponse:
"""Manually summarize old thread context while preserving the visible history."""
# Compaction writes only base-schema channels (messages + summary_text);
# every other channel — including middleware-contributed ones — is carried
@ -1180,7 +1191,7 @@ async def compact_thread(thread_id: str, body: ThreadCompactRequest, request: Re
# ---------------------------------------------------------------------------
@router.get("/{thread_id}/state", response_model=ThreadStateResponse)
@require_permission("threads", "read", owner_check=True)
async def get_thread_state(thread_id: str, request: Request) -> ThreadStateResponse:
async def get_thread_state(thread_id: ThreadId, request: Request) -> ThreadStateResponse:
"""Get the latest materialized graph state for a thread."""
# Resolve through the thread's assistant so custom middleware channels
# appear in the response instead of being dropped by the default schema.
@ -1222,7 +1233,7 @@ async def get_thread_state(thread_id: str, request: Request) -> ThreadStateRespo
@router.post("/{thread_id}/state", response_model=ThreadStateResponse)
@require_permission("threads", "write", owner_check=True, require_existing=True)
async def update_thread_state(thread_id: str, body: ThreadStateUpdateRequest, request: Request) -> ThreadStateResponse:
async def update_thread_state(thread_id: ThreadId, body: ThreadStateUpdateRequest, request: Request) -> ThreadStateResponse:
"""Replace selected thread-state fields through the materialized graph."""
from app.gateway.deps import get_thread_store
@ -1322,7 +1333,7 @@ def _checkpoint_run_durations(metadata: Any) -> dict[str, int]:
@router.post("/{thread_id}/history", response_model=list[HistoryEntry])
@require_permission("threads", "read", owner_check=True)
async def get_thread_history(
thread_id: str,
thread_id: ThreadId,
body: ThreadHistoryRequest,
request: Request,
background_tasks: BackgroundTasks,

View File

@ -35,6 +35,7 @@ from deerflow.uploads.manager import (
)
from deerflow.utils.file_conversion import CONVERTIBLE_EXTENSIONS, convert_file_to_markdown
from deerflow.utils.file_io import run_file_io
from deerflow.utils.thread_id import ThreadId
logger = logging.getLogger(__name__)
@ -299,7 +300,7 @@ def _auto_convert_documents_enabled(app_config: AppConfig) -> bool:
@router.post("", response_model=UploadResponse)
@require_permission("threads", "write", owner_check=True, require_existing=False)
async def upload_files(
thread_id: str,
thread_id: ThreadId,
request: Request,
files: list[UploadFile] = File(...),
config: AppConfig = Depends(get_config),
@ -445,7 +446,7 @@ async def upload_files(
@router.get("/limits", response_model=UploadLimits)
@require_permission("threads", "read", owner_check=True)
async def get_upload_limits(
thread_id: str,
thread_id: ThreadId,
request: Request,
config: AppConfig = Depends(get_config),
) -> UploadLimits:
@ -455,7 +456,7 @@ async def get_upload_limits(
@router.get("/list", response_model=UploadListResponse)
@require_permission("threads", "read", owner_check=True)
async def list_uploaded_files(thread_id: str, request: Request) -> UploadListResponse:
async def list_uploaded_files(thread_id: ThreadId, request: Request) -> UploadListResponse:
"""List all files in a thread's uploads directory."""
try:
result = await run_file_io(_list_uploaded_files_for_thread, thread_id, get_effective_user_id())
@ -467,7 +468,7 @@ async def list_uploaded_files(thread_id: str, request: Request) -> UploadListRes
@router.delete("/{filename}")
@require_permission("threads", "delete", owner_check=True, require_existing=True)
async def delete_uploaded_file(thread_id: str, filename: str, request: Request) -> dict:
async def delete_uploaded_file(thread_id: ThreadId, filename: str, request: Request) -> dict:
"""Delete a file from a thread's uploads directory."""
try:
return await run_file_io(_delete_uploaded_file_for_thread, thread_id, filename, get_effective_user_id())

View File

@ -4,10 +4,11 @@ from __future__ import annotations
from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field, ValidationInfo, field_validator
from pydantic import BaseModel, ConfigDict, Field, ValidationInfo, field_validator, model_validator
from pydantic_core import PydanticCustomError
from deerflow.runtime.stream_modes import RunStreamMode, UnsupportedStreamModeError, normalize_stream_modes
from deerflow.utils.thread_id import validate_thread_id
class RunCreateRequest(BaseModel):
@ -36,6 +37,19 @@ class RunCreateRequest(BaseModel):
if_not_exists: Literal["create"] = Field(default="create", description="Compatibility default; missing threads are created")
feedback_keys: None = Field(default=None, description="Compatibility placeholder; feedback key collection is not supported")
@model_validator(mode="after")
def validate_configurable_thread_id(self) -> RunCreateRequest:
"""Validate the stateless-run thread selector inside RunnableConfig."""
if not isinstance(self.config, dict):
return self
configurable = self.config.get("configurable")
if not isinstance(configurable, dict) or "thread_id" not in configurable:
return self
thread_id = configurable["thread_id"]
if thread_id is not None:
validate_thread_id(thread_id)
return self
@field_validator(
"webhook",
"on_completion",

View File

@ -71,6 +71,7 @@ from deerflow.runtime.secret_context import (
from deerflow.runtime.stream_modes import normalize_stream_modes
from deerflow.runtime.user_context import reset_current_user, set_current_user
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
from deerflow.utils.thread_id import validate_thread_id
logger = logging.getLogger(__name__)
@ -1062,6 +1063,11 @@ async def start_run(
request : Request
FastAPI request used to retrieve singletons from ``app.state``.
"""
try:
validate_thread_id(thread_id)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
body_config = getattr(body, "config", None)
config_metadata = body_config.get("metadata") if isinstance(body_config, dict) else None
try:

View File

@ -12,6 +12,7 @@ from fastapi import HTTPException
from deerflow.persistence.scheduled_task_runs import ActiveScheduledRunConflict
from deerflow.runtime import ConflictError, RunRecord
from deerflow.scheduler.schedules import next_run_at
from deerflow.utils.thread_id import validate_thread_id
logger = logging.getLogger(__name__)
@ -105,8 +106,41 @@ class ScheduledTaskService:
trigger: str,
) -> dict[str, Any]:
execution_thread_id = task.get("thread_id")
if task.get("context_mode") == "fresh_thread_per_run" or not execution_thread_id:
if task.get("context_mode") == "fresh_thread_per_run" or execution_thread_id is None:
execution_thread_id = str(uuid.uuid4())
try:
validate_thread_id(execution_thread_id)
except ValueError as exc:
# Rows persisted before the thread-id contract was centralized may
# hold IDs that were valid then (dots, unlimited length) but fail
# the canonical pattern now. Route through the normal failure
# bookkeeping instead of raising: an uncaught ValueError here would
# surface as HTTP 500 on manual trigger and, in the poller, abort
# the rest of the claimed batch every cycle while the task itself
# is never marked with last_error.
task_status = self._task_status_for_failure(task, trigger=trigger)
await self._task_repo.update_after_launch(
task["id"],
status=task_status,
next_run_at=next_run_at(
task["schedule_type"],
task["schedule_spec"],
task["timezone"],
now=now,
),
last_run_at=now,
last_run_id=None,
last_thread_id=execution_thread_id,
last_error=str(exc),
increment_run_count=False,
)
return {
"outcome": "failed",
"task_run_id": None,
"run_id": None,
"thread_id": execution_thread_id,
"error": str(exc),
}
# "skip" must hold for fresh-thread runs too, where every run gets a new
# thread and the same-thread multitask ConflictError below can never
# fire. Checked before creating this dispatch's own run row so the row

View File

@ -72,6 +72,7 @@ from deerflow.uploads.manager import (
upload_artifact_url,
upload_virtual_path,
)
from deerflow.utils.thread_id import resolve_thread_id, validate_thread_id
logger = logging.getLogger(__name__)
@ -525,6 +526,7 @@ class DeerFlowClient:
def get_goal(self, thread_id: str) -> dict:
"""Return the active goal for a thread, if any."""
validate_thread_id(thread_id)
checkpointer = self._get_thread_checkpointer()
goal = _run_async_from_sync(read_thread_goal(checkpointer, thread_id))
return {"goal": goal}
@ -537,6 +539,7 @@ class DeerFlowClient:
max_continuations: int = DEFAULT_MAX_GOAL_CONTINUATIONS,
) -> dict:
"""Set or replace a thread-scoped goal."""
validate_thread_id(thread_id)
checkpointer = self._get_thread_checkpointer()
goal = build_goal_state(objective, max_continuations=max_continuations)
@ -549,6 +552,7 @@ class DeerFlowClient:
def clear_goal(self, thread_id: str) -> dict:
"""Clear the active goal for a thread."""
validate_thread_id(thread_id)
checkpointer = self._get_thread_checkpointer()
async def _clear_goal() -> None:
@ -803,8 +807,7 @@ class DeerFlowClient:
Tool results also include ``"artifact"`` when the source ToolMessage has a non-None artifact.
- type="end" data={"usage": {"input_tokens": int, "output_tokens": int, "total_tokens": int}}
"""
if thread_id is None:
thread_id = str(uuid.uuid4())
thread_id = resolve_thread_id(thread_id)
config = self._get_runnable_config(thread_id, **kwargs)
inject_checkpoint_mode(config, self._checkpoint_channel_mode)
@ -1481,6 +1484,7 @@ class DeerFlowClient:
FileNotFoundError: If any file does not exist.
ValueError: If any supplied path exists but is not a regular file.
"""
validate_thread_id(thread_id)
from deerflow.utils.file_conversion import CONVERTIBLE_EXTENSIONS, convert_file_to_markdown
# Validate all files upfront to avoid partial uploads.
@ -1584,6 +1588,7 @@ class DeerFlowClient:
Dict with "files" and "count" keys, matching the Gateway API
``list_uploaded_files`` response.
"""
validate_thread_id(thread_id)
uploads_dir = get_uploads_dir(thread_id)
result = list_files_in_dir(uploads_dir)
return enrich_file_listing(result, thread_id)
@ -1603,6 +1608,7 @@ class DeerFlowClient:
FileNotFoundError: If the file does not exist.
PermissionError: If path traversal is detected.
"""
validate_thread_id(thread_id)
from deerflow.utils.file_conversion import CONVERTIBLE_EXTENSIONS
uploads_dir = get_uploads_dir(thread_id)
@ -1626,6 +1632,7 @@ class DeerFlowClient:
FileNotFoundError: If the artifact does not exist.
ValueError: If the path is invalid.
"""
validate_thread_id(thread_id)
try:
actual = get_paths().resolve_virtual_path(thread_id, path, user_id=get_effective_user_id())
except ValueError as exc:

View File

@ -6,11 +6,11 @@ import shutil
from pathlib import Path, PureWindowsPath
from deerflow.config.runtime_paths import runtime_home
from deerflow.utils.thread_id import validate_thread_id
# Virtual path prefix seen by agents inside the sandbox
VIRTUAL_PATH_PREFIX = "/mnt/user-data"
_SAFE_THREAD_ID_RE = re.compile(r"^[A-Za-z0-9_\-]+$")
_SAFE_USER_ID_RE = re.compile(r"^[A-Za-z0-9_\-]+$")
_SAFE_INTEGRATION_ID_RE = re.compile(r"^[A-Za-z0-9_.\-]+$")
_UNSAFE_USER_ID_CHAR_RE = re.compile(r"[^A-Za-z0-9_\-]")
@ -26,9 +26,7 @@ def _default_local_base_dir() -> Path:
def _validate_thread_id(thread_id: str) -> str:
"""Validate a thread ID before using it in filesystem paths."""
if not _SAFE_THREAD_ID_RE.match(thread_id):
raise ValueError(f"Invalid thread_id {thread_id!r}: only alphanumeric characters, hyphens, and underscores are allowed.")
return thread_id
return validate_thread_id(thread_id)
def _validate_user_id(user_id: str) -> str:

View File

@ -32,6 +32,7 @@ from typing import Any
from deerflow.runtime.events.store.base import RunEventStore
from deerflow.runtime.user_context import AUTO, _AutoSentinel
from deerflow.utils.thread_id import validate_thread_id
logger = logging.getLogger(__name__)
@ -56,7 +57,7 @@ class JsonlRunEventStore(RunEventStore):
return value
def _thread_dir(self, thread_id: str) -> Path:
self._validate_id(thread_id, "thread_id")
validate_thread_id(thread_id)
return self._base_dir / "threads" / thread_id / "runs"
def _run_file(self, thread_id: str, run_id: str) -> Path:

View File

@ -39,18 +39,29 @@ class Session:
Matches an existing thread by id first, then by exact title. Falls back to
the literal ref (treated as an id) when nothing matches, so an unknown id
still continues/creates that namespace.
still continues/creates that namespace provided it satisfies the
canonical thread ID contract.
"""
try:
threads = self.client.list_threads(limit=100).get("thread_list", [])
except Exception: # noqa: BLE001 - resolution is best-effort
return ref
return self._validated_literal_ref(ref)
if any(t.get("thread_id") == ref for t in threads):
return ref
for thread in threads:
if (thread.get("title") or "") == ref:
return thread.get("thread_id") or ref
return ref
return thread.get("thread_id") or self._validated_literal_ref(ref)
return self._validated_literal_ref(ref)
@staticmethod
def _validated_literal_ref(ref: str) -> str:
"""Validate a literal ref before it is adopted as a thread id."""
from deerflow.utils.thread_id import validate_thread_id
try:
return validate_thread_id(ref)
except ValueError as exc:
raise ValueError(f"Thread reference {ref!r} matches no existing thread and is not a valid thread id (expected 1-64 ASCII letters, digits, hyphens, or underscores).") from exc
def recent_threads(self, limit: int = 20) -> list[dict]:
return self.client.list_threads(limit=limit).get("thread_list", [])

View File

@ -7,13 +7,13 @@ Both Gateway and Client delegate to these functions.
import errno
import logging
import os
import re
import stat
from pathlib import Path
from urllib.parse import quote
from deerflow.config.paths import VIRTUAL_PATH_PREFIX, get_paths
from deerflow.runtime.user_context import get_effective_user_id
from deerflow.utils.thread_id import validate_thread_id
class PathTraversalError(ValueError):
@ -26,22 +26,10 @@ class UnsafeUploadPathError(ValueError):
logger = logging.getLogger(__name__)
# thread_id must be alphanumeric, hyphens, underscores, or dots only.
_SAFE_THREAD_ID = re.compile(r"^[a-zA-Z0-9._-]+$")
UPLOAD_STAGING_PREFIX = ".upload-"
UPLOAD_STAGING_SUFFIX = ".part"
def validate_thread_id(thread_id: str) -> None:
"""Reject thread IDs containing characters unsafe for filesystem paths.
Raises:
ValueError: If thread_id is empty or contains unsafe characters.
"""
if not thread_id or not _SAFE_THREAD_ID.match(thread_id):
raise ValueError(f"Invalid thread_id: {thread_id!r}")
def get_uploads_dir(thread_id: str, *, user_id: str | None = None) -> Path:
"""Return the uploads directory path for a thread (no side effects)."""
validate_thread_id(thread_id)

View File

@ -0,0 +1,37 @@
"""Canonical thread identifier validation shared across DeerFlow backends."""
from __future__ import annotations
import re
import uuid
from typing import Annotated
from pydantic import AfterValidator, StringConstraints
THREAD_ID_PATTERN = r"^[A-Za-z0-9_-]{1,64}$"
_THREAD_ID_RE = re.compile(THREAD_ID_PATTERN)
def validate_thread_id(thread_id: str) -> str:
"""Return a valid thread ID or raise ``ValueError``.
Thread IDs are caller-defined opaque identifiers, not necessarily UUIDs,
but they must be safe for every persistence and filesystem backend.
"""
if not isinstance(thread_id, str) or _THREAD_ID_RE.fullmatch(thread_id) is None:
raise ValueError("Invalid thread_id: expected 1-64 ASCII letters, digits, hyphens, or underscores")
return thread_id
def resolve_thread_id(thread_id: str | None) -> str:
"""Validate a supplied ID, generating a UUID only when it is ``None``."""
if thread_id is None:
return str(uuid.uuid4())
return validate_thread_id(thread_id)
ThreadId = Annotated[
str,
StringConstraints(min_length=1, max_length=64, pattern=THREAD_ID_PATTERN),
AfterValidator(validate_thread_id),
]

View File

@ -2748,3 +2748,17 @@ async def test_run_agent_full_mode_checks_selected_checkpoint_before_graph():
RunStatus.error,
error="Thread requires delta mode; materialize and convert its checkpoints before using full mode.",
)
@pytest.mark.asyncio
async def test_start_run_rejects_invalid_thread_id_before_resolving_dependencies():
from fastapi import HTTPException
from app.gateway.run_models import RunCreateRequest
from app.gateway.services import start_run
with pytest.raises(HTTPException) as exc_info:
await start_run(RunCreateRequest(), "thread.with.dot", SimpleNamespace())
assert exc_info.value.status_code == 422
assert "Invalid thread_id" in exc_info.value.detail

View File

@ -15,6 +15,33 @@ from blockbuster import BlockBuster
from kubernetes.client.rest import ApiException
def test_provisioner_thread_id_pattern_matches_gateway_contract(provisioner_module) -> None:
from deerflow.utils.thread_id import THREAD_ID_PATTERN
assert provisioner_module.SAFE_THREAD_ID_PATTERN == THREAD_ID_PATTERN
@pytest.mark.parametrize("thread_id", ["", "thread.with.dot", "../escape", "x" * 65])
def test_provisioner_rejects_noncanonical_thread_ids(provisioner_module, thread_id: str) -> None:
from pydantic import ValidationError
with pytest.raises(ValidationError):
provisioner_module.CreateSandboxRequest(
sandbox_id="sandbox-validation",
thread_id=thread_id,
)
@pytest.mark.parametrize("thread_id", ["a", "A1_b-2", "x" * 64])
def test_provisioner_accepts_canonical_thread_ids(provisioner_module, thread_id: str) -> None:
request = provisioner_module.CreateSandboxRequest(
sandbox_id="sandbox-validation",
thread_id=thread_id,
)
assert request.thread_id == thread_id
class _RecordingCoreV1:
def __init__(
self,

View File

@ -725,6 +725,20 @@ class TestMakeRunEventStore:
class TestJsonlRunEventStore:
@pytest.mark.anyio
@pytest.mark.parametrize("thread_id", ["", "thread.with.dot", "../escape", "x" * 65])
async def test_rejects_noncanonical_thread_ids(self, tmp_path, thread_id):
from deerflow.runtime.events.store.jsonl import JsonlRunEventStore
store = JsonlRunEventStore(base_dir=tmp_path / "jsonl")
with pytest.raises(ValueError, match="Invalid thread_id"):
await store.put(
thread_id=thread_id,
run_id="r1",
event_type="human_message",
category="message",
)
@pytest.mark.anyio
async def test_basic_crud(self, tmp_path):
from deerflow.runtime.events.store.jsonl import JsonlRunEventStore

View File

@ -324,6 +324,19 @@ def test_resolve_thread_id_handles_null_configurable():
assert runs._resolve_thread_id(RunCreateRequest(config={"configurable": {"thread_id": "t1"}})) == "t1"
@pytest.mark.parametrize(
"thread_id",
["", "thread.with.dot", "../escape", "x" * 65, 123],
)
def test_run_request_rejects_invalid_configurable_thread_id(thread_id):
from pydantic import ValidationError
from app.gateway.run_models import RunCreateRequest
with pytest.raises(ValidationError):
RunCreateRequest(config={"configurable": {"thread_id": thread_id}})
def test_build_run_config_handles_null_configurable():
"""A null ``configurable`` must also survive ``build_run_config``.

View File

@ -7,6 +7,31 @@ import pytest
from app.gateway.routers import scheduled_tasks
@pytest.mark.parametrize(
"model",
[
scheduled_tasks.ScheduledTaskCreateRequest,
scheduled_tasks.ScheduledTaskUpdateRequest,
],
)
@pytest.mark.parametrize("thread_id", ["", "thread.with.dot", "../escape", "x" * 65])
def test_scheduled_task_models_reject_invalid_thread_ids(model, thread_id):
from pydantic import ValidationError
kwargs = {"thread_id": thread_id}
if model is scheduled_tasks.ScheduledTaskCreateRequest:
kwargs.update(
title="Task",
prompt="Prompt",
schedule_type="cron",
schedule_spec={"cron": "0 * * * *"},
timezone="UTC",
)
with pytest.raises(ValidationError):
model(**kwargs)
class _Repo:
def __init__(self) -> None:
self.created = []

View File

@ -278,6 +278,112 @@ async def test_manual_overlap_conflict_returns_conflict():
assert run_repo.updated[-1][1]["status"] == "failed"
@pytest.mark.asyncio
async def test_dispatch_task_records_failure_for_legacy_invalid_thread_id():
"""Rows persisted before the thread-id contract was centralized may store
IDs that fail the canonical pattern (dots, >64 chars). Dispatch must record
the failure through normal bookkeeping instead of raising an uncaught
ValueError surfaces as HTTP 500 on manual trigger and, in the poller,
aborts the rest of the claimed batch every cycle."""
async def fake_launch(**_kwargs):
raise AssertionError("launch_run must not be called for an invalid thread_id")
task_repo = DummyTaskRepo(
[
{
"id": "task-legacy",
"user_id": "user-1",
"thread_id": "thread.with.dot",
"context_mode": "reuse_thread",
"assistant_id": "lead_agent",
"prompt": "Summarize thread",
"schedule_type": "cron",
"schedule_spec": {"cron": "0 9 * * *"},
"timezone": "UTC",
"status": "enabled",
}
]
)
run_repo = DummyRunRepo()
service = ScheduledTaskService(
task_repo=task_repo,
task_run_repo=run_repo,
launch_run=fake_launch,
poll_interval_seconds=5,
lease_seconds=120,
max_concurrent_runs=3,
)
result = await service.dispatch_task(
task_repo.rows[0],
now=datetime.now(UTC),
trigger="scheduled",
)
assert result["outcome"] == "failed"
assert result["task_run_id"] is None
assert result["run_id"] is None
assert "Invalid thread_id" in result["error"]
assert run_repo.created is None
assert task_repo.updated[1]["last_error"] == result["error"]
assert task_repo.updated[1]["last_thread_id"] == "thread.with.dot"
assert task_repo.updated[1]["increment_run_count"] is False
@pytest.mark.asyncio
async def test_run_once_continues_batch_after_invalid_thread_id():
"""A poison legacy row must not prevent later claimed tasks from dispatching."""
launched = []
async def fake_launch(**kwargs):
launched.append(kwargs)
return {"run_id": "run-ok", "thread_id": kwargs["thread_id"]}
task_repo = DummyTaskRepo(
[
{
"id": "task-legacy",
"user_id": "user-1",
"thread_id": "thread.with.dot",
"context_mode": "reuse_thread",
"assistant_id": "lead_agent",
"prompt": "Summarize thread",
"schedule_type": "cron",
"schedule_spec": {"cron": "0 9 * * *"},
"timezone": "UTC",
"status": "enabled",
},
{
"id": "task-valid",
"user_id": "user-1",
"thread_id": "thread-ok",
"context_mode": "reuse_thread",
"assistant_id": "lead_agent",
"prompt": "Summarize thread",
"schedule_type": "cron",
"schedule_spec": {"cron": "0 9 * * *"},
"timezone": "UTC",
"status": "enabled",
},
]
)
run_repo = DummyRunRepo()
service = ScheduledTaskService(
task_repo=task_repo,
task_run_repo=run_repo,
launch_run=fake_launch,
poll_interval_seconds=5,
lease_seconds=120,
max_concurrent_runs=3,
)
await service.run_once(now=datetime.now(UTC))
assert len(launched) == 1
assert launched[0]["thread_id"] == "thread-ok"
@pytest.mark.asyncio
async def test_handle_run_completion_persists_success():
task_repo = DummyTaskRepo(

View File

@ -715,7 +715,17 @@ def test_validate_thread_id_rejects_dot_traversal(thread_id):
def test_validate_thread_id_accepts_safe_ids():
support_bundle._validate_thread_id("thread-123")
support_bundle._validate_thread_id("a.b_c-1")
support_bundle._validate_thread_id("ab_c-1")
def test_validate_thread_id_rejects_noncanonical_ids():
"""The script's pattern is pinned byte-identical to the canonical
``THREAD_ID_PATTERN`` (see test_thread_id_validation.py) dotted or
over-length IDs are rejected even though they are not traversals."""
with pytest.raises(ValueError, match="Invalid thread_id"):
support_bundle._validate_thread_id("a.b_c-1")
with pytest.raises(ValueError, match="Invalid thread_id"):
support_bundle._validate_thread_id("x" * 65)
def test_main_reports_invalid_thread_id_without_traceback(tmp_path, capsys):

View File

@ -0,0 +1,162 @@
"""Canonical thread ID contract: every HTTP route and embedded-client entry
point that takes a ``thread_id`` must enforce ``ThreadId`` validation.
Two complementary guards:
1. A static sweep (AST over ``app/gateway/routers/*.py``) asserting every
route handler that declares a ``thread_id`` parameter annotates it
``ThreadId`` this is what prevents new routes from silently landing
with a raw ``str`` again (the suggestions/thread_runs/threads gaps).
2. A runtime sweep hitting every ``{thread_id}`` route with a non-canonical
ID and asserting a 422 whose error location names ``thread_id``.
Deliberate exceptions (RFC #4588):
- ``DELETE /api/threads/{thread_id}`` keeps ``thread_id: str`` as the
legacy-cleanup escape hatch.
- The browser websocket stream validates on upgrade; covered separately.
"""
from __future__ import annotations
import ast
import re
from pathlib import Path
import pytest
from _router_auth_helpers import make_authed_test_app
from fastapi.testclient import TestClient
ROUTERS_DIR = Path(__file__).resolve().parent.parent / "app" / "gateway" / "routers"
# (handler name) route handlers deliberately allowed to keep ``thread_id: str``.
STATIC_WHITELIST = {"delete_thread_data"}
# (method, path) routes deliberately excluded from the runtime 422 sweep.
RUNTIME_WHITELIST = {
("DELETE", "/api/threads/{thread_id}"), # legacy-cleanup escape hatch
}
BAD_THREAD_ID = "bad.thread.id"
_ROUTE_DECORATOR_RE = re.compile(r"router\.(get|post|delete|put|patch|websocket)")
def _iter_route_handlers(path: Path):
"""Yield (handler_name, has_thread_id_param, annotation) for route handlers."""
tree = ast.parse(path.read_text())
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
if not any(isinstance(dec, ast.Call) and isinstance(dec.func, ast.Attribute) and isinstance(dec.func.value, ast.Name) and _ROUTE_DECORATOR_RE.fullmatch(f"{dec.func.value.id}.{dec.func.attr}") for dec in node.decorator_list):
continue
for arg in (*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs):
if arg.arg == "thread_id":
annotation = ast.unparse(arg.annotation) if arg.annotation else None
yield node.name, annotation
def test_every_thread_id_route_handler_uses_canonical_type():
"""Static guard: no route handler may declare a bare ``str`` thread_id."""
violations = []
for path in sorted(ROUTERS_DIR.glob("*.py")):
for handler, annotation in _iter_route_handlers(path):
if handler in STATIC_WHITELIST:
continue
if annotation != "ThreadId":
violations.append(f"{path.name}:{handler} -> {annotation!r}")
assert not violations, "route handlers with non-canonical thread_id:\n" + "\n".join(violations)
def _collect_thread_id_routes():
"""Import every gateway router and collect (method, full_path) with {thread_id}."""
from app.gateway.routers import (
artifacts,
browser,
feedback,
runs,
scheduled_tasks,
skills,
suggestions,
thread_runs,
threads,
uploads,
)
routers = [artifacts, browser, feedback, runs, scheduled_tasks, skills, suggestions, thread_runs, threads, uploads]
cases = []
for module in routers:
for route in module.router.routes:
path = getattr(route, "path", "")
if "{thread_id}" not in path:
continue
methods = getattr(route, "methods", None)
if methods is None:
continue # websocket routes — covered by the dedicated test below
for method in sorted(methods):
if (method, path) in RUNTIME_WHITELIST:
continue
cases.append((module.__name__.rsplit(".", 1)[-1], method, path))
return cases
_THREAD_ID_ROUTES = _collect_thread_id_routes()
def test_browser_websocket_rejects_noncanonical_thread_id():
"""The browser stream websocket validates thread_id on upgrade."""
from starlette.websockets import WebSocketDisconnect
from app.gateway.routers import browser
app = make_authed_test_app()
app.include_router(browser.router)
with TestClient(app, raise_server_exceptions=False) as client:
with pytest.raises(WebSocketDisconnect):
with client.websocket_connect(f"/api/threads/{BAD_THREAD_ID}/browser/stream"):
pass
def test_sweep_covers_expected_surface():
"""Sanity: the sweep must actually see the known thread_id routes."""
assert len(_THREAD_ID_ROUTES) >= 30
assert any("suggestions" in name for name, _, _ in _THREAD_ID_ROUTES)
@pytest.mark.parametrize(
("router_name", "method", "path"),
_THREAD_ID_ROUTES,
ids=[f"{name}:{method}:{path}" for name, method, path in _THREAD_ID_ROUTES],
)
def test_noncanonical_thread_id_gets_422(router_name, method, path):
"""Runtime guard: a non-canonical thread_id yields 422 naming thread_id."""
import importlib
from unittest.mock import MagicMock
from app.gateway.deps import get_config
module = importlib.import_module(f"app.gateway.routers.{router_name}")
app = make_authed_test_app()
app.include_router(module.router)
# get_config 503s when no config.yaml exists (CI), and dependency solving
# precedes path-param validation — override it so the 422 contract is
# exercised regardless of the environment.
app.dependency_overrides[get_config] = MagicMock()
url = path.replace("{thread_id}", BAD_THREAD_ID)
# Other path params get a harmless canonical placeholder.
url = re.sub(r"\{(\w+)(?::path)?\}", "x", url)
with TestClient(app, raise_server_exceptions=False) as client:
if method == "GET":
response = client.get(url)
elif method == "DELETE":
response = client.delete(url)
elif method in {"POST", "PUT", "PATCH"}:
response = client.request(method, url, json={})
else: # WEBSOCKET etc. — not expected in the sweep
pytest.skip(f"unsupported method {method}")
assert response.status_code == 422, f"{method} {path} -> {response.status_code}: {response.text[:300]}"
detail = response.json()["detail"]
assert any("thread_id" in str(err.get("loc", ())) for err in detail), f"422 did not name thread_id: {detail}"

View File

@ -0,0 +1,90 @@
from __future__ import annotations
import pytest
@pytest.mark.parametrize(
"thread_id",
[
"a",
"A1_b-2",
"x" * 64,
],
)
def test_validate_thread_id_accepts_canonical_ids(thread_id: str) -> None:
from deerflow.utils.thread_id import validate_thread_id
assert validate_thread_id(thread_id) == thread_id
@pytest.mark.parametrize(
"thread_id",
[
"",
"x" * 65,
"thread.with.dot",
"../escape",
"has space",
"line\nbreak",
"线程",
],
)
def test_validate_thread_id_rejects_noncanonical_ids(thread_id: str) -> None:
from deerflow.utils.thread_id import validate_thread_id
with pytest.raises(ValueError, match="Invalid thread_id"):
validate_thread_id(thread_id)
@pytest.mark.parametrize("thread_id", [1, {}, []])
def test_validate_thread_id_rejects_non_strings(thread_id: object) -> None:
from deerflow.utils.thread_id import validate_thread_id
with pytest.raises(ValueError, match="Invalid thread_id"):
validate_thread_id(thread_id) # type: ignore[arg-type]
@pytest.mark.parametrize("method_name", ["upload_files", "list_uploads", "delete_upload", "get_artifact"])
def test_client_mutating_and_read_entry_points_validate_thread_id(method_name: str) -> None:
"""RFC #4588: the embedded client validates on all mutating entry points.
Validation is the first statement of each method, so a bare ``__new__``
instance is enough no config, sandbox, or event loop is touched.
"""
from deerflow.client import DeerFlowClient
client = DeerFlowClient.__new__(DeerFlowClient)
method = getattr(client, method_name)
args = {"upload_files": (["x"],), "list_uploads": (), "delete_upload": ("f.txt",), "get_artifact": ("mnt/user-data/outputs/f.txt",)}[method_name]
with pytest.raises(ValueError, match="Invalid thread_id"):
method("bad.thread.id", *args)
def test_support_bundle_thread_id_pattern_matches_canonical() -> None:
"""scripts/support_bundle.py keeps its own copy of the pattern (it must
run with a broken venv) pin it byte-identical to the canonical one."""
import importlib.util
from pathlib import Path
from deerflow.utils.thread_id import THREAD_ID_PATTERN
script = Path(__file__).resolve().parents[2] / "scripts" / "support_bundle.py"
spec = importlib.util.spec_from_file_location("support_bundle", script)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
assert module.SAFE_THREAD_ID_RE.pattern == THREAD_ID_PATTERN
def test_tui_resume_literal_ref_validated() -> None:
"""The TUI /resume fallback adopts a literal ref as thread id only when
it satisfies the canonical contract."""
from deerflow.tui.session import Session
session = Session.__new__(Session)
session.client = type("StubClient", (), {"list_threads": lambda self, limit: {"thread_list": []}})()
assert session.resolve_ref("valid-id_1") == "valid-id_1"
with pytest.raises(ValueError, match="not a valid thread id"):
session.resolve_ref("bad.thread.id")

View File

@ -410,18 +410,66 @@ def test_delete_thread_route_rejects_invalid_thread_id(tmp_path):
assert response.status_code == 404
def test_delete_thread_route_returns_422_for_route_safe_invalid_id(tmp_path):
paths = Paths(tmp_path)
def test_delete_thread_route_cleans_legacy_metadata_without_resolving_unsafe_path():
app, store, _checkpointer = _build_thread_app()
legacy_thread_id = "legacy.thread"
app = make_authed_test_app()
app.include_router(threads.router)
asyncio.run(
store.aput(
THREADS_NS,
legacy_thread_id,
{
"thread_id": legacy_thread_id,
"status": "idle",
"created_at": "",
"updated_at": "",
"metadata": {},
},
)
)
with patch("app.gateway.routers.threads.get_paths", return_value=paths):
with TestClient(app) as client:
response = client.delete("/api/threads/thread.with.dot")
with (
patch(
"app.gateway.routers.threads._delete_thread_data",
side_effect=AssertionError("legacy thread ID must not reach filesystem cleanup"),
),
TestClient(app) as client,
):
response = client.delete(f"/api/threads/{legacy_thread_id}")
assert response.status_code == 200
assert "Skipped local data cleanup" in response.json()["message"]
assert asyncio.run(store.aget(THREADS_NS, legacy_thread_id)) is None
def test_legacy_thread_metadata_mutation_is_rejected():
app, store, _checkpointer = _build_thread_app()
legacy_thread_id = "legacy.thread"
asyncio.run(
store.aput(
THREADS_NS,
legacy_thread_id,
{
"thread_id": legacy_thread_id,
"status": "idle",
"created_at": "",
"updated_at": "",
"metadata": {"original": True},
},
)
)
with TestClient(app) as client:
response = client.patch(
f"/api/threads/{legacy_thread_id}",
json={"metadata": {"mutated": True}},
)
assert response.status_code == 422
assert "Invalid thread_id" in response.json()["detail"]
record = asyncio.run(store.aget(THREADS_NS, legacy_thread_id))
assert record is not None
assert record.value["metadata"] == {"original": True}
def test_delete_thread_data_returns_generic_500_error(tmp_path):
@ -490,6 +538,31 @@ def test_create_thread_returns_iso_timestamps() -> None:
assert body["created_at"] == body["updated_at"]
@pytest.mark.parametrize(
"thread_id",
["", "thread.with.dot", "../escape", "x" * 65],
)
def test_create_thread_rejects_invalid_explicit_thread_id_before_persistence(thread_id: str) -> None:
app, store, checkpointer = _build_thread_app()
with TestClient(app) as client:
response = client.post("/api/threads", json={"thread_id": thread_id})
assert response.status_code == 422
assert asyncio.run(store.asearch(THREADS_NS)) == []
assert not checkpointer.storage
def test_create_thread_preserves_valid_explicit_thread_id() -> None:
app, _store, _checkpointer = _build_thread_app()
with TestClient(app) as client:
response = client.post("/api/threads", json={"thread_id": "caller_thread-1"})
assert response.status_code == 200
assert response.json()["thread_id"] == "caller_thread-1"
def test_create_thread_returns_existing_when_insert_loses_race() -> None:
"""A concurrent create that loses the INSERT race stays idempotent.

View File

@ -100,7 +100,7 @@ if not (1 <= SANDBOX_CONTAINER_PORT <= 65535):
raise RuntimeError(f"Invalid SANDBOX_CONTAINER_PORT={SANDBOX_CONTAINER_PORT}; expected a value in [1, 65535]")
if SANDBOX_SERVICE_TYPE not in {"NodePort", "ClusterIP"}:
raise RuntimeError(f"Invalid SANDBOX_SERVICE_TYPE={SANDBOX_SERVICE_TYPE!r}; expected 'NodePort' or 'ClusterIP'")
SAFE_THREAD_ID_PATTERN = r"^[A-Za-z0-9_\-]+$"
SAFE_THREAD_ID_PATTERN = r"^[A-Za-z0-9_-]{1,64}$"
SAFE_USER_ID_PATTERN = r"^[A-Za-z0-9_\-]+$"
DEFAULT_USER_ID = "default"
MAX_EXTRA_MOUNTS = 9

View File

@ -63,7 +63,11 @@ SECRET_FLAG_RE = re.compile(r"(?i)^--?[\w.-]*(?:api[_-]?key|token|secret|passwor
HEADER_KEY_RE = re.compile(r"(?i)header")
POSIX_HOME_RE = re.compile(r"(?<![\w.-])(/Users|/home)/([^/\s:]+)")
WINDOWS_HOME_RE = re.compile(r"(?i)([A-Z]:\\Users\\)([^\\\s:]+)")
SAFE_THREAD_ID_RE = re.compile(r"^[a-zA-Z0-9._-]+$")
# Must stay byte-identical to deerflow.utils.thread_id.THREAD_ID_PATTERN
# (canonical thread ID contract); pinned by a parity test in backend/tests.
# Kept as a local copy because this script must run even in environments
# where the backend venv is broken.
SAFE_THREAD_ID_RE = re.compile(r"^[A-Za-z0-9_-]{1,64}$")
DOCTOR_STATUS_RE = re.compile(r"Status:\s*(\d+)\s+error\(s\),\s*(\d+)\s+warning\(s\)", re.IGNORECASE)
ATTENTION_SIGNAL_NAMES = {
"doctor_failed",