Compare commits

..
Author SHA1 Message Date
hjlarry 138dbe6d0d fix(web): localize workflow archive logs 2026-07-09 15:15:39 +08:00
hjlarry 571bff9796 Merge remote-tracking branch 'myori/main' into p441
# Conflicts:
#	api/openapi/markdown/console-openapi.md
#	packages/contracts/generated/api/console/orpc.gen.ts
#	web/app/components/header/account-setting/index.tsx
#	web/app/components/main-nav/components/account-section.tsx
#	web/contract/router.ts
2026-07-09 15:07:22 +08:00
hjlarry 68763928f2 fix(web): include workflow archive contract 2026-07-01 09:45:37 +08:00
hjlarry ed3ce357ca fix(web): reorder main nav account classes 2026-07-01 09:24:22 +08:00
autofix-ci[bot]andGitHub 747b9582fa [autofix.ci] apply automated fixes 2026-07-01 01:12:26 +00:00
hjlarry 0d3a9a6de1 Merge remote-tracking branch 'myori/main' into p441
# Conflicts:
#	packages/contracts/generated/api/console/orpc.gen.ts
2026-07-01 09:07:22 +08:00
autofix-ci[bot]andGitHub 8ae8897a27 [autofix.ci] apply automated fixes 2026-06-30 09:51:23 +00:00
hjlarry e1daa6d8cc fix migration version 2026-06-30 17:41:52 +08:00
hjlarry bb7c16e3d0 fix CI 2026-06-30 17:41:42 +08:00
hjlarry 3fa24cac6a feat: only paid user can see the archive notice 2026-06-30 17:27:19 +08:00
hjlarry 187436a167 feat: sandbox upgrade hint 2026-06-30 17:06:08 +08:00
hjlarry fde2042cfe fix: archived logs only saas and admin can view 2026-06-30 16:48:47 +08:00
hjlarry add719a6e1 Merge remote-tracking branch 'myori/main' into p441 2026-06-30 16:22:01 +08:00
hjlarry c7feb18446 feat: add archived log notice 2026-06-30 10:58:57 +08:00
hjlarry 6b13371c8d fix: add more upsert archive bundle 2026-06-30 09:51:41 +08:00
hjlarry 0e47bf1c8f fix CI 2026-06-28 15:01:44 +08:00
hjlarry 2848c183ad Merge remote-tracking branch 'myori/main' into p441 2026-06-28 15:01:28 +08:00
hjlarry ca2443b4f6 feat: use seperate export bucket 2026-06-28 14:56:08 +08:00
hjlarry b4f89fe36c fix: display fail reason 2026-06-27 22:30:46 +08:00
hjlarry 6bd7cbea24 feat: add a fake pagnition 2026-06-27 22:10:06 +08:00
hjlarry 47a2524f5b fix: improve display 2026-06-27 07:01:00 +08:00
hjlarry 3d3b94c75d fix: improve download 2026-06-26 23:30:02 +08:00
hjlarry cc9d95773a feat: use csv for exports 2026-06-26 22:48:55 +08:00
hjlarry 7308304fe0 feat: add celery download task 2026-06-26 10:43:52 +08:00
hjlarry cc99e4ac57 fix: improve prepare download display 2026-06-26 09:42:48 +08:00
hjlarry 8b23d2e71a feat: add archive log api 2026-06-26 08:54:15 +08:00
hjlarry 76d93e04ce feat: add backfill archive logs command 2026-06-25 17:05:07 +08:00
hjlarry 42a32a0e0c feat: add archive logs basic backend 2026-06-25 15:31:17 +08:00
hjlarry 02fbb36785 feat: add frontend of archived workflow logs 2026-06-25 13:48:43 +08:00
613 changed files with 5978 additions and 15253 deletions
+2
View File
@@ -125,6 +125,8 @@ All of Dify's offerings come with corresponding APIs, so you could effortlessly
- **Dify for enterprise / organizations<br/>**
We provide additional enterprise-centric features. [Send us an email](mailto:[email protected]?subject=%5BGitHub%5DBusiness%20License%20Inquiry) to discuss your enterprise needs. <br/>
> For startups and small businesses using AWS, check out [Dify Premium on AWS Marketplace](https://aws.amazon.com/marketplace/pp/prodview-t22mebxzwjhu6) and deploy it to your own AWS VPC with one click. It's an affordable AMI offering with the option to create apps with custom logo and branding.
## Staying ahead
Star Dify on GitHub and be instantly notified of new releases.
-3
View File
@@ -663,9 +663,6 @@ PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600
PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400
INNER_API_KEY_FOR_PLUGIN=QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1
# Dify Agent backend
AGENT_BACKEND_BASE_URL=http://localhost:5050
# Marketplace configuration
MARKETPLACE_ENABLED=true
MARKETPLACE_API_URL=https://marketplace.dify.ai
+2
View File
@@ -26,6 +26,7 @@ from .rbac import migrate_dataset_permissions_to_rbac, migrate_member_roles_to_r
from .retention import (
archive_workflow_runs,
archive_workflow_runs_plan,
backfill_workflow_run_archive_bundles,
clean_expired_messages,
clean_workflow_runs,
cleanup_orphaned_draft_variables,
@@ -54,6 +55,7 @@ __all__ = [
"archive_workflow_runs",
"archive_workflow_runs_plan",
"backfill_plugin_auto_upgrade",
"backfill_workflow_run_archive_bundles",
"clean_expired_messages",
"clean_workflow_runs",
"cleanup_orphaned_draft_variables",
+121 -122
View File
@@ -1,12 +1,10 @@
import datetime
import logging
import time
from collections.abc import Callable
from typing import TypedDict
import click
import sqlalchemy as sa
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_database import db
from libs.datetime_utils import naive_utc_now
@@ -14,7 +12,6 @@ from services.clear_free_plan_tenant_expired_logs import ClearFreePlanTenantExpi
from services.retention.conversation.messages_clean_policy import create_message_clean_policy
from services.retention.conversation.messages_clean_service import MessagesCleanService
from services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs import WorkflowRunCleanup
from services.retention.workflow_run.db_retry import run_with_db_retry
from services.retention.workflow_run.tenant_prefix import tenant_prefix_condition
from tasks.remove_app_and_related_data_task import delete_draft_variables_batch
@@ -38,12 +35,6 @@ class WorkflowRunArchiveTenantPlan(TypedDict):
unpaid_tenant_ids: list[str]
class WorkflowRunArchivePrefixStats(TypedDict):
tenant_ids: list[str]
workflow_runs: int
workflow_node_executions: int
def _normalize_utc_datetime(value: datetime.datetime) -> datetime.datetime:
if value.tzinfo is None:
return value.replace(tzinfo=datetime.UTC)
@@ -65,8 +56,16 @@ def _parse_tenant_prefixes(prefixes: str | None) -> list[str]:
return sorted(set(parsed))
def _parse_comma_separated_ids(raw_ids: str | None, *, param_name: str) -> list[str] | None:
if raw_ids is None:
return None
parsed = sorted({raw_id.strip() for raw_id in raw_ids.split(",") if raw_id.strip()})
if not parsed:
raise click.BadParameter(f"{param_name} must not be empty")
return parsed
def _get_archive_candidate_tenant_ids_by_prefix(
session: Session,
prefix: str,
*,
start_from: datetime.datetime | None,
@@ -85,7 +84,7 @@ def _get_archive_candidate_tenant_ids_by_prefix(
if start_from is not None:
conditions.append(WorkflowRun.created_at >= start_from)
tenant_ids = session.scalars(
tenant_ids = db.session.scalars(
sa.select(WorkflowRun.tenant_id).where(*conditions).distinct().order_by(WorkflowRun.tenant_id)
).all()
return list(tenant_ids)
@@ -112,80 +111,8 @@ def _filter_paid_workflow_archive_tenant_ids(tenant_ids: list[str]) -> tuple[lis
return paid_tenant_ids, unpaid_tenant_ids
def _run_archive_command_db_retry[T](operation_name: str, operation: Callable[[], T]) -> T:
return run_with_db_retry(operation_name, operation, logger=logger)
def _get_archive_candidate_tenant_ids_with_retry(
session_maker: sessionmaker[Session],
prefix: str,
*,
start_from: datetime.datetime | None,
end_before: datetime.datetime,
) -> list[str]:
def fetch_tenant_ids() -> list[str]:
with session_maker() as session:
return _get_archive_candidate_tenant_ids_by_prefix(
session,
prefix,
start_from=start_from,
end_before=end_before,
)
return _run_archive_command_db_retry(f"workflow archive tenant resolve for prefix {prefix}", fetch_tenant_ids)
def _get_archive_plan_prefix_stats(
session_maker: sessionmaker[Session],
prefix: str,
*,
start_from: datetime.datetime | None,
end_before: datetime.datetime,
) -> WorkflowRunArchivePrefixStats:
from graphon.enums import WorkflowExecutionStatus
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun
from services.retention.workflow_run.archive_paid_plan_workflow_run import WorkflowRunArchiver
def fetch_prefix_stats() -> WorkflowRunArchivePrefixStats:
with session_maker() as session:
tenant_ids = _get_archive_candidate_tenant_ids_by_prefix(
session,
prefix,
start_from=start_from,
end_before=end_before,
)
run_conditions = [
WorkflowRun.created_at < end_before,
WorkflowRun.status.in_(WorkflowExecutionStatus.ended_values()),
WorkflowRun.type.in_(WorkflowRunArchiver.ARCHIVED_TYPE),
tenant_prefix_condition(WorkflowRun.tenant_id, prefix),
]
if start_from is not None:
run_conditions.append(WorkflowRun.created_at >= start_from)
workflow_runs = (
session.scalar(sa.select(sa.func.count()).select_from(WorkflowRun).where(*run_conditions)) or 0
)
candidate_runs = sa.select(WorkflowRun.id).where(*run_conditions).subquery()
workflow_node_executions = (
session.scalar(
sa.select(sa.func.count())
.select_from(WorkflowNodeExecutionModel)
.join(candidate_runs, WorkflowNodeExecutionModel.workflow_run_id == candidate_runs.c.id)
)
or 0
)
return WorkflowRunArchivePrefixStats(
tenant_ids=tenant_ids,
workflow_runs=workflow_runs,
workflow_node_executions=workflow_node_executions,
)
return _run_archive_command_db_retry(f"workflow archive plan for prefix {prefix}", fetch_prefix_stats)
def _resolve_archive_tenant_ids_from_plan(
*,
session_maker: sessionmaker[Session],
tenant_ids: str | None,
tenant_prefixes: list[str],
start_from: datetime.datetime | None,
@@ -204,8 +131,7 @@ def _resolve_archive_tenant_ids_from_plan(
requested_tenant_ids = []
for prefix in tenant_prefixes:
requested_tenant_ids.extend(
_get_archive_candidate_tenant_ids_with_retry(
session_maker,
_get_archive_candidate_tenant_ids_by_prefix(
prefix,
start_from=start_from,
end_before=end_before,
@@ -226,21 +152,6 @@ def _resolve_archive_tenant_ids_from_plan(
)
def _safe_remove_scoped_session(context: str) -> None:
try:
db.session.remove()
except Exception:
logger.warning("Ignoring DB scoped-session cleanup error after %s", context, exc_info=True)
try:
db.session.registry.clear()
except Exception:
logger.warning("Ignoring DB scoped-session registry cleanup error after %s", context, exc_info=True)
try:
db.engine.dispose()
except Exception:
logger.warning("Ignoring DB engine dispose error after %s", context, exc_info=True)
def _resolve_archive_time_range(
*,
before_days: int,
@@ -447,6 +358,10 @@ def archive_workflow_runs_plan(
supported workflow types, and the requested created_at window. V2 bundle archive
does not maintain per-run archive logs, so this plan reports source-table volume.
"""
from graphon.enums import WorkflowExecutionStatus
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun
from services.retention.workflow_run.archive_paid_plan_workflow_run import WorkflowRunArchiver
before_days, start_from, end_before = _resolve_archive_time_range(
before_days=before_days,
from_days_ago=from_days_ago,
@@ -458,25 +373,37 @@ def archive_workflow_runs_plan(
if include_archived:
click.echo(click.style("--include-archived is a no-op for V2 bundle archive plans.", fg="yellow"))
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
rows: list[WorkflowRunArchivePlanRow] = []
for prefix in _HEX_PREFIXES:
try:
prefix_stats = _get_archive_plan_prefix_stats(
session_maker,
prefix,
start_from=start_from,
end_before=plan_end_before,
)
except Exception as exc:
logger.exception("Failed to build workflow archive plan for prefix %s", prefix)
raise click.ClickException(f"Failed to build workflow archive plan for prefix {prefix}.") from exc
tenant_ids = prefix_stats["tenant_ids"]
workflow_runs = prefix_stats["workflow_runs"]
workflow_node_executions = prefix_stats["workflow_node_executions"]
tenant_ids = _get_archive_candidate_tenant_ids_by_prefix(
prefix,
start_from=start_from,
end_before=plan_end_before,
)
total_tenants = len(tenant_ids)
paid_tenant_ids, unpaid_tenant_ids = _filter_paid_workflow_archive_tenant_ids(tenant_ids)
run_conditions = [
WorkflowRun.created_at < plan_end_before,
WorkflowRun.status.in_(WorkflowExecutionStatus.ended_values()),
WorkflowRun.type.in_(WorkflowRunArchiver.ARCHIVED_TYPE),
tenant_prefix_condition(WorkflowRun.tenant_id, prefix),
]
if start_from is not None:
run_conditions.append(WorkflowRun.created_at >= start_from)
workflow_runs = (
db.session.scalar(sa.select(sa.func.count()).select_from(WorkflowRun).where(*run_conditions)) or 0
)
candidate_runs = sa.select(WorkflowRun.id).where(*run_conditions).subquery()
workflow_node_executions = (
db.session.scalar(
sa.select(sa.func.count())
.select_from(WorkflowNodeExecutionModel)
.join(candidate_runs, WorkflowNodeExecutionModel.workflow_run_id == candidate_runs.c.id)
)
or 0
)
rows.append(
WorkflowRunArchivePlanRow(
tenant_prefix=prefix,
@@ -656,18 +583,17 @@ def archive_workflow_runs(
)
)
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
try:
tenant_plan = _resolve_archive_tenant_ids_from_plan(
session_maker=session_maker,
tenant_ids=tenant_ids,
tenant_prefixes=parsed_tenant_prefixes,
start_from=start_from,
end_before=plan_end_before,
)
except Exception as exc:
except Exception:
logger.exception("Failed to resolve workflow archive tenant plan")
raise click.ClickException("Failed to resolve workflow archive tenant plan.") from exc
click.echo(click.style("Failed to resolve workflow archive tenant plan.", fg="red"))
return
planned_tenant_ids = tenant_plan["archive_tenant_ids"]
planned_paid_tenant_ids = tenant_plan["paid_tenant_ids"] if planned_tenant_ids is not None else None
@@ -699,10 +625,7 @@ def archive_workflow_runs(
dry_run=dry_run,
delete_after_archive=delete_after_archive,
)
try:
summary = archiver.run()
finally:
_safe_remove_scoped_session("archive workflow run command")
summary = archiver.run()
click.echo(
click.style(
f"Summary: processed={summary.total_runs_processed}, archived={summary.runs_archived}, "
@@ -725,6 +648,82 @@ def archive_workflow_runs(
)
@click.command(
"backfill-workflow-run-archive-bundles",
help="Backfill workflow-run archive bundle DB index from object-storage manifests.",
)
@click.option("--tenant-ids", default=None, help="Optional comma-separated tenant IDs.")
@click.option(
"--tenant-prefixes",
default=None,
help="Optional comma-separated tenant ID first hex digits, e.g. 0,1,a,f.",
)
@click.option("--year", default=None, type=click.IntRange(min=1, max=9999), help="Optional archive year filter.")
@click.option("--month", default=None, type=click.IntRange(min=1, max=12), help="Optional archive month filter.")
@click.option("--limit", default=None, type=click.IntRange(min=1), help="Maximum number of manifests to process.")
@click.option("--dry-run", is_flag=True, help="Preview without writing workflow_run_archive_bundles.")
def backfill_workflow_run_archive_bundles(
tenant_ids: str | None,
tenant_prefixes: str | None,
year: int | None,
month: int | None,
limit: int | None,
dry_run: bool,
) -> None:
"""
Reconcile `workflow_run_archive_bundles` from V2 archive manifests.
This command is meant for bootstrapping the listing/download index after deploy or repairing index drift. The R2
manifests remain the source of truth; this command only mirrors their query metadata into the database.
"""
from services.retention.workflow_run.archive_bundle_index import WorkflowRunArchiveBundleIndexBackfill
if tenant_ids and tenant_prefixes:
raise click.UsageError("Choose either --tenant-ids or --tenant-prefixes, not both.")
if month is not None and year is None:
raise click.UsageError("--month must be used with --year.")
parsed_tenant_ids = _parse_comma_separated_ids(tenant_ids, param_name="tenant-ids")
parsed_tenant_prefixes = _parse_tenant_prefixes(tenant_prefixes)
if not parsed_tenant_ids and not parsed_tenant_prefixes:
click.echo(
click.style(
"No tenant scope supplied; scanning the full workflow-runs/v2/ archive prefix.",
fg="yellow",
)
)
started_at = datetime.datetime.now(datetime.UTC)
click.echo(click.style(f"Starting archive bundle index backfill at {started_at.isoformat()}.", fg="white"))
backfill = WorkflowRunArchiveBundleIndexBackfill()
summary = backfill.run(
tenant_ids=parsed_tenant_ids,
tenant_prefixes=parsed_tenant_prefixes or None,
year=year,
month=month,
limit=limit,
dry_run=dry_run,
)
status = "completed with failures" if summary.bundles_failed else "completed successfully"
fg = "red" if summary.bundles_failed else "green"
action = "would_upsert" if dry_run else "upserted"
action_count = summary.bundles_processed if dry_run else summary.bundles_upserted
click.echo(
click.style(
f"Backfill {status}. manifests_found={summary.manifests_found} "
f"bundles_processed={summary.bundles_processed} {action}={action_count} "
f"bundles_failed={summary.bundles_failed} runs={summary.workflow_run_count} rows={summary.row_count} "
f"archive_bytes={summary.archive_bytes} duration={summary.elapsed_time:.2f}s",
fg=fg,
)
)
for error in summary.errors[:10]:
click.echo(click.style(f" failed {error}", fg="red"))
if len(summary.errors) > 10:
click.echo(click.style(f" ... and {len(summary.errors) - 10} more failures", fg="red"))
def _echo_bundle_archive_operation_summary(summary) -> None:
status = "completed successfully" if summary.bundles_failed == 0 else "completed with failures"
fg = "green" if summary.bundles_failed == 0 else "red"
+4 -3
View File
@@ -25,10 +25,11 @@ class AgentBackendConfig(BaseSettings):
AGENT_SHELL_ENABLED: bool = Field(
description=(
"Inject the dify.shell layer (sandboxed bash workspace) into Agent runs. "
"Requires the agent backend to be wired with a shellctl entrypoint before "
"shell-using Agent runs are executed."
"Requires the agent backend to be wired with a shellctl entrypoint; keep it "
"off until shellctl is deployed, otherwise every agent run that includes the "
"shell layer will fail."
),
default=True,
default=False,
)
AGENT_APP_TEXT_DELTA_DEBOUNCE_SECONDS: NonNegativeFloat = Field(
+1 -4
View File
@@ -363,10 +363,7 @@ class FileAccessConfig(BaseSettings):
INTERNAL_FILES_URL: str = Field(
description="Internal base URL for file access within Docker network,"
" used for plugin daemon and internal service communication."
" Explicit INTERNAL_FILES_URL takes precedence; otherwise SERVER_CONSOLE_API_URL is used,"
" then FILES_URL.",
validation_alias=AliasChoices("INTERNAL_FILES_URL", "SERVER_CONSOLE_API_URL"),
alias_priority=1,
" Falls back to FILES_URL if not specified.",
default="",
)
+2
View File
@@ -43,6 +43,7 @@ from . import (
setup,
spec,
version,
workflow_run_archive,
)
from .agent import composer as agent_composer
from .agent import roster as agent_roster
@@ -238,6 +239,7 @@ __all__ = [
"workflow_draft_variable",
"workflow_node_output_inspector",
"workflow_run",
"workflow_run_archive",
"workflow_statistic",
"workflow_trigger",
"workspace",
@@ -0,0 +1,221 @@
import datetime
from http import HTTPStatus
from flask import redirect
from flask_restx import Resource
from pydantic import BaseModel, Field
from werkzeug.exceptions import Conflict, NotFound
from controllers.common.fields import RedirectResponse
from controllers.common.schema import register_response_schema_models, register_schema_models
from controllers.console import console_ns
from controllers.console.wraps import (
RBACPermission,
RBACResourceScope,
account_initialization_required,
is_admin_or_owner_required,
rbac_permission_required,
setup_required,
)
from extensions.ext_database import db
from fields.base import ResponseModel
from libs.archive_storage import get_export_storage
from libs.helper import dump_response
from libs.login import current_account_with_tenant, login_required
from services.retention.workflow_run.archive_download_preparation import ARCHIVE_DOWNLOAD_MIME_TYPE
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
)
from services.retention.workflow_run.archive_log_service import (
WorkflowRunArchiveDownloadNotReadyError,
WorkflowRunArchiveDownloadTaskNotFoundError,
WorkflowRunArchiveNotFoundError,
create_workflow_run_archive_download_task,
get_ready_workflow_run_archive_download_task,
get_workflow_run_archive_download_task,
list_workflow_run_archives,
)
class WorkflowRunArchiveDownloadPayload(BaseModel):
"""Request body for preparing one monthly workflow-run archive download."""
year: int = Field(ge=1)
month: int = Field(ge=1, le=12)
class WorkflowRunArchiveSummaryResponse(ResponseModel):
archived_month_count: int
workflow_run_count: int
archive_bytes: int
latest_archived_at: datetime.datetime | None = None
class WorkflowRunArchiveDownloadTaskResponse(ResponseModel):
download_id: str
year: int
month: int
bundle_count: int
archive_bytes: int
status: WorkflowRunArchiveDownloadStatus
file_name: str | None = None
file_size_bytes: int | None = None
error: str | None = None
created_at: datetime.datetime
updated_at: datetime.datetime
expires_at: datetime.datetime
started_at: datetime.datetime | None = None
finished_at: datetime.datetime | None = None
class WorkflowRunArchiveMonthResponse(ResponseModel):
year: int
month: int
bundle_count: int
workflow_run_count: int
row_count: int
archive_bytes: int
latest_archived_at: datetime.datetime
download_task: WorkflowRunArchiveDownloadTaskResponse | None = None
class WorkflowRunArchiveListResponse(ResponseModel):
summary: WorkflowRunArchiveSummaryResponse
months: list[WorkflowRunArchiveMonthResponse]
register_schema_models(console_ns, WorkflowRunArchiveDownloadPayload)
register_response_schema_models(
console_ns,
WorkflowRunArchiveSummaryResponse,
WorkflowRunArchiveMonthResponse,
WorkflowRunArchiveListResponse,
WorkflowRunArchiveDownloadTaskResponse,
RedirectResponse,
)
def _current_ids() -> tuple[str, str]:
"""Return current `(tenant_id, account_id)` or raise when no workspace is selected."""
current_user, current_tenant_id = current_account_with_tenant()
if not current_tenant_id:
raise NotFound("Current workspace not found")
return current_tenant_id, current_user.id
def _presigned_url_expires_in(expires_at: datetime.datetime) -> int:
"""Keep the storage URL no longer-lived than the Redis task and cap it for browser downloads."""
expires_at_utc = expires_at if expires_at.tzinfo else expires_at.replace(tzinfo=datetime.UTC)
remaining_seconds = int((expires_at_utc - datetime.datetime.now(datetime.UTC)).total_seconds())
return max(1, min(3600, remaining_seconds))
@console_ns.route("/workflow-run-archives")
class WorkflowRunArchivesApi(Resource):
@console_ns.doc("list_workflow_run_archives")
@console_ns.doc(description="List monthly workflow-run archive metadata for the current workspace")
@console_ns.response(200, "Success", console_ns.models[WorkflowRunArchiveListResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self):
tenant_id, _ = _current_ids()
return dump_response(WorkflowRunArchiveListResponse, list_workflow_run_archives(db.session(), tenant_id))
@console_ns.route("/workflow-run-archives/downloads")
class WorkflowRunArchiveDownloadsApi(Resource):
@console_ns.doc("create_workflow_run_archive_download")
@console_ns.doc(description="Create or return a temporary workflow-run archive download task")
@console_ns.expect(console_ns.models[WorkflowRunArchiveDownloadPayload.__name__])
@console_ns.response(
HTTPStatus.ACCEPTED,
"Download task accepted",
console_ns.models[WorkflowRunArchiveDownloadTaskResponse.__name__],
)
@setup_required
@login_required
@account_initialization_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def post(self):
tenant_id, account_id = _current_ids()
payload = WorkflowRunArchiveDownloadPayload.model_validate(console_ns.payload or {})
try:
task = create_workflow_run_archive_download_task(
db.session(),
tenant_id=tenant_id,
requested_by=account_id,
year=payload.year,
month=payload.month,
)
except WorkflowRunArchiveNotFoundError as exc:
raise NotFound(str(exc)) from exc
return dump_response(WorkflowRunArchiveDownloadTaskResponse, task), HTTPStatus.ACCEPTED
@console_ns.route("/workflow-run-archives/downloads/<string:download_id>")
class WorkflowRunArchiveDownloadApi(Resource):
@console_ns.doc("get_workflow_run_archive_download")
@console_ns.doc(description="Get a temporary workflow-run archive download task")
@console_ns.response(200, "Success", console_ns.models[WorkflowRunArchiveDownloadTaskResponse.__name__])
@setup_required
@login_required
@account_initialization_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self, download_id: str):
tenant_id, _ = _current_ids()
try:
task = get_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id)
except WorkflowRunArchiveDownloadTaskNotFoundError as exc:
raise NotFound(str(exc)) from exc
return dump_response(WorkflowRunArchiveDownloadTaskResponse, task)
@console_ns.route("/workflow-run-archives/downloads/<string:download_id>/file")
class WorkflowRunArchiveDownloadFileApi(Resource):
@console_ns.doc("download_workflow_run_archive_file")
@console_ns.doc(description="Redirect to a prepared workflow-run archive ZIP file")
@console_ns.response(
302,
"Redirect to pre-signed archive storage URL",
console_ns.models[RedirectResponse.__name__],
)
@console_ns.response(409, "Download task is not ready")
@setup_required
@login_required
@account_initialization_required
@is_admin_or_owner_required
@rbac_permission_required(
RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False
)
def get(self, download_id: str):
tenant_id, _ = _current_ids()
try:
task = get_ready_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id)
except WorkflowRunArchiveDownloadTaskNotFoundError as exc:
raise NotFound(str(exc)) from exc
except WorkflowRunArchiveDownloadNotReadyError as exc:
raise Conflict(str(exc)) from exc
storage_key = task.storage_key
if storage_key is None:
raise Conflict(f"Workflow run archive download is not ready: {download_id}")
storage = get_export_storage()
presigned_url = storage.generate_presigned_url(
storage_key,
expires_in=_presigned_url_expires_in(task.expires_at),
filename=task.file_name,
content_type=ARCHIVE_DOWNLOAD_MIME_TYPE,
)
return redirect(presigned_url, code=HTTPStatus.FOUND)
+4 -2
View File
@@ -1,3 +1,5 @@
from mimetypes import guess_extension
from flask import request
from flask_restx import Resource
from flask_restx.api import HTTPStatus
@@ -6,7 +8,7 @@ from werkzeug.exceptions import Forbidden
import services
from core.tools.signature import verify_plugin_file_signature
from core.tools.tool_file_manager import ToolFileManager, resolve_extension
from core.tools.tool_file_manager import ToolFileManager
from core.workflow.file_reference import build_file_reference
from fields.file_fields import FileResponse
@@ -108,7 +110,7 @@ class PluginUploadFileApi(Resource):
conversation_id=args.conversation_id,
)
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
extension = guess_extension(tool_file.mimetype) or ".bin"
preview_url = ToolFileManager.sign_file(tool_file_id=tool_file.id, extension=extension)
# Create a dictionary with all the necessary attributes
@@ -476,7 +476,6 @@ class PluginDownloadFileRequestApi(Resource):
user_from=payload.user_from,
invoke_from=payload.invoke_from,
file_mapping=payload.file.model_dump(mode="python", exclude_none=True),
for_external=payload.for_external,
)
return BaseBackwardsInvocationResponse(
data={
@@ -333,7 +333,6 @@ class AgentAppGenerator(MessageBasedAppGenerator):
target=self._generate_worker,
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"session": db.session(),
"context": context,
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
+4 -42
View File
@@ -27,7 +27,6 @@ from clients.agent_backend import (
AgentBackendInternalEventType,
AgentBackendRunClient,
AgentBackendRunEventAdapter,
AgentBackendRunFailedInternalEvent,
AgentBackendRunSucceededInternalEvent,
AgentBackendStreamInternalEvent,
extract_runtime_layer_specs,
@@ -58,14 +57,6 @@ from core.workflow.nodes.agent_v2.ask_human_resume import build_deferred_tool_re
from extensions.ext_database import db
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, PromptMessage, UserPromptMessage
from graphon.model_runtime.errors.invoke import (
InvokeAuthorizationError,
InvokeBadRequestError,
InvokeConnectionError,
InvokeError,
InvokeRateLimitError,
InvokeServerUnavailableError,
)
from models.agent_config_entities import AgentSoulConfig
from models.enums import CreatorUserRole
from models.model import MessageAgentThought
@@ -80,22 +71,6 @@ class _DefaultSessionScopeSnapshotId:
_DEFAULT_SESSION_SCOPE_SNAPSHOT_ID = _DefaultSessionScopeSnapshotId()
_AGENT_BACKEND_INVOKE_ERROR_BY_REASON: Mapping[str, type[InvokeError]] = {
"InvokeAuthorizationError": InvokeAuthorizationError,
"InvokeBadRequestError": InvokeBadRequestError,
"CredentialsValidateFailedError": InvokeBadRequestError,
"InvokeConnectionError": InvokeConnectionError,
"InvokeRateLimitError": InvokeRateLimitError,
"InvokeServerUnavailableError": InvokeServerUnavailableError,
}
def _agent_backend_failure_to_exception(event: AgentBackendRunFailedInternalEvent) -> Exception:
err_cls = _AGENT_BACKEND_INVOKE_ERROR_BY_REASON.get(event.reason or "")
if err_cls is not None:
return err_cls(event.error)
return AgentBackendError(event.error or "Agent backend run did not complete successfully.")
def _prompt_messages_from_query(user_query: str | None) -> list[PromptMessage]:
if not user_query:
@@ -437,15 +412,12 @@ class _AgentProcessRecorder:
def _lookup_tool_thought(self, *, index: int, tool_call_id: str | None) -> str | None:
if tool_call_id and tool_call_id in self._tool_by_call_id:
return self._tool_by_call_id[tool_call_id]
if index < 0:
return None
return self._tool_by_index.get(index)
def _remember_tool_thought(
self, *, index: int, tool_call_id: str | None, tool_name: str | None, thought_id: str
) -> None:
if index >= 0:
self._tool_by_index[index] = thought_id
self._tool_by_index[index] = thought_id
if tool_call_id:
self._tool_by_call_id[tool_call_id] = thought_id
if tool_name:
@@ -461,10 +433,6 @@ class _AgentProcessRecorder:
return None
def _mark_tool_observed(self, thought_id: str) -> None:
self._tool_by_index = {index: value for index, value in self._tool_by_index.items() if value != thought_id}
self._tool_by_call_id = {
tool_call_id: value for tool_call_id, value in self._tool_by_call_id.items() if value != thought_id
}
for open_thought_ids in self._open_tool_by_name.values():
open_thought_ids.discard(thought_id)
@@ -562,12 +530,7 @@ def _event_index(data: dict[str, Any]) -> int:
def _string_or_none(value: Any) -> str | None:
if not isinstance(value, str):
return None
normalized = value.strip()
if not normalized or normalized.lower() in {"none", "null"}:
return None
return normalized
return value if isinstance(value, str) and value else None
def _json_or_text(value: Any) -> str | None:
@@ -689,9 +652,8 @@ class AgentAppRunner:
return
if not isinstance(terminal, AgentBackendRunSucceededInternalEvent):
if isinstance(terminal, AgentBackendRunFailedInternalEvent):
raise _agent_backend_failure_to_exception(terminal)
raise AgentBackendError("Agent backend run did not complete successfully.")
error = getattr(terminal, "error", None) or "Agent backend run did not complete successfully."
raise AgentBackendError(str(error))
answer = self._terminal_output_to_answer(terminal.output)
try:
@@ -205,7 +205,6 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
target=self._generate_worker,
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"session": db.session(),
"context": context,
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
@@ -8,7 +8,7 @@ from pydantic import JsonValue
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.task_entities import AppBlockingResponse, AppStreamResponse
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from graphon.model_runtime.errors.invoke import InvokeError, InvokeRateLimitError
from graphon.model_runtime.errors.invoke import InvokeError
logger = logging.getLogger(__name__)
@@ -127,7 +127,6 @@ class AppGenerateResponseConverter[TBlockingResponse: AppBlockingResponse](ABC):
},
ModelCurrentlyNotSupportError: {"code": "model_currently_not_support", "status": 400},
InvokeError: {"code": "completion_request_error", "status": 400},
InvokeRateLimitError: {"code": "rate_limit_error", "status": 429},
}
# Determine the response based on the type of exception
@@ -420,7 +420,6 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
message.total_price = usage.total_price
message.currency = usage.currency
self._task_state.llm_result.usage.latency = message.provider_response_latency
self._task_state.metadata.usage = self._task_state.llm_result.usage
message.message_metadata = self._task_state.metadata.model_dump_json()
if trace_manager:
+1 -1
View File
@@ -27,7 +27,7 @@ class ProviderCredentialsCache:
try:
cached_provider_credentials = cached_provider_credentials.decode("utf-8")
cached_provider_credentials = json.loads(cached_provider_credentials)
except (JSONDecodeError, UnicodeDecodeError):
except JSONDecodeError:
return None
return dict(cached_provider_credentials)
+1 -1
View File
@@ -24,7 +24,7 @@ class ProviderCredentialsCache(ABC):
try:
cached_credentials = cached_credentials.decode("utf-8")
return dict(json.loads(cached_credentials))
except (JSONDecodeError, UnicodeDecodeError):
except JSONDecodeError:
return None
return None
+1 -1
View File
@@ -30,7 +30,7 @@ class ToolParameterCache:
try:
cached_tool_parameter = cached_tool_parameter.decode("utf-8")
cached_tool_parameter = json.loads(cached_tool_parameter)
except (JSONDecodeError, UnicodeDecodeError):
except JSONDecodeError:
return None
return dict(cached_tool_parameter)
-1
View File
@@ -276,7 +276,6 @@ class RequestRequestDownloadFile(BaseModel):
"validation",
]
file: RequestDownloadFileMapping
for_external: bool = True
model_config = ConfigDict(extra="forbid")
+3 -12
View File
@@ -4,7 +4,6 @@ import hmac
import logging
import os
import time
import urllib.parse
from collections.abc import Generator
from mimetypes import guess_extension, guess_type
from uuid import uuid4
@@ -27,7 +26,7 @@ logger = logging.getLogger(__name__)
class ToolFileManager:
@staticmethod
def _build_graph_file_reference(tool_file: ToolFile) -> File:
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
extension = guess_extension(tool_file.mimetype) or ".bin"
return File(
file_type=get_file_type_by_mime_type(tool_file.mimetype),
transfer_method=FileTransferMethod.TOOL_FILE,
@@ -71,7 +70,7 @@ class ToolFileManager:
mimetype: str,
filename: str | None = None,
) -> ToolFile:
extension = resolve_extension(filename=filename, mimetype=mimetype)
extension = guess_extension(mimetype) or ".bin"
unique_name = uuid4().hex
unique_filename = f"{unique_name}{extension}"
# default just as before
@@ -121,8 +120,7 @@ class ToolFileManager:
or response.headers.get("Content-Type", "").split(";")[0].strip()
or "application/octet-stream"
)
url_filename = os.path.basename(urllib.parse.urlparse(file_url).path)
extension = resolve_extension(filename=url_filename, mimetype=mimetype)
extension = guess_extension(mimetype) or ".bin"
unique_name = uuid4().hex
filename = f"{unique_name}{extension}"
filepath = f"tools/{tenant_id}/{filename}"
@@ -222,11 +220,4 @@ def _factory() -> ToolFileManager:
return ToolFileManager()
def resolve_extension(*, filename: str | None, mimetype: str) -> str:
filename_extension = os.path.splitext(filename or "")[1].lower()
if filename_extension:
return filename_extension
return guess_extension(mimetype) or ".bin"
set_tool_file_manager_factory(_factory)
+4 -5
View File
@@ -3,6 +3,7 @@ import re
from collections.abc import Generator
from datetime import date, datetime
from decimal import Decimal
from mimetypes import guess_extension
from typing import Any
from uuid import UUID
@@ -10,7 +11,7 @@ import numpy as np
import pytz
from core.tools.entities.tool_entities import ToolInvokeMessage
from core.tools.tool_file_manager import ToolFileManager, resolve_extension
from core.tools.tool_file_manager import ToolFileManager
from core.workflow.file_reference import parse_file_reference
from graphon.file import File, FileTransferMethod, FileType
from libs.login import current_user
@@ -90,8 +91,7 @@ class ToolFileMessageTransformer:
conversation_id=conversation_id,
)
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
url = cls.get_tool_file_url(tool_file_id=tool_file.id, extension=extension)
url = f"/files/tools/{tool_file.id}{guess_extension(tool_file.mimetype) or '.png'}"
meta = cls._with_tool_file_meta(
message.meta,
tool_file_id=str(tool_file.id),
@@ -136,8 +136,7 @@ class ToolFileMessageTransformer:
filename=filename,
)
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
url = cls.get_tool_file_url(tool_file_id=tool_file.id, extension=extension)
url = cls.get_tool_file_url(tool_file_id=tool_file.id, extension=guess_extension(tool_file.mimetype))
meta = cls._with_tool_file_meta(meta, tool_file_id=str(tool_file.id))
# check if file is image
@@ -7,7 +7,6 @@ from typing import TYPE_CHECKING, Any, override
from agenton.compositor import CompositorSessionSnapshot
from clients.agent_backend import (
AgentBackendAgentMessageDeltaInternalEvent,
AgentBackendDeferredToolCallInternalEvent,
AgentBackendError,
AgentBackendHTTPError,
@@ -482,10 +481,6 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
if isinstance(internal_event, AgentBackendStreamInternalEvent):
self._record_stream_metadata(metadata, internal_event)
continue
if internal_event.type == AgentBackendInternalEventType.AGENT_MESSAGE_DELTA:
if isinstance(internal_event, AgentBackendAgentMessageDeltaInternalEvent):
self._record_agent_message_delta_metadata(metadata, internal_event)
continue
metadata["agent_backend"] = {
**dict(metadata.get("agent_backend") or {}),
"stream_event_count": stream_event_count,
@@ -739,17 +734,6 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
agent_backend["usage"] = dict(usage)
metadata["agent_backend"] = agent_backend
@staticmethod
def _record_agent_message_delta_metadata(
metadata: dict[str, Any], event: AgentBackendAgentMessageDeltaInternalEvent
) -> None:
agent_backend = dict(metadata.get("agent_backend") or {})
agent_backend["agent_message_delta_count"] = int(agent_backend.get("agent_message_delta_count") or 0) + 1
agent_backend["agent_message_delta_length"] = int(agent_backend.get("agent_message_delta_length") or 0) + len(
event.delta
)
metadata["agent_backend"] = agent_backend
@classmethod
@override
def _extract_variable_selector_to_variable_mapping(
@@ -10,13 +10,13 @@ trustworthy metadata.
from __future__ import annotations
from mimetypes import guess_extension
from uuid import UUID
from sqlalchemy import select
from sqlalchemy.exc import DataError, SQLAlchemyError
from core.db.session_factory import session_factory
from core.tools.tool_file_manager import resolve_extension
from core.workflow.file_reference import build_file_reference
from graphon.file import File, FileTransferMethod, get_file_type_by_mime_type
from models.tools import ToolFile
@@ -46,7 +46,7 @@ def reback_tool_file_output(*, tenant_id: str, tool_file_id: str) -> File | None
return None
mime_type = tool_file.mimetype or ""
extension = resolve_extension(filename=tool_file.name, mimetype=mime_type)
extension = guess_extension(mime_type) or ".bin"
return File(
type=get_file_type_by_mime_type(mime_type),
transfer_method=FileTransferMethod.TOOL_FILE,
+1
View File
@@ -156,6 +156,7 @@ def init_app(app: DifyApp) -> Celery:
"tasks.generate_summary_index_task", # summary index generation
"tasks.regenerate_summary_index_task", # summary index regeneration
"tasks.app_generate.resume_agent_app_task", # ENG-635: Agent v2 chat ask_human resume
"tasks.workflow_run_archive_download_tasks", # workflow-run archive download preparation
]
day = dify_config.CELERY_BEAT_SCHEDULER_TIME
+2
View File
@@ -7,6 +7,7 @@ def init_app(app: DifyApp):
archive_workflow_runs,
archive_workflow_runs_plan,
backfill_plugin_auto_upgrade,
backfill_workflow_run_archive_bundles,
clean_expired_messages,
clean_workflow_runs,
cleanup_orphaned_draft_variables,
@@ -77,6 +78,7 @@ def init_app(app: DifyApp):
install_rag_pipeline_plugins,
archive_workflow_runs_plan,
archive_workflow_runs,
backfill_workflow_run_archive_bundles,
delete_archived_workflow_runs,
restore_workflow_runs,
clean_workflow_runs,
+2 -8
View File
@@ -3,7 +3,6 @@
from __future__ import annotations
import mimetypes
import os
import uuid
from collections.abc import Mapping, Sequence
from typing import Any, Literal, NotRequired, TypedDict, assert_never, cast
@@ -286,7 +285,7 @@ def _build_from_remote_url(
raise ValueError("Invalid file url")
mime_type, filename, file_size = get_remote_file_info(url)
extension = os.path.splitext(filename)[1].lower() or mimetypes.guess_extension(mime_type) or ".bin"
extension = mimetypes.guess_extension(mime_type) or ("." + filename.split(".")[-1] if "." in filename else ".bin")
detected_file_type = standardize_file_type(extension=extension, mime_type=mime_type)
file_type = _resolve_file_type(
detected_file_type=detected_file_type,
@@ -327,12 +326,7 @@ def _build_from_tool_file(
if tool_file is None:
raise ValueError(f"ToolFile {tool_file_id} not found")
extension = (
os.path.splitext(tool_file.name)[1].lower()
or mimetypes.guess_extension(tool_file.mimetype)
or os.path.splitext(tool_file.file_key)[1].lower()
or ".bin"
)
extension = "." + tool_file.file_key.split(".")[-1] if "." in tool_file.file_key else ".bin"
detected_file_type = standardize_file_type(extension=extension, mime_type=tool_file.mimetype)
file_type = _resolve_file_type(
detected_file_type=detected_file_type,
+1 -11
View File
@@ -1,10 +1,9 @@
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
from uuid import uuid4
from pydantic import Field, computed_field, field_validator
from pydantic import Field, field_validator
from core.entities.execution_extra_content import ExecutionExtraContentDomainModel
from fields.base import ResponseModel
@@ -56,19 +55,10 @@ class MessageListItem(ResponseModel):
created_at: int | None = None
agent_thoughts: list[AgentThought]
message_files: list[MessageFile]
message_tokens: int = 0
answer_tokens: int = 0
provider_response_latency: float = 0
total_price: Decimal | None = None
currency: str | None = None
status: str
error: str | None = None
extra_contents: list[ExecutionExtraContentDomainModel]
@computed_field
def total_tokens(self) -> int:
return self.message_tokens + self.answer_tokens
@field_validator("inputs", mode="before")
@classmethod
def _normalize_inputs(cls, value: JSONValueType) -> JSONValueType:
+17 -2
View File
@@ -11,6 +11,7 @@ import hashlib
import logging
from collections.abc import Generator
from typing import Any, cast
from urllib.parse import quote
import boto3
import orjson
@@ -197,13 +198,22 @@ class ArchiveStorage:
except ClientError as e:
raise ArchiveStorageError(f"Failed to delete object '{key}': {e}")
def generate_presigned_url(self, key: str, expires_in: int = 3600) -> str:
def generate_presigned_url(
self,
key: str,
expires_in: int = 3600,
*,
filename: str | None = None,
content_type: str | None = None,
) -> str:
"""
Generate a pre-signed URL for downloading an object.
Args:
key: Object key (path) within the bucket
expires_in: URL validity duration in seconds (default: 1 hour)
filename: Optional browser download filename
content_type: Optional response content type
Returns:
Pre-signed URL string.
@@ -211,10 +221,15 @@ class ArchiveStorage:
Raises:
ArchiveStorageError: If generation fails
"""
params = {"Bucket": self.bucket, "Key": key}
if filename:
params["ResponseContentDisposition"] = f"attachment; filename*=UTF-8''{quote(filename)}"
if content_type:
params["ResponseContentType"] = content_type
try:
return self.client.generate_presigned_url(
ClientMethod="get_object",
Params={"Bucket": self.bucket, "Key": key},
Params=params,
ExpiresIn=expires_in,
)
except ClientError as e:
+1
View File
@@ -22,6 +22,7 @@ logger = logging.getLogger(__name__)
CSRF_WHITE_LIST = [
re.compile(r"/console/api/apps/[a-f0-9-]+/workflows/draft"),
re.compile(r"/console/api/workflow-run-archives/downloads/[a-f0-9]+/file"),
]
@@ -0,0 +1,61 @@
"""add workflow run archive bundle index table
Revision ID: 7a1c2d9e4b60
Revises: c3d4e5f6a7b8
Create Date: 2026-06-25 15:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
import models
# revision identifiers, used by Alembic.
revision = "7a1c2d9e4b60"
down_revision = "c3d4e5f6a7b8"
branch_labels = None
depends_on = None
def _uuid_column(name: str, **kwargs):
if op.get_bind().dialect.name == "postgresql":
kwargs.setdefault("server_default", sa.text("uuidv7()"))
return sa.Column(name, models.types.StringUUID(), **kwargs)
def upgrade() -> None:
op.create_table(
"workflow_run_archive_bundles",
_uuid_column("id", nullable=False),
sa.Column("tenant_id", models.types.StringUUID(), nullable=False),
sa.Column("year", sa.Integer(), nullable=False),
sa.Column("month", sa.Integer(), nullable=False),
sa.Column("shard", sa.String(length=32), nullable=False),
sa.Column("bundle_id", sa.String(length=64), nullable=False),
sa.Column("workflow_run_count", sa.Integer(), nullable=False),
sa.Column("row_count", sa.BigInteger(), nullable=False),
sa.Column("archive_bytes", sa.BigInteger(), nullable=False),
sa.Column("archived_at", sa.DateTime(), nullable=False),
sa.Column("created_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
sa.Column("updated_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
sa.PrimaryKeyConstraint("id", name="workflow_run_archive_bundle_pkey"),
sa.UniqueConstraint(
"tenant_id",
"year",
"month",
"shard",
"bundle_id",
name="workflow_run_archive_bundle_identity_uq",
),
)
op.create_index(
"workflow_run_archive_bundle_tenant_month_idx",
"workflow_run_archive_bundles",
["tenant_id", "year", "month"],
)
def downgrade() -> None:
op.drop_index("workflow_run_archive_bundle_tenant_month_idx", table_name="workflow_run_archive_bundles")
op.drop_table("workflow_run_archive_bundles")
+2
View File
@@ -144,6 +144,7 @@ from .workflow import (
WorkflowNodeExecutionTriggeredFrom,
WorkflowPause,
WorkflowRun,
WorkflowRunArchiveBundle,
WorkflowType,
resolve_workflow_kind,
)
@@ -282,6 +283,7 @@ __all__ = [
"WorkflowNodeExecutionTriggeredFrom",
"WorkflowPause",
"WorkflowRun",
"WorkflowRunArchiveBundle",
"WorkflowRunTriggeredFrom",
"WorkflowSchedulePlan",
"WorkflowToolProvider",
+34
View File
@@ -1446,6 +1446,40 @@ class WorkflowArchiveLog(TypeBase):
}
class WorkflowRunArchiveBundle(DefaultFieldsDCMixin, TypeBase):
"""
Query index for one immutable V2 workflow-run archive bundle.
R2 manifest objects remain the recoverable archive source of truth. This table stores the small subset needed to
list tenant/month archives and locate bundles without listing object storage online. Missing rows can be rebuilt
from existing manifests by a backfill/reconciliation command.
"""
__tablename__ = "workflow_run_archive_bundles"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="workflow_run_archive_bundle_pkey"),
sa.UniqueConstraint(
"tenant_id",
"year",
"month",
"shard",
"bundle_id",
name="workflow_run_archive_bundle_identity_uq",
),
sa.Index("workflow_run_archive_bundle_tenant_month_idx", "tenant_id", "year", "month"),
)
tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
year: Mapped[int] = mapped_column(sa.Integer, nullable=False)
month: Mapped[int] = mapped_column(sa.Integer, nullable=False)
shard: Mapped[str] = mapped_column(String(32), nullable=False)
bundle_id: Mapped[str] = mapped_column(String(64), nullable=False)
workflow_run_count: Mapped[int] = mapped_column(sa.Integer, nullable=False)
row_count: Mapped[int] = mapped_column(sa.BigInteger, nullable=False)
archive_bytes: Mapped[int] = mapped_column(sa.BigInteger, nullable=False)
archived_at: Mapped[datetime] = mapped_column(DateTime, nullable=False)
class ConversationVariable(TypeBase):
__tablename__ = "workflow_conversation_variables"
+120 -6
View File
@@ -9690,6 +9690,61 @@ Suggest example workflow-generator instructions for the tenant
| 200 | Suggestions generated successfully | **application/json**: [GeneratorResponse](#generatorresponse)<br> |
| 400 | Invalid request parameters | |
### [GET] /workflow-run-archives
List monthly workflow-run archive metadata for the current workspace
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [WorkflowRunArchiveListResponse](#workflowrunarchivelistresponse)<br> |
### [POST] /workflow-run-archives/downloads
Create or return a temporary workflow-run archive download task
#### Request Body
| Required | Schema |
| -------- | ------ |
| Yes | **application/json**: [WorkflowRunArchiveDownloadPayload](#workflowrunarchivedownloadpayload)<br> |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 202 | Download task accepted | **application/json**: [WorkflowRunArchiveDownloadTaskResponse](#workflowrunarchivedownloadtaskresponse)<br> |
### [GET] /workflow-run-archives/downloads/{download_id}
Get a temporary workflow-run archive download task
#### Parameters
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| download_id | path | | Yes | string |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 200 | Success | **application/json**: [WorkflowRunArchiveDownloadTaskResponse](#workflowrunarchivedownloadtaskresponse)<br> |
### [GET] /workflow-run-archives/downloads/{download_id}/file
Redirect to a prepared workflow-run archive ZIP file
#### Parameters
| Name | Located in | Description | Required | Schema |
| ---- | ---------- | ----------- | -------- | ------ |
| download_id | path | | Yes | string |
#### Responses
| Code | Description | Schema |
| ---- | ----------- | ------ |
| 302 | Redirect to pre-signed archive storage URL | **application/json**: [RedirectResponse](#redirectresponse)<br> |
| 409 | Download task is not ready | |
### [GET] /workflow/{workflow_run_id}/events
**Get workflow execution events stream after resume**
@@ -17614,25 +17669,19 @@ Built-in tool icons are URL strings; API-based tool icons are provider-defined p
| ---- | ---- | ----------- | -------- |
| agent_thoughts | [ [AgentThought](#agentthought) ] | | Yes |
| answer | string | | Yes |
| answer_tokens | integer | | No |
| conversation_id | string | | Yes |
| created_at | integer | | No |
| currency | string | | No |
| error | string | | No |
| extra_contents | [ [HumanInputContent](#humaninputcontent) ] | | Yes |
| feedback | [SimpleFeedback](#simplefeedback) | | No |
| id | string | | Yes |
| inputs | object | | Yes |
| message_files | [ [MessageFile](#messagefile) ] | | Yes |
| message_tokens | integer | | No |
| metadata | [JSONValueType](#jsonvaluetype) | | No |
| parent_message_id | string | | No |
| provider_response_latency | number | | No |
| query | string | | Yes |
| retriever_resources | [ [RetrieverResource](#retrieverresource) ] | | Yes |
| status | string | | Yes |
| total_price | string | | No |
| total_tokens | integer | | Yes |
#### ExternalApiTemplateListQuery
@@ -23300,6 +23349,71 @@ tenant's default model. The underlying generator never raises — an empty
| result | string | | Yes |
| updated_at | integer | | Yes |
#### WorkflowRunArchiveDownloadPayload
Request body for preparing one monthly workflow-run archive download.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| month | integer | | Yes |
| year | integer | | Yes |
#### WorkflowRunArchiveDownloadStatus
Lifecycle state for an asynchronous archive download request.
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| WorkflowRunArchiveDownloadStatus | string | Lifecycle state for an asynchronous archive download request. | |
#### WorkflowRunArchiveDownloadTaskResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| archive_bytes | integer | | Yes |
| bundle_count | integer | | Yes |
| created_at | dateTime | | Yes |
| download_id | string | | Yes |
| error | string | | No |
| expires_at | dateTime | | Yes |
| file_name | string | | No |
| file_size_bytes | integer | | No |
| finished_at | string | | No |
| month | integer | | Yes |
| started_at | string | | No |
| status | [WorkflowRunArchiveDownloadStatus](#workflowrunarchivedownloadstatus) | | Yes |
| updated_at | dateTime | | Yes |
| year | integer | | Yes |
#### WorkflowRunArchiveListResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| months | [ [WorkflowRunArchiveMonthResponse](#workflowrunarchivemonthresponse) ] | | Yes |
| summary | [WorkflowRunArchiveSummaryResponse](#workflowrunarchivesummaryresponse) | | Yes |
#### WorkflowRunArchiveMonthResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| archive_bytes | integer | | Yes |
| bundle_count | integer | | Yes |
| download_task | [WorkflowRunArchiveDownloadTaskResponse](#workflowrunarchivedownloadtaskresponse) | | No |
| latest_archived_at | dateTime | | Yes |
| month | integer | | Yes |
| row_count | integer | | Yes |
| workflow_run_count | integer | | Yes |
| year | integer | | Yes |
#### WorkflowRunArchiveSummaryResponse
| Name | Type | Description | Required |
| ---- | ---- | ----------- | -------- |
| archive_bytes | integer | | Yes |
| archived_month_count | integer | | Yes |
| latest_archived_at | string | | No |
| workflow_run_count | integer | | Yes |
#### WorkflowRunCountQuery
| Name | Type | Description | Required |
-6
View File
@@ -3467,24 +3467,18 @@ Model class for i18n object.
| ---- | ---- | ----------- | -------- |
| agent_thoughts | [ [AgentThought](#agentthought) ] | | Yes |
| answer | string | | Yes |
| answer_tokens | integer | | No |
| conversation_id | string | | Yes |
| created_at | integer | | No |
| currency | string | | No |
| error | string | | No |
| extra_contents | [ [HumanInputContent](#humaninputcontent) ] | | Yes |
| feedback | [SimpleFeedback](#simplefeedback) | | No |
| id | string | | Yes |
| inputs | object | | Yes |
| message_files | [ [MessageFile](#messagefile) ] | | Yes |
| message_tokens | integer | | No |
| parent_message_id | string | | No |
| provider_response_latency | number | | No |
| query | string | | Yes |
| retriever_resources | [ [RetrieverResource](#retrieverresource) ] | | Yes |
| status | string | | Yes |
| total_price | string | | No |
| total_tokens | integer | | Yes |
#### MessageListQuery
-6
View File
@@ -1685,25 +1685,19 @@ in form definiton, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| agent_thoughts | [ [AgentThought](#agentthought) ] | | Yes |
| answer | string | | Yes |
| answer_tokens | integer | | No |
| conversation_id | string | | Yes |
| created_at | integer | | No |
| currency | string | | No |
| error | string | | No |
| extra_contents | [ [HumanInputContent](#humaninputcontent) ] | | Yes |
| feedback | [SimpleFeedback](#simplefeedback) | | No |
| id | string | | Yes |
| inputs | object | | Yes |
| message_files | [ [MessageFile](#messagefile) ] | | Yes |
| message_tokens | integer | | No |
| metadata | [JSONValueType](#jsonvaluetype) | | No |
| parent_message_id | string | | No |
| provider_response_latency | number | | No |
| query | string | | Yes |
| retriever_resources | [ [RetrieverResource](#retrieverresource) ] | | Yes |
| status | string | | Yes |
| total_price | string | | No |
| total_tokens | integer | | Yes |
#### WebModelConfigResponse
@@ -1,6 +1,5 @@
from __future__ import annotations
import logging
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import cast
@@ -260,16 +259,14 @@ def test_get_project_url_success(trace_instance: AliyunDataTrace):
assert trace_instance.get_project_url() == "project-url"
def test_get_project_url_error(
trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
):
def test_get_project_url_error(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(trace_instance.trace_client, "get_project_url", MagicMock(side_effect=Exception("boom")))
logger_mock = MagicMock()
monkeypatch.setattr(aliyun_trace_module, "logger", logger_mock)
caplog.set_level(logging.INFO, logger=aliyun_trace_module.logger.name)
with pytest.raises(ValueError, match=r"Aliyun get project url failed: boom"):
trace_instance.get_project_url()
assert "Aliyun get project url failed: boom" in caplog.text
logger_mock.info.assert_called()
def test_workflow_trace_adds_workflow_and_node_spans(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch):
@@ -2,7 +2,6 @@
from __future__ import annotations
import logging
import sys
import types
from types import SimpleNamespace
@@ -88,6 +87,7 @@ class PatchedCoreComponents(TypedDict):
tracer: MagicMock
span: MagicMock
tracer_provider: MagicMock
logger: MagicMock
trace_api: Any
@@ -148,6 +148,9 @@ def patch_core_components(monkeypatch: pytest.MonkeyPatch) -> PatchedCoreCompone
resource = MagicMock(name="resource")
monkeypatch.setattr(client_module, "Resource", MagicMock(return_value=resource))
logger_mock = MagicMock(name="tencent_logger")
monkeypatch.setattr(client_module, "logger", logger_mock)
trace_api_stub = SimpleNamespace(
set_span_in_context=MagicMock(name="set_span_in_context", return_value="trace-context"),
NonRecordingSpan=MagicMock(name="non_recording_span", side_effect=lambda ctx: f"non-{ctx}"),
@@ -171,6 +174,7 @@ def patch_core_components(monkeypatch: pytest.MonkeyPatch) -> PatchedCoreCompone
"tracer": tracer,
"span": span,
"tracer_provider": tracer_provider,
"logger": logger_mock,
"trace_api": trace_api_stub,
}
@@ -264,15 +268,14 @@ def test_record_methods_skip_when_histogram_missing() -> None:
client.record_trace_duration(0.5)
def test_record_llm_duration_handles_exceptions(caplog: pytest.LogCaptureFixture) -> None:
def test_record_llm_duration_handles_exceptions(patch_core_components: PatchedCoreComponents) -> None:
client = _build_client()
client.hist_llm_duration = MagicMock(name="hist_llm_duration")
client.hist_llm_duration.record.side_effect = RuntimeError("boom")
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client.record_llm_duration(0.2)
assert "[Tencent APM] Failed to record LLM duration" in caplog.text
logger = patch_core_components["logger"]
logger.debug.assert_called()
def test_create_and_export_span_sets_attributes(patch_core_components: PatchedCoreComponents) -> None:
@@ -325,15 +328,12 @@ def test_create_and_export_span_uses_parent_context(patch_core_components: Patch
trace_api.set_span_in_context.assert_called_once()
def test_create_and_export_span_exception_logs_error(
patch_core_components: PatchedCoreComponents, caplog: pytest.LogCaptureFixture
) -> None:
def test_create_and_export_span_exception_logs_error(patch_core_components: PatchedCoreComponents) -> None:
client = _build_client()
span = patch_core_components["span"]
span.get_span_context.return_value = _make_span_context(span_id=2)
client.tracer.start_span.side_effect = RuntimeError("boom")
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client._create_and_export_span(
SpanData(
trace_id=1,
@@ -346,10 +346,8 @@ def test_create_and_export_span_exception_logs_error(
end_time=1,
)
)
error_records = [record for record in caplog.records if record.levelno == logging.ERROR]
assert len(error_records) == 1
assert error_records[0].getMessage() == "[Tencent APM] Error creating span: span"
logger = patch_core_components["logger"]
logger.exception.assert_called_once()
def test_api_check_connects_successfully(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -425,18 +423,23 @@ def test_shutdown_flushes_all_components(patch_core_components: PatchedCoreCompo
metric_reader.shutdown.assert_called_once()
def test_shutdown_logs_when_meter_provider_fails(caplog: pytest.LogCaptureFixture) -> None:
def test_shutdown_logs_when_meter_provider_fails(patch_core_components: PatchedCoreComponents) -> None:
client = _build_client()
meter_provider = meter_provider_instances[-1]
meter_provider.shutdown.side_effect = RuntimeError("boom")
assert client.metric_reader is not None
client.metric_reader.shutdown.side_effect = RuntimeError("boom")
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client.shutdown()
assert "[Tencent APM] Error shutting down meter provider" in caplog.text
assert "[Tencent APM] Error shutting down metric reader" in caplog.text
logger = patch_core_components["logger"]
logger.debug.assert_any_call(
"[Tencent APM] Error shutting down meter provider",
exc_info=True,
)
logger.debug.assert_any_call(
"[Tencent APM] Error shutting down metric reader",
exc_info=True,
)
def test_metrics_initialization_failure_sets_histogram_attributes(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -453,11 +456,10 @@ def test_metrics_initialization_failure_sets_histogram_attributes(monkeypatch: p
assert client.metric_reader is None
def test_add_span_logs_exception(monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture) -> None:
def test_add_span_logs_exception(monkeypatch: pytest.MonkeyPatch, patch_core_components: PatchedCoreComponents) -> None:
client = _build_client()
monkeypatch.setattr(client, "_create_and_export_span", MagicMock(side_effect=RuntimeError("boom")))
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client.add_span(
SpanData(
trace_id=1,
@@ -471,9 +473,8 @@ def test_add_span_logs_exception(monkeypatch: pytest.MonkeyPatch, caplog: pytest
)
)
error_records = [record for record in caplog.records if record.levelno == logging.ERROR]
assert len(error_records) == 1
assert error_records[0].getMessage() == "[Tencent APM] Failed to create span: span"
logger = patch_core_components["logger"]
logger.exception.assert_called_once()
def test_create_and_export_span_converts_attribute_types(patch_core_components: PatchedCoreComponents) -> None:
@@ -534,20 +535,16 @@ def test_record_trace_duration_converts_attributes() -> None:
],
)
def test_record_methods_handle_exceptions(
method: str, attr_name: str, args: tuple[object, ...], caplog: pytest.LogCaptureFixture
method: str, attr_name: str, args: tuple[object, ...], patch_core_components: PatchedCoreComponents
) -> None:
client = _build_client()
hist_mock = MagicMock(name=attr_name)
hist_mock.record.side_effect = RuntimeError("boom")
setattr(client, attr_name, hist_mock)
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
getattr(client, method)(*args)
assert any(
record.levelno == logging.DEBUG and record.getMessage().startswith("[Tencent APM] Failed to record")
for record in caplog.records
)
logger = patch_core_components["logger"]
logger.debug.assert_called()
def test_metrics_initializes_grpc_metric_exporter() -> None:
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-api"
version = "1.16.0-rc1"
version = "1.15.0"
requires-python = "~=3.12.0"
dependencies = [
+1 -2
View File
@@ -45,7 +45,6 @@ class FileRequestService:
user_from: UserFrom | str,
invoke_from: InvokeFrom | str,
file_mapping: Mapping[str, Any],
for_external: bool = True,
) -> DownloadFileRequestResult:
"""Resolve one file mapping into signed download metadata.
@@ -62,7 +61,7 @@ class FileRequestService:
)
with bind_file_access_scope(scope):
file = self._build_file(mapping=file_mapping, tenant_id=tenant_id)
download_url = file_helpers.resolve_file_url(file, for_external=for_external)
download_url = file_helpers.resolve_file_url(file, for_external=True)
if not download_url:
raise ValueError("file does not support signed download")
@@ -0,0 +1,347 @@
"""
Workflow-run archive bundle index helpers.
Archive manifests in object storage remain the recoverable source of truth. This module mirrors their small query
surface into `workflow_run_archive_bundles` so console listing and download jobs can avoid listing R2 on request.
The backfill path is intentionally idempotent: every manifest is decoded, checked against the V2 schema markers, and
upserted by immutable bundle identity.
"""
import datetime
import json
import logging
import time
from collections.abc import Sequence
from dataclasses import dataclass, field
from typing import TypedDict, cast
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_database import db
from libs.archive_storage import ArchiveStorage, get_archive_storage
from models.workflow import WorkflowRunArchiveBundle
from services.retention.workflow_run.constants import (
ARCHIVE_BUNDLE_FORMAT,
ARCHIVE_BUNDLE_MANIFEST_NAME,
ARCHIVE_BUNDLE_SCHEMA_VERSION,
)
logger = logging.getLogger(__name__)
ARCHIVE_BUNDLE_ROOT_PREFIX = "workflow-runs/v2/"
class ArchiveBundleTableManifestEntry(TypedDict):
row_count: int
checksum: str
size_bytes: int
object_key: str
class ArchiveBundleManifest(TypedDict):
schema_version: str
archive_format: str
tenant_id: str
tenant_prefix: str
year: int
month: int
shard: str
bundle_id: str
object_prefix: str
workflow_run_count: int
workflow_node_execution_count: int
min_created_at: str
max_created_at: str
min_run_id: str
max_run_id: str
archived_at: str
tables: dict[str, ArchiveBundleTableManifestEntry]
run_ids: list[str]
@dataclass(frozen=True)
class ArchiveBundleIndexValues:
"""Computed DB-index values derived from one manifest."""
row_count: int
archive_bytes: int
archived_at: datetime.datetime
@dataclass
class ArchiveBundleIndexBackfillSummary:
"""Aggregate result for a manifest-to-DB-index reconciliation run."""
manifests_found: int = 0
bundles_processed: int = 0
bundles_upserted: int = 0
bundles_failed: int = 0
workflow_run_count: int = 0
row_count: int = 0
archive_bytes: int = 0
elapsed_time: float = 0.0
errors: list[str] = field(default_factory=list)
def decode_archive_bundle_manifest(manifest_data: bytes) -> ArchiveBundleManifest:
"""Decode raw manifest bytes into the V2 archive manifest shape."""
return cast(ArchiveBundleManifest, json.loads(manifest_data.decode("utf-8")))
def parse_archive_manifest_datetime(value: str) -> datetime.datetime:
"""Parse manifest datetimes and normalize timezone-aware values to naive UTC for DB storage."""
parsed = datetime.datetime.fromisoformat(value)
if parsed.tzinfo is None:
return parsed
return parsed.astimezone(datetime.UTC).replace(tzinfo=None)
def calculate_archive_bundle_index_values(
manifest: ArchiveBundleManifest,
manifest_size_bytes: int,
) -> ArchiveBundleIndexValues:
"""Calculate row count, stored bytes, and archived timestamp for the DB index."""
_validate_archive_bundle_manifest(manifest)
row_count = sum(entry["row_count"] for entry in manifest["tables"].values())
archive_bytes = manifest_size_bytes + sum(entry["size_bytes"] for entry in manifest["tables"].values())
return ArchiveBundleIndexValues(
row_count=row_count,
archive_bytes=archive_bytes,
archived_at=parse_archive_manifest_datetime(manifest["archived_at"]),
)
def upsert_archive_bundle_index_from_manifest(
session: Session,
manifest: ArchiveBundleManifest,
manifest_size_bytes: int,
) -> WorkflowRunArchiveBundle:
"""
Persist one archive manifest into `workflow_run_archive_bundles`.
The caller owns transaction boundaries. Re-running this function for the same manifest is safe and refreshes the
mutable metrics derived from object sizes and row counts.
"""
values = calculate_archive_bundle_index_values(manifest, manifest_size_bytes)
existing = session.scalar(
select(WorkflowRunArchiveBundle).where(
WorkflowRunArchiveBundle.tenant_id == manifest["tenant_id"],
WorkflowRunArchiveBundle.year == manifest["year"],
WorkflowRunArchiveBundle.month == manifest["month"],
WorkflowRunArchiveBundle.shard == manifest["shard"],
WorkflowRunArchiveBundle.bundle_id == manifest["bundle_id"],
)
)
if existing is None:
bundle = WorkflowRunArchiveBundle(
tenant_id=manifest["tenant_id"],
year=manifest["year"],
month=manifest["month"],
shard=manifest["shard"],
bundle_id=manifest["bundle_id"],
workflow_run_count=manifest["workflow_run_count"],
row_count=values.row_count,
archive_bytes=values.archive_bytes,
archived_at=values.archived_at,
)
session.add(bundle)
return bundle
existing.workflow_run_count = manifest["workflow_run_count"]
existing.row_count = values.row_count
existing.archive_bytes = values.archive_bytes
existing.archived_at = values.archived_at
return existing
class WorkflowRunArchiveBundleIndexBackfill:
"""
Rebuild the DB bundle index by scanning object-store manifests.
Tenant IDs are the cheapest scope because they map directly to the object prefix. Tenant prefixes are supported for
rollout reconciliation, but they still require listing all tenants under that prefix and filtering keys locally.
"""
storage: ArchiveStorage | None
session_factory: sessionmaker[Session]
def __init__(
self,
*,
storage: ArchiveStorage | None = None,
session_factory: sessionmaker[Session] | None = None,
) -> None:
self.storage = storage
self.session_factory = session_factory or sessionmaker(bind=db.engine, expire_on_commit=False)
def run(
self,
*,
tenant_ids: Sequence[str] | None = None,
tenant_prefixes: Sequence[str] | None = None,
year: int | None = None,
month: int | None = None,
limit: int | None = None,
dry_run: bool = False,
) -> ArchiveBundleIndexBackfillSummary:
"""Scan matching manifest objects and idempotently upsert their DB index rows."""
start_time = time.time()
summary = ArchiveBundleIndexBackfillSummary()
storage = self.storage or get_archive_storage()
manifest_keys = self._list_manifest_keys(
storage,
tenant_ids=tenant_ids,
tenant_prefixes=tenant_prefixes,
year=year,
month=month,
)
summary.manifests_found = len(manifest_keys)
if limit is not None:
manifest_keys = manifest_keys[:limit]
for manifest_key in manifest_keys:
try:
manifest_data = storage.get_object(manifest_key)
manifest = decode_archive_bundle_manifest(manifest_data)
self._validate_manifest_scope(
manifest,
manifest_key=manifest_key,
tenant_ids=tenant_ids,
tenant_prefixes=tenant_prefixes,
year=year,
month=month,
)
values = calculate_archive_bundle_index_values(manifest, len(manifest_data))
summary.bundles_processed += 1
summary.workflow_run_count += manifest["workflow_run_count"]
summary.row_count += values.row_count
summary.archive_bytes += values.archive_bytes
if dry_run:
continue
with self.session_factory() as session:
upsert_archive_bundle_index_from_manifest(session, manifest, len(manifest_data))
session.commit()
summary.bundles_upserted += 1
except Exception as exc:
logger.warning("Failed to backfill workflow archive bundle index from %s", manifest_key, exc_info=True)
summary.bundles_failed += 1
summary.errors.append(f"{manifest_key}: {exc}")
summary.elapsed_time = time.time() - start_time
return summary
@classmethod
def _list_manifest_keys(
cls,
storage: ArchiveStorage,
*,
tenant_ids: Sequence[str] | None,
tenant_prefixes: Sequence[str] | None,
year: int | None,
month: int | None,
) -> list[str]:
prefixes = cls._list_prefixes(tenant_ids=tenant_ids, tenant_prefixes=tenant_prefixes, year=year, month=month)
keys: list[str] = []
for prefix in prefixes:
keys.extend(storage.list_objects(prefix))
return sorted(
key
for key in keys
if key.endswith(f"/{ARCHIVE_BUNDLE_MANIFEST_NAME}")
and cls._manifest_key_matches_scope(
key,
tenant_ids=tenant_ids,
tenant_prefixes=tenant_prefixes,
year=year,
month=month,
)
)
@staticmethod
def _list_prefixes(
*,
tenant_ids: Sequence[str] | None,
tenant_prefixes: Sequence[str] | None,
year: int | None,
month: int | None,
) -> list[str]:
if tenant_ids:
prefixes = []
for tenant_id in sorted(set(tenant_ids)):
prefix = f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix={tenant_id[0].lower()}/tenant_id={tenant_id}/"
if year is not None:
prefix += f"year={year:04d}/"
if month is not None:
prefix += f"month={month:02d}/"
prefixes.append(prefix)
return prefixes
if tenant_prefixes:
return [
f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix={tenant_prefix}/"
for tenant_prefix in sorted(set(tenant_prefixes))
]
return [ARCHIVE_BUNDLE_ROOT_PREFIX]
@staticmethod
def _manifest_key_matches_scope(
key: str,
*,
tenant_ids: Sequence[str] | None,
tenant_prefixes: Sequence[str] | None,
year: int | None,
month: int | None,
) -> bool:
if tenant_ids and _extract_key_part(key, "tenant_id") not in set(tenant_ids):
return False
if tenant_prefixes and _extract_key_part(key, "tenant_prefix") not in set(tenant_prefixes):
return False
if year is not None and _extract_key_part(key, "year") != f"{year:04d}":
return False
if month is not None and _extract_key_part(key, "month") != f"{month:02d}":
return False
return True
@staticmethod
def _validate_manifest_scope(
manifest: ArchiveBundleManifest,
*,
manifest_key: str,
tenant_ids: Sequence[str] | None,
tenant_prefixes: Sequence[str] | None,
year: int | None,
month: int | None,
) -> None:
expected_object_prefix = manifest_key.removesuffix(f"/{ARCHIVE_BUNDLE_MANIFEST_NAME}")
if manifest["object_prefix"] != expected_object_prefix:
raise ValueError(
f"manifest object_prefix mismatch: expected={expected_object_prefix}, "
f"actual={manifest['object_prefix']}"
)
if tenant_ids and manifest["tenant_id"] not in tenant_ids:
raise ValueError(f"manifest tenant_id is outside requested scope: {manifest['tenant_id']}")
if tenant_prefixes and manifest["tenant_prefix"] not in tenant_prefixes:
raise ValueError(f"manifest tenant_prefix is outside requested scope: {manifest['tenant_prefix']}")
if year is not None and manifest["year"] != year:
raise ValueError(f"manifest year is outside requested scope: {manifest['year']}")
if month is not None and manifest["month"] != month:
raise ValueError(f"manifest month is outside requested scope: {manifest['month']}")
def _validate_archive_bundle_manifest(manifest: ArchiveBundleManifest) -> None:
if manifest["schema_version"] != ARCHIVE_BUNDLE_SCHEMA_VERSION:
raise ValueError(f"unsupported archive bundle schema version: {manifest['schema_version']}")
if manifest["archive_format"] != ARCHIVE_BUNDLE_FORMAT:
raise ValueError(f"unsupported archive bundle format: {manifest['archive_format']}")
def _extract_key_part(key: str, name: str) -> str | None:
prefix = f"{name}="
for part in key.split("/"):
if part.startswith(prefix):
return part[len(prefix) :]
return None
@@ -0,0 +1,325 @@
"""
Prepare monthly workflow-run archive downloads.
Console requests create a short-lived Redis task and Celery runs this module in the background. The DB bundle index is
the online lookup source: this preparer never lists archive storage, and it validates the indexed bundle set against the
stable download id before packaging archive Parquet objects into one user-facing CSV ZIP file.
"""
import datetime
import hashlib
import io
import logging
import zipfile
from collections.abc import Sequence
from typing import cast
import pyarrow.csv as pa_csv
import pyarrow.parquet as pq
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_database import db
from libs.archive_storage import ArchiveStorage, get_archive_storage, get_export_storage
from models.workflow import WorkflowRunArchiveBundle
from services.retention.workflow_run.archive_bundle_index import (
ARCHIVE_BUNDLE_ROOT_PREFIX,
ArchiveBundleManifest,
ArchiveBundleTableManifestEntry,
decode_archive_bundle_manifest,
)
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
WorkflowRunArchiveDownloadTask,
WorkflowRunArchiveDownloadTaskCache,
build_archive_download_id,
)
from services.retention.workflow_run.constants import (
ARCHIVE_BUNDLE_FORMAT,
ARCHIVE_BUNDLE_MANIFEST_NAME,
ARCHIVE_BUNDLE_SCHEMA_VERSION,
)
logger = logging.getLogger(__name__)
ARCHIVE_DOWNLOAD_ROOT_PREFIX = "workflow-runs/downloads/v1/"
ARCHIVE_DOWNLOAD_MIME_TYPE = "application/zip"
class WorkflowRunArchiveDownloadPreparer:
"""
Build one ready-to-download CSV ZIP for a Redis archive download task.
The output object is deterministic for a given `download_id`, so retrying a failed task overwrites the same
temporary object instead of creating unbounded duplicate files. Source archive bundles are read from the archive
bucket, while the prepared ZIP is written to the export bucket so object lifecycle policies can expire downloads
without touching long-lived archives.
"""
archive_storage: ArchiveStorage | None
download_storage: ArchiveStorage | None
cache: WorkflowRunArchiveDownloadTaskCache
session_factory: sessionmaker[Session]
def __init__(
self,
*,
storage: ArchiveStorage | None = None,
archive_storage: ArchiveStorage | None = None,
download_storage: ArchiveStorage | None = None,
cache: WorkflowRunArchiveDownloadTaskCache | None = None,
session_factory: sessionmaker[Session] | None = None,
) -> None:
self.archive_storage = archive_storage or storage
self.download_storage = download_storage or storage
self.cache = cache or WorkflowRunArchiveDownloadTaskCache()
self.session_factory = session_factory or sessionmaker(bind=db.engine, expire_on_commit=False)
def prepare(self, *, tenant_id: str, download_id: str) -> WorkflowRunArchiveDownloadTask | None:
"""Prepare a ZIP for an existing Redis task and persist terminal task state."""
task = self.cache.get(tenant_id=tenant_id, download_id=download_id)
if task is None:
logger.info("Workflow run archive download task expired before preparation: %s", download_id)
return None
if task.status == WorkflowRunArchiveDownloadStatus.READY:
return task
if task.status == WorkflowRunArchiveDownloadStatus.FAILED:
logger.info("Skipping failed workflow run archive download task: %s", download_id)
return task
processing_task = self._mark_processing(task)
try:
archive_storage = self.archive_storage or get_archive_storage()
download_storage = self.download_storage or get_export_storage()
bundles = self._get_task_bundles(processing_task)
payload = self._build_zip_payload(archive_storage, processing_task, bundles)
storage_key = build_archive_download_storage_key(processing_task)
download_storage.put_object(storage_key, payload)
return self._mark_ready(processing_task, storage_key=storage_key, file_size_bytes=len(payload))
except Exception as exc:
logger.exception("Failed to prepare workflow run archive download %s", download_id)
return self._mark_failed(processing_task, error=str(exc))
def _get_task_bundles(self, task: WorkflowRunArchiveDownloadTask) -> list[WorkflowRunArchiveBundle]:
with self.session_factory() as session:
return _list_task_bundles(session, task)
def _build_zip_payload(
self,
storage: ArchiveStorage,
task: WorkflowRunArchiveDownloadTask,
bundles: Sequence[WorkflowRunArchiveBundle],
) -> bytes:
zip_root = f"workflow-run-logs-{task.year:04d}-{task.month:02d}"
csv_buffers: dict[str, io.BytesIO] = {}
csv_headers_written: set[str] = set()
for bundle in bundles:
object_prefix = _build_archive_bundle_object_prefix(task, bundle)
_, manifest = _load_and_validate_manifest(storage, task, bundle, object_prefix)
for table_name in sorted(manifest["tables"]):
entry = manifest["tables"][table_name]
object_key = entry["object_key"]
table_payload = storage.get_object(object_key)
_validate_table_payload(object_key=object_key, entry=entry, payload=table_payload)
csv_payload = _parquet_payload_to_csv(
table_payload,
include_header=table_name not in csv_headers_written,
)
if not csv_payload:
continue
csv_buffers.setdefault(table_name, io.BytesIO()).write(csv_payload)
csv_headers_written.add(table_name)
buffer = io.BytesIO()
with zipfile.ZipFile(buffer, mode="w", compression=zipfile.ZIP_DEFLATED) as archive:
for table_name, csv_buffer in sorted(csv_buffers.items()):
archive.writestr(f"{zip_root}/{table_name}.csv", csv_buffer.getvalue())
return buffer.getvalue()
def _mark_processing(self, task: WorkflowRunArchiveDownloadTask) -> WorkflowRunArchiveDownloadTask:
now = datetime.datetime.now(datetime.UTC)
processing_task = task.model_copy(
update={
"status": WorkflowRunArchiveDownloadStatus.PROCESSING,
"error": None,
"updated_at": now,
"started_at": task.started_at or now,
}
)
self.cache.save(processing_task)
return processing_task
def _mark_ready(
self,
task: WorkflowRunArchiveDownloadTask,
*,
storage_key: str,
file_size_bytes: int,
) -> WorkflowRunArchiveDownloadTask:
now = datetime.datetime.now(datetime.UTC)
ready_task = task.model_copy(
update={
"status": WorkflowRunArchiveDownloadStatus.READY,
"file_name": build_archive_download_file_name(task),
"storage_key": storage_key,
"file_size_bytes": file_size_bytes,
"error": None,
"updated_at": now,
"finished_at": now,
}
)
self.cache.save(ready_task)
return ready_task
def _mark_failed(self, task: WorkflowRunArchiveDownloadTask, *, error: str) -> WorkflowRunArchiveDownloadTask:
now = datetime.datetime.now(datetime.UTC)
failed_task = task.model_copy(
update={
"status": WorkflowRunArchiveDownloadStatus.FAILED,
"error": error,
"updated_at": now,
"finished_at": now,
}
)
self.cache.save(failed_task)
return failed_task
def build_archive_download_file_name(task: WorkflowRunArchiveDownloadTask) -> str:
"""Return the browser download filename for one monthly archive."""
return f"workflow-run-logs-{task.year:04d}-{task.month:02d}.zip"
def build_archive_download_storage_key(task: WorkflowRunArchiveDownloadTask) -> str:
"""Return the deterministic object-store key for a prepared download ZIP."""
return (
f"{ARCHIVE_DOWNLOAD_ROOT_PREFIX}tenant_prefix={task.tenant_id[0].lower()}/tenant_id={task.tenant_id}/"
f"year={task.year:04d}/month={task.month:02d}/{task.download_id}.zip"
)
def _list_task_bundles(session: Session, task: WorkflowRunArchiveDownloadTask) -> list[WorkflowRunArchiveBundle]:
stmt = (
select(WorkflowRunArchiveBundle)
.where(
WorkflowRunArchiveBundle.tenant_id == task.tenant_id,
WorkflowRunArchiveBundle.year == task.year,
WorkflowRunArchiveBundle.month == task.month,
)
.order_by(WorkflowRunArchiveBundle.shard, WorkflowRunArchiveBundle.bundle_id)
)
indexed_bundles = list(session.scalars(stmt))
if task.bundle_refs:
requested_refs = [(ref.shard, ref.bundle_id) for ref in task.bundle_refs]
else:
requested_bundle_ids = set(task.bundle_ids)
requested_refs = [
(bundle.shard, bundle.bundle_id) for bundle in indexed_bundles if bundle.bundle_id in requested_bundle_ids
]
bundle_by_ref = {(bundle.shard, bundle.bundle_id): bundle for bundle in indexed_bundles}
missing_refs = [ref for ref in requested_refs if ref not in bundle_by_ref]
if missing_refs:
raise ValueError(f"archive bundle index is missing requested bundles: {missing_refs}")
bundles = [bundle_by_ref[ref] for ref in requested_refs]
if len(bundles) != task.bundle_count:
raise ValueError(f"archive bundle count changed: expected={task.bundle_count}, actual={len(bundles)}")
download_id = build_archive_download_id(
tenant_id=task.tenant_id,
year=task.year,
month=task.month,
bundle_refs=requested_refs,
)
if download_id != task.download_id:
raise ValueError("archive download id no longer matches indexed bundle set")
return bundles
def _build_archive_bundle_object_prefix(
task: WorkflowRunArchiveDownloadTask,
bundle: WorkflowRunArchiveBundle,
) -> str:
return (
f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix={task.tenant_id[0].lower()}/tenant_id={task.tenant_id}/"
f"year={task.year:04d}/month={task.month:02d}/shard={bundle.shard}/bundle={bundle.bundle_id}"
)
def _load_and_validate_manifest(
storage: ArchiveStorage,
task: WorkflowRunArchiveDownloadTask,
bundle: WorkflowRunArchiveBundle,
object_prefix: str,
) -> tuple[bytes, ArchiveBundleManifest]:
manifest_key = f"{object_prefix}/{ARCHIVE_BUNDLE_MANIFEST_NAME}"
manifest_data = storage.get_object(manifest_key)
manifest = decode_archive_bundle_manifest(manifest_data)
_validate_manifest(task=task, bundle=bundle, manifest=manifest, object_prefix=object_prefix)
return manifest_data, manifest
def _validate_manifest(
*,
task: WorkflowRunArchiveDownloadTask,
bundle: WorkflowRunArchiveBundle,
manifest: ArchiveBundleManifest,
object_prefix: str,
) -> None:
if manifest["schema_version"] != ARCHIVE_BUNDLE_SCHEMA_VERSION:
raise ValueError(f"unsupported archive bundle schema version: {manifest['schema_version']}")
if manifest["archive_format"] != ARCHIVE_BUNDLE_FORMAT:
raise ValueError(f"unsupported archive bundle format: {manifest['archive_format']}")
if manifest["tenant_id"] != task.tenant_id:
raise ValueError(f"manifest tenant_id mismatch: expected={task.tenant_id}, actual={manifest['tenant_id']}")
if manifest["year"] != task.year:
raise ValueError(f"manifest year mismatch: expected={task.year}, actual={manifest['year']}")
if manifest["month"] != task.month:
raise ValueError(f"manifest month mismatch: expected={task.month}, actual={manifest['month']}")
if manifest["shard"] != bundle.shard:
raise ValueError(f"manifest shard mismatch: expected={bundle.shard}, actual={manifest['shard']}")
if manifest["bundle_id"] != bundle.bundle_id:
raise ValueError(f"manifest bundle_id mismatch: expected={bundle.bundle_id}, actual={manifest['bundle_id']}")
if manifest["object_prefix"] != object_prefix:
raise ValueError(
f"manifest object_prefix mismatch: expected={object_prefix}, actual={manifest['object_prefix']}"
)
if not manifest["tables"]:
raise ValueError("manifest tables must not be empty")
for table_name, raw_entry in manifest["tables"].items():
entry = cast(ArchiveBundleTableManifestEntry, raw_entry)
expected_object_key = f"{object_prefix}/{table_name}.parquet"
if entry["object_key"] != expected_object_key:
raise ValueError(
f"manifest object_key mismatch for {table_name}: "
f"expected={expected_object_key}, actual={entry['object_key']}"
)
def _validate_table_payload(
*,
object_key: str,
entry: ArchiveBundleTableManifestEntry,
payload: bytes,
) -> None:
if len(payload) != entry["size_bytes"]:
raise ValueError(f"archive object size mismatch for {object_key}")
checksum = hashlib.md5(payload).hexdigest()
if checksum != entry["checksum"]:
raise ValueError(f"archive object checksum mismatch for {object_key}")
def _parquet_payload_to_csv(payload: bytes, *, include_header: bool) -> bytes:
table = pq.read_table(io.BytesIO(payload))
if table.num_columns == 0:
return b""
buffer = io.BytesIO()
pa_csv.write_csv(
table,
buffer,
write_options=pa_csv.WriteOptions(include_header=include_header),
)
return buffer.getvalue()
@@ -0,0 +1,178 @@
"""Redis-backed temporary state for workflow-run archive downloads."""
import datetime
import hashlib
import json
import logging
from collections.abc import Sequence
from enum import StrEnum
from pydantic import BaseModel, ConfigDict, Field
from extensions.ext_redis import RedisClientWrapper, redis_client
logger = logging.getLogger(__name__)
ARCHIVE_DOWNLOAD_FORMAT_VERSION = "v1"
DEFAULT_ARCHIVE_DOWNLOAD_TASK_TTL_SECONDS = 24 * 60 * 60
_CACHE_KEY_PREFIX = "workflow_run_archive_download"
class WorkflowRunArchiveDownloadStatus(StrEnum):
"""Lifecycle state for an asynchronous archive download request."""
PENDING = "pending"
PROCESSING = "processing"
READY = "ready"
FAILED = "failed"
class WorkflowRunArchiveBundleRef(BaseModel):
"""Immutable object-store identity for one bundle included in a download task."""
model_config = ConfigDict(extra="forbid")
shard: str
bundle_id: str
class WorkflowRunArchiveDownloadTask(BaseModel):
"""Temporary Redis payload for a monthly archive download request."""
model_config = ConfigDict(extra="forbid")
download_id: str
tenant_id: str
requested_by: str
year: int = Field(ge=1)
month: int = Field(ge=1, le=12)
bundle_ids: list[str]
bundle_refs: list[WorkflowRunArchiveBundleRef] = Field(default_factory=list)
bundle_count: int = Field(ge=0)
archive_bytes: int = Field(ge=0)
status: WorkflowRunArchiveDownloadStatus
file_name: str | None = None
storage_key: str | None = None
file_size_bytes: int | None = Field(default=None, ge=0)
celery_task_id: str | None = None
error: str | None = None
created_at: datetime.datetime
updated_at: datetime.datetime
expires_at: datetime.datetime
started_at: datetime.datetime | None = None
finished_at: datetime.datetime | None = None
class WorkflowRunArchiveDownloadTaskCache:
"""Store ephemeral archive download task state in Redis with a TTL."""
_redis: RedisClientWrapper
def __init__(self, redis: RedisClientWrapper = redis_client) -> None:
self._redis = redis
def get(self, *, tenant_id: str, download_id: str) -> WorkflowRunArchiveDownloadTask | None:
raw = self._redis.get(self._cache_key(tenant_id=tenant_id, download_id=download_id))
if raw is None:
return None
data = raw.decode("utf-8") if isinstance(raw, bytes | bytearray) else raw
try:
return WorkflowRunArchiveDownloadTask.model_validate_json(data)
except ValueError:
logger.warning("Malformed workflow run archive download task cache entry: %s", download_id)
return None
def save(self, task: WorkflowRunArchiveDownloadTask) -> None:
ttl_seconds = self._ttl_seconds(task.expires_at)
self._redis.setex(
self._cache_key(tenant_id=task.tenant_id, download_id=task.download_id),
ttl_seconds,
task.model_dump_json(),
)
def create_if_absent(self, task: WorkflowRunArchiveDownloadTask) -> bool:
ttl_seconds = self._ttl_seconds(task.expires_at)
result = self._redis.set(
self._cache_key(tenant_id=task.tenant_id, download_id=task.download_id),
task.model_dump_json(),
ex=ttl_seconds,
nx=True,
)
return bool(result)
def delete(self, *, tenant_id: str, download_id: str) -> None:
self._redis.delete(self._cache_key(tenant_id=tenant_id, download_id=download_id))
@staticmethod
def _cache_key(*, tenant_id: str, download_id: str) -> str:
return f"{_CACHE_KEY_PREFIX}:{tenant_id}:{download_id}"
@staticmethod
def _ttl_seconds(expires_at: datetime.datetime) -> int:
expires_at_utc = expires_at if expires_at.tzinfo else expires_at.replace(tzinfo=datetime.UTC)
remaining = expires_at_utc - datetime.datetime.now(datetime.UTC)
return max(int(remaining.total_seconds()), 1)
def build_pending_archive_download_task(
*,
tenant_id: str,
requested_by: str,
year: int,
month: int,
bundle_ids: Sequence[str],
bundle_refs: Sequence[tuple[str, str]] = (),
archive_bytes: int,
download_id: str,
ttl_seconds: int = DEFAULT_ARCHIVE_DOWNLOAD_TASK_TTL_SECONDS,
now: datetime.datetime | None = None,
) -> WorkflowRunArchiveDownloadTask:
"""Create the Redis payload stored when the console starts an archive download."""
created_at = now or datetime.datetime.now(datetime.UTC)
if created_at.tzinfo is None:
created_at = created_at.replace(tzinfo=datetime.UTC)
normalized_bundle_ids = list(bundle_ids)
normalized_bundle_refs = [
WorkflowRunArchiveBundleRef(shard=shard, bundle_id=bundle_id) for shard, bundle_id in bundle_refs
]
return WorkflowRunArchiveDownloadTask(
download_id=download_id,
tenant_id=tenant_id,
requested_by=requested_by,
year=year,
month=month,
bundle_ids=normalized_bundle_ids,
bundle_refs=normalized_bundle_refs,
bundle_count=len(normalized_bundle_ids),
archive_bytes=archive_bytes,
status=WorkflowRunArchiveDownloadStatus.PENDING,
created_at=created_at,
updated_at=created_at,
expires_at=created_at + datetime.timedelta(seconds=ttl_seconds),
)
def build_archive_download_id(
*,
tenant_id: str,
year: int,
month: int,
bundle_refs: Sequence[tuple[str, str]],
download_format_version: str = ARCHIVE_DOWNLOAD_FORMAT_VERSION,
) -> str:
"""Build a stable id for the exact archive download content."""
if not bundle_refs:
raise ValueError("bundle_refs must not be empty")
normalized_refs = sorted(f"{shard}:{bundle_id}" for shard, bundle_id in bundle_refs)
payload = json.dumps(
{
"tenant_id": tenant_id,
"year": year,
"month": month,
"bundle_refs": normalized_refs,
"download_format_version": download_format_version,
},
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()[:32]
@@ -0,0 +1,293 @@
"""
Console-facing workflow-run archive queries.
The object store remains the recoverable archive source of truth. This module only reads the DB bundle index and writes
temporary Redis download-task state, so console requests never list R2 online.
"""
import datetime
import logging
import uuid
from collections.abc import Callable
from dataclasses import dataclass
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.workflow import WorkflowRunArchiveBundle
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
WorkflowRunArchiveDownloadTask,
WorkflowRunArchiveDownloadTaskCache,
build_archive_download_id,
build_pending_archive_download_task,
)
logger = logging.getLogger(__name__)
ArchiveDownloadTaskDispatcher = Callable[
[WorkflowRunArchiveDownloadTask, WorkflowRunArchiveDownloadTaskCache],
WorkflowRunArchiveDownloadTask,
]
@dataclass(frozen=True)
class WorkflowRunArchiveMonth:
"""Aggregated archive metadata for one tenant/month."""
year: int
month: int
bundle_count: int
workflow_run_count: int
row_count: int
archive_bytes: int
latest_archived_at: datetime.datetime
download_task: WorkflowRunArchiveDownloadTask | None
@dataclass(frozen=True)
class WorkflowRunArchiveSummary:
"""Top-level archive totals shown on the console page."""
archived_month_count: int
workflow_run_count: int
archive_bytes: int
latest_archived_at: datetime.datetime | None
@dataclass(frozen=True)
class WorkflowRunArchiveList:
"""Console response model before controller serialization."""
summary: WorkflowRunArchiveSummary
months: list[WorkflowRunArchiveMonth]
class WorkflowRunArchiveNotFoundError(Exception):
"""Raised when no archive bundles exist for a requested tenant/month."""
class WorkflowRunArchiveDownloadTaskNotFoundError(Exception):
"""Raised when the temporary Redis task has expired or never existed."""
class WorkflowRunArchiveDownloadNotReadyError(Exception):
"""Raised when a cached download task has not produced a file yet."""
def list_workflow_run_archives(
session: Session,
tenant_id: str,
*,
cache: WorkflowRunArchiveDownloadTaskCache | None = None,
) -> WorkflowRunArchiveList:
"""Return monthly archive metadata for one tenant from the DB bundle index."""
stmt = (
select(WorkflowRunArchiveBundle)
.where(WorkflowRunArchiveBundle.tenant_id == tenant_id)
.order_by(
WorkflowRunArchiveBundle.year.desc(),
WorkflowRunArchiveBundle.month.desc(),
WorkflowRunArchiveBundle.shard,
WorkflowRunArchiveBundle.bundle_id,
)
)
month_bundles: dict[tuple[int, int], list[WorkflowRunArchiveBundle]] = {}
for bundle in session.scalars(stmt):
month_bundles.setdefault((bundle.year, bundle.month), []).append(bundle)
task_cache = cache or WorkflowRunArchiveDownloadTaskCache()
months: list[WorkflowRunArchiveMonth] = []
for (year, month), bundles in month_bundles.items():
bundle_refs = [(bundle.shard, bundle.bundle_id) for bundle in bundles]
months.append(
WorkflowRunArchiveMonth(
year=year,
month=month,
bundle_count=len(bundles),
workflow_run_count=sum(bundle.workflow_run_count for bundle in bundles),
row_count=sum(bundle.row_count for bundle in bundles),
archive_bytes=sum(bundle.archive_bytes for bundle in bundles),
latest_archived_at=max(bundle.archived_at for bundle in bundles),
download_task=_get_cached_month_download_task(
task_cache,
tenant_id=tenant_id,
year=year,
month=month,
bundle_refs=bundle_refs,
),
)
)
latest_archived_at = max((month.latest_archived_at for month in months), default=None)
return WorkflowRunArchiveList(
summary=WorkflowRunArchiveSummary(
archived_month_count=len(months),
workflow_run_count=sum(month.workflow_run_count for month in months),
archive_bytes=sum(month.archive_bytes for month in months),
latest_archived_at=latest_archived_at,
),
months=months,
)
def _get_cached_month_download_task(
cache: WorkflowRunArchiveDownloadTaskCache,
*,
tenant_id: str,
year: int,
month: int,
bundle_refs: list[tuple[str, str]],
) -> WorkflowRunArchiveDownloadTask | None:
if not bundle_refs:
return None
download_id = build_archive_download_id(
tenant_id=tenant_id,
year=year,
month=month,
bundle_refs=bundle_refs,
)
try:
return cache.get(tenant_id=tenant_id, download_id=download_id)
except Exception:
logger.warning("Failed to read cached workflow run archive download task: %s", download_id, exc_info=True)
return None
def create_workflow_run_archive_download_task(
session: Session,
*,
tenant_id: str,
requested_by: str,
year: int,
month: int,
cache: WorkflowRunArchiveDownloadTaskCache | None = None,
dispatcher: ArchiveDownloadTaskDispatcher | None = None,
) -> WorkflowRunArchiveDownloadTask:
"""
Create or return the idempotent Redis task for downloading one tenant/month archive.
The task identity is based on the exact ordered bundle set currently indexed for the month. If the month receives a
new bundle later, the next request gets a different download id and prepares a fresh file.
"""
bundles = _list_archive_bundles(session, tenant_id=tenant_id, year=year, month=month)
if not bundles:
raise WorkflowRunArchiveNotFoundError(f"Workflow run archive not found: {year:04d}-{month:02d}")
bundle_refs = [(bundle.shard, bundle.bundle_id) for bundle in bundles]
download_id = build_archive_download_id(
tenant_id=tenant_id,
year=year,
month=month,
bundle_refs=bundle_refs,
)
task = build_pending_archive_download_task(
tenant_id=tenant_id,
requested_by=requested_by,
year=year,
month=month,
bundle_ids=[bundle.bundle_id for bundle in bundles],
bundle_refs=bundle_refs,
archive_bytes=sum(bundle.archive_bytes for bundle in bundles),
download_id=download_id,
)
task_cache = cache or WorkflowRunArchiveDownloadTaskCache()
dispatch = dispatcher or _dispatch_workflow_run_archive_download_task
if task_cache.create_if_absent(task):
return dispatch(task, task_cache)
existing = task_cache.get(tenant_id=tenant_id, download_id=download_id)
if existing is not None:
if existing.status == WorkflowRunArchiveDownloadStatus.FAILED:
task_cache.save(task)
return dispatch(task, task_cache)
if existing.status == WorkflowRunArchiveDownloadStatus.PENDING and not existing.celery_task_id:
return dispatch(existing, task_cache)
return existing
task_cache.save(task)
return dispatch(task, task_cache)
def get_workflow_run_archive_download_task(
*,
tenant_id: str,
download_id: str,
cache: WorkflowRunArchiveDownloadTaskCache | None = None,
) -> WorkflowRunArchiveDownloadTask:
"""Return a cached archive download task or raise when the TTL has expired."""
task_cache = cache or WorkflowRunArchiveDownloadTaskCache()
task = task_cache.get(tenant_id=tenant_id, download_id=download_id)
if task is None:
raise WorkflowRunArchiveDownloadTaskNotFoundError(f"Workflow run archive download not found: {download_id}")
return task
def get_ready_workflow_run_archive_download_task(
*,
tenant_id: str,
download_id: str,
cache: WorkflowRunArchiveDownloadTaskCache | None = None,
) -> WorkflowRunArchiveDownloadTask:
"""Return a ready cached archive download task or raise when the file is not available."""
task = get_workflow_run_archive_download_task(tenant_id=tenant_id, download_id=download_id, cache=cache)
if task.status != WorkflowRunArchiveDownloadStatus.READY or not task.storage_key or not task.file_name:
raise WorkflowRunArchiveDownloadNotReadyError(f"Workflow run archive download is not ready: {download_id}")
return task
def _list_archive_bundles(
session: Session,
*,
tenant_id: str,
year: int,
month: int,
) -> list[WorkflowRunArchiveBundle]:
stmt = (
select(WorkflowRunArchiveBundle)
.where(
WorkflowRunArchiveBundle.tenant_id == tenant_id,
WorkflowRunArchiveBundle.year == year,
WorkflowRunArchiveBundle.month == month,
)
.order_by(WorkflowRunArchiveBundle.shard, WorkflowRunArchiveBundle.bundle_id)
)
return list(session.scalars(stmt))
def _dispatch_workflow_run_archive_download_task(
task: WorkflowRunArchiveDownloadTask,
cache: WorkflowRunArchiveDownloadTaskCache,
) -> WorkflowRunArchiveDownloadTask:
"""
Enqueue background ZIP preparation and persist the Celery id before the worker can start.
The Redis task key is the idempotency boundary. We generate the Celery id in the API process, save it on the task,
then submit with that exact id so duplicate console requests keep seeing one logical download request.
"""
from tasks.workflow_run_archive_download_tasks import prepare_workflow_run_archive_download_task
now = datetime.datetime.now(datetime.UTC)
celery_task_id = uuid.uuid4().hex
queued_task = task.model_copy(update={"celery_task_id": celery_task_id, "updated_at": now})
cache.save(queued_task)
try:
prepare_workflow_run_archive_download_task.apply_async(
args=(queued_task.tenant_id, queued_task.download_id),
task_id=celery_task_id,
)
except Exception:
failure_time = datetime.datetime.now(datetime.UTC)
failed_task = queued_task.model_copy(
update={
"status": WorkflowRunArchiveDownloadStatus.FAILED,
"error": "Failed to enqueue archive download task.",
"updated_at": failure_time,
"finished_at": failure_time,
}
)
cache.save(failed_task)
logger.exception("Failed to enqueue workflow run archive download task %s", queued_task.download_id)
return failed_task
return queued_task
@@ -5,8 +5,8 @@ This service archives workflow run logs for paid plan users older than the confi
90 days) to S3-compatible storage.
Archive V2 writes bundle-level Parquet objects. A bundle contains many workflow runs and their related table rows.
Bundle metadata lives in the object-store manifest instead of a database table, so archive/delete/restore does not move
the large-table retention problem into another OLTP table.
Bundle metadata lives in the object-store manifest as the recoverable source of truth. Completed bundles are also
mirrored into a small database index so console listing and download jobs do not list object storage online.
Archive campaigns should use fixed absolute UTC windows for every tenant-prefix/shard execution. Relative windows are
evaluated at process start and are not safe for multi-day rollout because each command would scan a different window.
@@ -29,12 +29,12 @@ import hashlib
import json
import logging
import time
from collections.abc import Callable, Sequence
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from enum import Enum
from threading import Lock
from typing import Any, NotRequired, TypedDict, TypeVar, cast
from typing import Any, NotRequired, TypedDict, cast
import click
import pyarrow as pa
@@ -64,16 +64,20 @@ from repositories.api_workflow_node_execution_repository import DifyAPIWorkflowN
from repositories.api_workflow_run_repository import APIWorkflowRunRepository
from repositories.sqlalchemy_workflow_trigger_log_repository import SQLAlchemyWorkflowTriggerLogRepository
from services.billing_service import BillingService
from services.retention.workflow_run.archive_bundle_index import (
ArchiveBundleManifest,
ArchiveBundleTableManifestEntry,
decode_archive_bundle_manifest,
upsert_archive_bundle_index_from_manifest,
)
from services.retention.workflow_run.constants import (
ARCHIVE_BUNDLE_FORMAT,
ARCHIVE_BUNDLE_INDEX_NAME,
ARCHIVE_BUNDLE_MANIFEST_NAME,
ARCHIVE_BUNDLE_SCHEMA_VERSION,
)
from services.retention.workflow_run.db_retry import is_retryable_db_disconnect, run_with_db_retry
logger = logging.getLogger(__name__)
T = TypeVar("T")
class TableStatsManifestEntry(TypedDict):
@@ -210,8 +214,6 @@ class WorkflowRunArchiver:
"workflow_pause_reasons",
"workflow_trigger_logs",
]
DB_RETRY_ATTEMPTS = 3
DB_RETRY_DELAYS_SECONDS = (1.0, 2.0)
start_from: datetime.datetime | None
end_before: datetime.datetime
@@ -459,40 +461,16 @@ class WorkflowRunArchiver:
"""Fetch a batch of workflow runs to archive."""
repo = self._get_workflow_run_repo()
tenant_ids = list(tenant_scope) if tenant_scope is not None else self.tenant_ids or None
return self._run_with_db_retry(
"workflow run batch fetch",
lambda: repo.get_runs_batch_by_time_range(
start_from=self.start_from,
end_before=self.end_before,
last_seen=last_seen,
batch_size=self.batch_size,
run_types=self.ARCHIVED_TYPE,
tenant_ids=tenant_ids,
tenant_prefixes=None if tenant_ids else self.tenant_prefixes or None,
run_shard_index=self.run_shard_index,
run_shard_total=self.run_shard_total,
),
)
@staticmethod
def _is_retryable_db_disconnect(exc: BaseException) -> bool:
return is_retryable_db_disconnect(exc)
@staticmethod
def _safe_rollback(session: Session, bundle_id: str) -> None:
try:
session.rollback()
except Exception:
logger.warning("Failed to rollback archive session for bundle %s", bundle_id, exc_info=True)
def _run_with_db_retry(self, operation_name: str, operation: Callable[[], T]) -> T:
return run_with_db_retry(
operation_name,
operation,
logger=logger,
attempts=self.DB_RETRY_ATTEMPTS,
delays_seconds=self.DB_RETRY_DELAYS_SECONDS,
return repo.get_runs_batch_by_time_range(
start_from=self.start_from,
end_before=self.end_before,
last_seen=last_seen,
batch_size=self.batch_size,
run_types=self.ARCHIVED_TYPE,
tenant_ids=tenant_ids,
tenant_prefixes=None if tenant_ids else self.tenant_prefixes or None,
run_shard_index=self.run_shard_index,
run_shard_total=self.run_shard_total,
)
def _tenant_scan_scopes(self) -> list[list[str] | None]:
@@ -560,14 +538,16 @@ class WorkflowRunArchiver:
if self.workers == 1 or len(bundle_groups) == 1:
results: list[ArchiveResult] = []
for bundle_runs in bundle_groups:
results.append(self._archive_bundle_with_retry(session_maker, storage, bundle_runs))
with session_maker() as session:
results.append(self._archive_bundle(session, storage, bundle_runs))
return results
results = []
max_workers = min(self.workers, len(bundle_groups))
def archive_in_worker(bundle_runs: Sequence[WorkflowRun]) -> ArchiveResult:
return self._archive_bundle_with_retry(session_maker, storage, bundle_runs)
with session_maker() as session:
return self._archive_bundle(session, storage, bundle_runs)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
futures = [executor.submit(archive_in_worker, bundle_runs) for bundle_runs in bundle_groups]
@@ -575,39 +555,6 @@ class WorkflowRunArchiver:
results.append(future.result())
return results
def _archive_bundle_with_retry(
self,
session_maker: sessionmaker[Session],
storage: ArchiveStorage | None,
runs: Sequence[WorkflowRun],
) -> ArchiveResult:
identity = self._build_bundle_identity(runs)
try:
return self._run_with_db_retry(
f"archive workflow run bundle {identity.bundle_id}",
lambda: self._archive_bundle_once(session_maker, storage, runs),
)
except Exception as exc:
logger.exception("Failed to archive workflow run bundle %s after retries", identity.bundle_id)
return ArchiveResult(
bundle_id=identity.bundle_id,
tenant_id=identity.tenant_id,
object_prefix=identity.object_prefix,
run_count=len(runs),
success=False,
error=str(exc),
)
def _archive_bundle_once(
self,
session_maker: sessionmaker[Session],
storage: ArchiveStorage | None,
runs: Sequence[WorkflowRun],
) -> ArchiveResult:
with session_maker() as session:
return self._archive_bundle(session, storage, runs)
def _archive_bundle(
self,
session: Session,
@@ -634,6 +581,7 @@ class WorkflowRunArchiver:
raise ArchiveStorageNotConfiguredError("Archive storage not configured")
if storage.object_exists(self._get_manifest_object_key(identity)):
self._write_bundle_index(storage, identity)
self._sync_existing_bundle_index(session, storage, identity)
result.success = True
result.skipped = True
result.error = "bundle already archived"
@@ -657,6 +605,7 @@ class WorkflowRunArchiver:
result.run_count = len(runs)
if storage.object_exists(self._get_manifest_object_key(identity)):
self._write_bundle_index(storage, identity)
self._sync_existing_bundle_index(session, storage, identity)
result.success = True
result.skipped = True
result.error = "filtered bundle already archived"
@@ -688,6 +637,8 @@ class WorkflowRunArchiver:
storage.put_object(self._get_table_object_key(identity, table_name), payload)
storage.put_object(self._get_manifest_object_key(identity), manifest_data)
self._merge_bundle_manifest_into_index(storage, identity, [run.id for run in runs])
manifest = decode_archive_bundle_manifest(manifest_data)
upsert_archive_bundle_index_from_manifest(session, manifest, len(manifest_data))
session.commit()
logger.info(
@@ -704,16 +655,30 @@ class WorkflowRunArchiver:
result.success = True
except Exception as e:
if self._is_retryable_db_disconnect(e):
self._safe_rollback(session, identity.bundle_id)
raise
logger.exception("Failed to archive workflow run bundle %s", identity.bundle_id)
result.error = str(e)
self._safe_rollback(session, identity.bundle_id)
session.rollback()
result.elapsed_time = time.time() - start_time
return result
def _sync_existing_bundle_index(
self,
session: Session,
storage: ArchiveStorage,
identity: ArchiveBundleIdentity,
) -> None:
"""Best-effort DB index sync for a bundle whose manifest already exists in archive storage."""
manifest_key = self._get_manifest_object_key(identity)
try:
manifest_data = storage.get_object(manifest_key)
manifest = decode_archive_bundle_manifest(manifest_data)
upsert_archive_bundle_index_from_manifest(session, manifest, len(manifest_data))
session.commit()
except Exception:
session.rollback()
logger.warning("Failed to sync workflow archive bundle index for %s", manifest_key, exc_info=True)
def _lock_runs_for_archive(
self,
session: Session,
@@ -841,9 +806,9 @@ class WorkflowRunArchiver:
identity: ArchiveBundleIdentity,
runs: Sequence[WorkflowRun],
table_stats: list[TableStats],
) -> ArchiveManifestDict:
) -> ArchiveBundleManifest:
"""Generate a manifest for the archived workflow run bundle."""
tables: dict[str, TableStatsManifestEntry] = {
tables: dict[str, ArchiveBundleTableManifestEntry] = {
stat.table_name: {
"row_count": stat.row_count,
"checksum": stat.checksum,
@@ -856,8 +821,8 @@ class WorkflowRunArchiver:
end_before = self.end_before
if end_before is None:
raise ValueError("archive window end must be set")
formatted_end_before = self._format_window_datetime(end_before)
if formatted_end_before is None:
archive_window_end = self._format_window_datetime(end_before)
if archive_window_end is None:
raise ValueError("archive window end must be set")
return ArchiveManifestDict(
schema_version=ARCHIVE_BUNDLE_SCHEMA_VERSION,
@@ -878,7 +843,7 @@ class WorkflowRunArchiver:
archived_at=datetime.datetime.now(datetime.UTC).isoformat(),
campaign_id=self.campaign_id,
archive_window_start=self._format_window_datetime(self.start_from),
archive_window_end=formatted_end_before,
archive_window_end=archive_window_end,
run_shard=identity.shard,
tables=tables,
run_ids=[run.id for run in sorted_runs],
@@ -1,9 +1,10 @@
"""
Maintain V2 workflow-run archive bundles.
Archive V2 keeps bundle metadata in object-store manifests, not in a database table. This module discovers bundles by
listing `manifest.json` objects, uses object-store marker files for delete/restore state, and only touches the database
for source-table validation, deletion, and restoration.
Archive V2 keeps object-store manifests as the recoverable bundle source of truth. This maintenance module still
discovers delete/restore targets by listing `manifest.json` objects and uses object-store marker files for
delete/restore state. The separate database bundle index is intended for console listing and download jobs, not as the
source of truth for destructive maintenance.
Each bundle is processed in its own database transaction. A failed bundle leaves source rows unchanged unless the
transaction has already committed; marker handling makes the next run able to reconcile the common committed-but-marker
@@ -1,66 +0,0 @@
import logging
import time
from collections.abc import Callable
from sqlalchemy.exc import DBAPIError
from sqlalchemy.exc import OperationalError as SQLAlchemyOperationalError
DEFAULT_DB_RETRY_ATTEMPTS = 3
DEFAULT_DB_RETRY_DELAYS_SECONDS = (1.0, 2.0)
_DB_DISCONNECT_PATTERNS = (
"server closed the connection unexpectedly",
"connection already closed",
"closed the connection",
"connection not open",
"terminating connection",
"connection reset",
"broken pipe",
"connection invalidated",
)
def is_retryable_db_disconnect(exc: BaseException) -> bool:
if isinstance(exc, DBAPIError) and exc.connection_invalidated:
return True
if not _is_db_operational_error(exc):
return False
original_exception = exc.orig if isinstance(exc, DBAPIError) else None
message = f"{exc} {original_exception or ''}".lower()
return any(pattern in message for pattern in _DB_DISCONNECT_PATTERNS)
def run_with_db_retry[T](
operation_name: str,
operation: Callable[[], T],
*,
logger: logging.Logger,
attempts: int = DEFAULT_DB_RETRY_ATTEMPTS,
delays_seconds: tuple[float, ...] = DEFAULT_DB_RETRY_DELAYS_SECONDS,
) -> T:
for attempt in range(1, attempts + 1):
try:
return operation()
except Exception as exc:
if not is_retryable_db_disconnect(exc) or attempt == attempts:
raise
delay = delays_seconds[min(attempt - 1, len(delays_seconds) - 1)]
logger.warning(
"Retrying %s after retryable DB disconnect (attempt %s/%s, sleep %.1fs)",
operation_name,
attempt,
attempts,
delay,
exc_info=True,
)
time.sleep(delay)
raise RuntimeError(f"{operation_name} did not complete")
def _is_db_operational_error(exc: BaseException) -> bool:
if isinstance(exc, SQLAlchemyOperationalError):
return True
return exc.__class__.__name__ == "OperationalError" and exc.__class__.__module__.startswith(("psycopg", "psycopg2"))
+7 -6
View File
@@ -6,7 +6,7 @@ import uuid
from datetime import UTC, datetime
from typing import TypedDict, cast
from sqlalchemy import select, update
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
@@ -912,11 +912,12 @@ class SummaryIndexService:
# Disable summary records (don't delete)
now = naive_utc_now()
session.execute(
update(DocumentSegmentSummary)
.where(DocumentSegmentSummary.id.in_(s.id for s in summaries))
.values(enabled=False, disabled_at=now, disabled_by=disabled_by)
)
for summary in summaries:
summary.enabled = False
summary.disabled_at = now
summary.disabled_by = disabled_by
session.add(summary)
session.commit()
logger.info("Disabled %s summary records for dataset %s", len(summaries), dataset.id)
@@ -0,0 +1,18 @@
"""Celery tasks for preparing workflow-run archive downloads."""
import logging
from celery import shared_task
from services.retention.workflow_run.archive_download_preparation import WorkflowRunArchiveDownloadPreparer
logger = logging.getLogger(__name__)
WORKFLOW_RUN_ARCHIVE_DOWNLOAD_QUEUE = "workflow_archive"
@shared_task(queue=WORKFLOW_RUN_ARCHIVE_DOWNLOAD_QUEUE)
def prepare_workflow_run_archive_download_task(tenant_id: str, download_id: str) -> None:
"""Prepare a cached workflow-run archive download in the background."""
logger.info("Preparing workflow run archive download: tenant=%s download_id=%s", tenant_id, download_id)
WorkflowRunArchiveDownloadPreparer().prepare(tenant_id=tenant_id, download_id=download_id)
@@ -6,10 +6,9 @@ from unittest.mock import ANY, MagicMock, patch
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
from sqlalchemy.exc import OperationalError
from models.workflow import WorkflowRunArchiveBundle
from services.retention.workflow_run.archive_paid_plan_workflow_run import (
ArchiveResult,
ArchiveSummary,
WorkflowRunArchiver,
)
@@ -34,30 +33,6 @@ class FakeArchiveStorage:
return sorted(key for key in self.objects if key.startswith(prefix))
def _db_disconnect_error() -> OperationalError:
return OperationalError(
"select 1",
{},
RuntimeError("server closed the connection unexpectedly"),
connection_invalidated=True,
)
def _run(run_id: str = "run-1"):
run = MagicMock()
run.id = run_id
run.tenant_id = "tenant-1"
run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0)
return run
def _session_context(session):
context = MagicMock()
context.__enter__.return_value = session
context.__exit__.return_value = False
return context
class TestWorkflowRunArchiverInit:
def test_start_from_without_end_before_raises(self):
with pytest.raises(ValueError, match="start_from and end_before must be provided together"):
@@ -165,32 +140,6 @@ class TestWorkflowRunArchiverInit:
repo.get_runs_batch_by_time_range.assert_called_once()
assert repo.get_runs_batch_by_time_range.call_args.kwargs["tenant_ids"] == ["tenant-b"]
def test_get_runs_batch_retries_retryable_db_disconnect(self):
repo = MagicMock()
repo.get_runs_batch_by_time_range.side_effect = [_db_disconnect_error(), []]
archiver = WorkflowRunArchiver(workflow_run_repo=repo)
with patch("services.retention.workflow_run.db_retry.time.sleep") as sleep:
runs = archiver._get_runs_batch(None)
assert runs == []
assert repo.get_runs_batch_by_time_range.call_count == 2
sleep.assert_called_once_with(1.0)
def test_get_runs_batch_does_not_retry_non_db_broken_pipe_error(self):
repo = MagicMock()
repo.get_runs_batch_by_time_range.side_effect = RuntimeError("broken pipe")
archiver = WorkflowRunArchiver(workflow_run_repo=repo)
with (
patch("services.retention.workflow_run.db_retry.time.sleep") as sleep,
pytest.raises(RuntimeError, match="broken pipe"),
):
archiver._get_runs_batch(None)
repo.get_runs_batch_by_time_range.assert_called_once()
sleep.assert_not_called()
def test_start_message_includes_shard(self):
archiver = WorkflowRunArchiver(tenant_prefixes=["0"], run_shard_index=1, run_shard_total=4)
@@ -403,72 +352,6 @@ class TestDryRunArchive:
assert summary.table_stats["workflow_app_logs"].size_bytes == 32
class TestArchiveDbRetry:
def test_archive_bundle_groups_retries_with_fresh_session(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
session_maker = MagicMock(
side_effect=[
_session_context(MagicMock(name="session-1")),
_session_context(MagicMock(name="session-2")),
]
)
success = ArchiveResult(
bundle_id=archiver._build_bundle_identity([run]).bundle_id,
tenant_id=run.tenant_id,
object_prefix=archiver._build_bundle_identity([run]).object_prefix,
run_count=1,
success=True,
)
with (
patch.object(archiver, "_archive_bundle", side_effect=[_db_disconnect_error(), success]) as archive_bundle,
patch("services.retention.workflow_run.db_retry.time.sleep") as sleep,
):
results = archiver._archive_bundle_groups(session_maker, MagicMock(), [[run]])
assert results == [success]
assert archive_bundle.call_count == 2
assert session_maker.call_count == 2
sleep.assert_called_once_with(1.0)
def test_archive_bundle_groups_returns_failed_result_after_retry_exhaustion(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
session_maker = MagicMock(
side_effect=[
_session_context(MagicMock(name="session-1")),
_session_context(MagicMock(name="session-2")),
_session_context(MagicMock(name="session-3")),
]
)
with (
patch.object(archiver, "_archive_bundle", side_effect=[_db_disconnect_error()] * 3) as archive_bundle,
patch("services.retention.workflow_run.db_retry.time.sleep") as sleep,
):
results = archiver._archive_bundle_groups(session_maker, MagicMock(), [[run]])
assert len(results) == 1
assert results[0].success is False
assert "server closed the connection unexpectedly" in (results[0].error or "")
assert archive_bundle.call_count == archiver.DB_RETRY_ATTEMPTS
assert session_maker.call_count == archiver.DB_RETRY_ATTEMPTS
assert sleep.call_count == archiver.DB_RETRY_ATTEMPTS - 1
def test_archive_bundle_uses_safe_rollback_when_failure_rolls_back_badly(self):
archiver = WorkflowRunArchiver(days=90, dry_run=True)
session = MagicMock()
session.rollback.side_effect = RuntimeError("rollback failed")
with patch.object(archiver, "_extract_bundle_data", side_effect=RuntimeError("extract failed")):
result = archiver._archive_bundle(session, None, [_run()])
assert result.success is False
assert result.error == "extract failed"
session.rollback.assert_called_once()
class TestArchiveRunIdempotency:
def _index_payload(self, archiver: WorkflowRunArchiver, run_ids: list[str], run) -> tuple[str, bytes]:
identity = archiver._build_bundle_identity([run])
@@ -518,6 +401,37 @@ class TestArchiveRunIdempotency:
assert result.skipped is True
assert result.error == "bundle already archived"
def test_successful_bundle_persists_archive_index(self):
archiver = WorkflowRunArchiver(days=90)
run = MagicMock()
run.id = str(uuid.uuid4())
run.tenant_id = str(uuid.uuid4())
run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0)
session = MagicMock()
session.scalar.return_value = None
storage = MagicMock()
storage.object_exists.return_value = False
table_data = {
"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}],
"workflow_node_executions": [{"id": str(uuid.uuid4()), "workflow_run_id": run.id}],
}
with (
patch.object(archiver, "_lock_runs_for_archive", return_value=[run]),
patch.object(archiver, "_extract_bundle_data", return_value=table_data),
):
result = archiver._archive_bundle(session, storage, [run])
archived_bundle = session.add.call_args.args[0]
assert result.success is True
assert isinstance(archived_bundle, WorkflowRunArchiveBundle)
assert archived_bundle.tenant_id == run.tenant_id
assert archived_bundle.year == 2025
assert archived_bundle.month == 3
assert archived_bundle.workflow_run_count == 1
assert archived_bundle.row_count == 2
session.commit.assert_called_once()
def test_index_skips_all_already_archived_runs(self):
archiver = WorkflowRunArchiver(days=90)
run = MagicMock()
@@ -1,139 +0,0 @@
import datetime
from unittest.mock import MagicMock
import click
import pytest
from sqlalchemy.exc import OperationalError
from commands import retention
def _db_disconnect_error() -> OperationalError:
return OperationalError(
"select 1",
{},
RuntimeError("server closed the connection unexpectedly"),
connection_invalidated=True,
)
def _session_context(session):
context = MagicMock()
context.__enter__.return_value = session
context.__exit__.return_value = False
return context
def test_resolve_archive_tenant_ids_from_plan_uses_explicit_sessions(monkeypatch):
end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC)
sessions = [MagicMock(name="session-a"), MagicMock(name="session-b")]
session_maker = MagicMock(side_effect=[_session_context(sessions[0]), _session_context(sessions[1])])
calls = []
def get_candidate_tenants(session, prefix, *, start_from, end_before):
calls.append((session, prefix, start_from, end_before))
return [f"{prefix}-paid", f"{prefix}-free"]
monkeypatch.setattr(retention, "_get_archive_candidate_tenant_ids_by_prefix", get_candidate_tenants)
monkeypatch.setattr(
retention,
"_filter_paid_workflow_archive_tenant_ids",
lambda tenant_ids: (["a-paid", "b-paid"], ["a-free", "b-free"]),
)
tenant_plan = retention._resolve_archive_tenant_ids_from_plan(
session_maker=session_maker,
tenant_ids=None,
tenant_prefixes=["a", "b"],
start_from=None,
end_before=end_before,
)
assert tenant_plan["archive_tenant_ids"] == ["a-paid", "b-paid"]
assert tenant_plan["paid_tenant_ids"] == ["a-paid", "b-paid"]
assert tenant_plan["unpaid_tenant_ids"] == ["a-free", "b-free"]
assert calls == [
(sessions[0], "a", None, end_before),
(sessions[1], "b", None, end_before),
]
def test_safe_remove_scoped_session_discards_registry_and_disposes_after_remove_error(monkeypatch):
fake_db = MagicMock()
fake_db.session.remove.side_effect = RuntimeError("server closed the connection unexpectedly")
monkeypatch.setattr(retention, "db", fake_db)
retention._safe_remove_scoped_session("archive workflow run command")
fake_db.session.remove.assert_called_once()
fake_db.session.registry.clear.assert_called_once()
fake_db.engine.dispose.assert_called_once()
def test_archive_command_db_retry_retries_retryable_db_disconnect(monkeypatch):
operation = MagicMock(side_effect=[_db_disconnect_error(), "ok"])
sleep = MagicMock()
monkeypatch.setattr("services.retention.workflow_run.db_retry.time.sleep", sleep)
result = retention._run_archive_command_db_retry("archive plan", operation)
assert result == "ok"
assert operation.call_count == 2
sleep.assert_called_once_with(1.0)
def test_archive_plan_prefix_stats_retries_count_query_with_fresh_session(monkeypatch):
end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC)
sessions = [MagicMock(name="session-1"), MagicMock(name="session-2")]
sessions[0].scalar.side_effect = _db_disconnect_error()
sessions[1].scalar.side_effect = [7, 9]
session_maker = MagicMock(side_effect=[_session_context(sessions[0]), _session_context(sessions[1])])
sleep = MagicMock()
monkeypatch.setattr(
retention,
"_get_archive_candidate_tenant_ids_by_prefix",
lambda session, prefix, *, start_from, end_before: [f"{prefix}-tenant"],
)
monkeypatch.setattr("services.retention.workflow_run.db_retry.time.sleep", sleep)
stats = retention._get_archive_plan_prefix_stats(
session_maker,
"a",
start_from=None,
end_before=end_before,
)
assert stats["tenant_ids"] == ["a-tenant"]
assert stats["workflow_runs"] == 7
assert stats["workflow_node_executions"] == 9
assert session_maker.call_count == 2
sleep.assert_called_once_with(1.0)
def test_archive_workflow_runs_raises_click_exception_when_tenant_plan_fails(monkeypatch):
fake_db = MagicMock()
monkeypatch.setattr(retention, "db", fake_db)
monkeypatch.setattr(
retention,
"_resolve_archive_tenant_ids_from_plan",
MagicMock(side_effect=RuntimeError("tenant plan failed")),
)
with pytest.raises(click.ClickException, match="Failed to resolve workflow archive tenant plan"):
retention.archive_workflow_runs.callback(
tenant_ids="tenant-1",
tenant_prefixes=None,
before_days=90,
from_days_ago=None,
to_days_ago=None,
start_from=None,
end_before=None,
batch_size=10000,
workers=1,
run_shard_index=None,
run_shard_total=None,
limit=None,
dry_run=True,
delete_after_archive=False,
)
@@ -75,7 +75,6 @@ def test_dify_config(monkeypatch: pytest.MonkeyPatch):
# default values
assert config.EDITION == "SELF_HOSTED"
assert config.API_COMPRESSION_ENABLED is False
assert config.AGENT_SHELL_ENABLED is True
assert config.SENTRY_TRACES_SAMPLE_RATE == 1.0
assert config.TEMPLATE_TRANSFORM_MAX_LENGTH == 400_000
@@ -111,25 +110,6 @@ def test_http_timeout_defaults(monkeypatch: pytest.MonkeyPatch):
assert config.HTTP_REQUEST_MAX_WRITE_TIMEOUT == 600
def test_internal_files_url_falls_back_to_server_console_api_url(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
monkeypatch.setenv("SERVER_CONSOLE_API_URL", "http://api:5001")
config = DifyConfig(_env_file=None)
assert config.INTERNAL_FILES_URL == "http://api:5001"
def test_internal_files_url_prefers_explicit_value(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
monkeypatch.setenv("INTERNAL_FILES_URL", "http://files-internal:5001")
monkeypatch.setenv("SERVER_CONSOLE_API_URL", "http://api:5001")
config = DifyConfig(_env_file=None)
assert config.INTERNAL_FILES_URL == "http://files-internal:5001"
# NOTE: If there is a `.env` file in your Workspace, this test might not succeed as expected.
# This is due to `pymilvus` loading all the variables from the `.env` file into `os.environ`.
def test_flask_configs(monkeypatch: pytest.MonkeyPatch):
@@ -49,8 +49,6 @@ def make_message():
msg.query = "hello"
msg.re_sign_file_url_answer = ""
msg.user_feedback = MagicMock(rating=None)
msg.total_price = None
msg.currency = None
msg.status = "normal"
msg.error = None
return msg
@@ -34,11 +34,11 @@ class DummyFile:
class DummyToolFile:
def __init__(self, name="test.txt", mimetype="text/plain"):
def __init__(self):
self.id = "file-id"
self.name = name
self.name = "test.txt"
self.size = 10
self.mimetype = mimetype
self.mimetype = "text/plain"
self.original_url = "http://original"
self.user_id = "user-1"
self.tenant_id = "tenant-1"
@@ -56,7 +56,7 @@ class TestPluginUploadFileApi:
mock_get_user,
mock_verify_signature,
):
dummy_file = DummyFile(filename="report.docx", mimetype="application/octet-stream")
dummy_file = DummyFile()
module.request = fake_request(
{
@@ -71,10 +71,7 @@ class TestPluginUploadFileApi:
)
tool_file_manager_instance = mock_tool_file_manager.return_value
tool_file_manager_instance.create_file_by_raw.return_value = DummyToolFile(
name="report.docx",
mimetype="application/octet-stream",
)
tool_file_manager_instance.create_file_by_raw.return_value = DummyToolFile()
mock_tool_file_manager.sign_file.return_value = "signed-url"
@@ -87,12 +84,10 @@ class TestPluginUploadFileApi:
assert result["id"] == "file-id"
assert result["reference"] == build_file_reference(record_id="file-id")
assert result["preview_url"] == "signed-url"
assert result["extension"] == ".docx"
mock_verify_signature.assert_called_once()
assert mock_verify_signature.call_args.kwargs["conversation_id"] == "conversation-1"
tool_file_manager_instance.create_file_by_raw.assert_called_once()
assert tool_file_manager_instance.create_file_by_raw.call_args.kwargs["conversation_id"] == "conversation-1"
mock_tool_file_manager.sign_file.assert_called_once_with(tool_file_id="file-id", extension=".docx")
def test_missing_file(self):
module.request = fake_request(
@@ -318,7 +318,6 @@ class TestPluginDownloadFileRequestApi:
mock_payload.user_id = "user-id"
mock_payload.user_from = "account"
mock_payload.invoke_from = "debugger"
mock_payload.for_external = False
reference = build_file_reference(record_id="tool-file-1")
mock_payload.file.model_dump.return_value = {
"transfer_method": "tool_file",
@@ -334,7 +333,6 @@ class TestPluginDownloadFileRequestApi:
user_from="account",
invoke_from="debugger",
file_mapping={"transfer_method": "tool_file", "reference": reference},
for_external=False,
)
assert result["data"] == {
"filename": "report.pdf",
@@ -37,7 +37,6 @@ from pydantic_ai.messages import (
from clients.agent_backend import (
AgentBackendError,
AgentBackendRunEventAdapter,
AgentBackendRunFailedInternalEvent,
AgentBackendStreamInternalEvent,
FakeAgentBackendRunClient,
FakeAgentBackendScenario,
@@ -55,7 +54,6 @@ from core.app.entities.queue_entities import (
QueueMessageEndEvent,
)
from core.workflow.nodes.agent_v2.ask_human_resume import AskHumanResumeOutcome
from graphon.model_runtime.errors.invoke import InvokeRateLimitError
from models.agent_config_entities import AgentSoulConfig
from models.model import MessageAgentThought
@@ -1041,130 +1039,6 @@ def test_tool_result_without_call_id_matches_unique_open_tool_name(monkeypatch):
assert rows[0].observation == "Knowledge base search results: browser skill"
def test_repeated_tool_calls_without_call_id_or_index_create_distinct_rows(monkeypatch):
fake_session = _FakeDbSession()
monkeypatch.setattr(app_runner_module.db, "session", fake_session)
qm = _FakeQueueManager()
recorder = app_runner_module._AgentProcessRecorder(
dify_context=_dify_ctx(),
message_id="msg-1",
queue_manager=qm, # type: ignore[arg-type]
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_call",
"part": {
"part_kind": "tool-call",
"tool_name": "shell_run",
"args": {"script": "lookup find"},
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_result",
"part": {
"part_kind": "tool-return",
"tool_name": "shell_run",
"content": "find output",
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_call",
"part": {
"part_kind": "tool-call",
"tool_name": "shell_run",
"args": {"script": "lookup out"},
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_result",
"part": {
"part_kind": "tool-return",
"tool_name": "shell_run",
"content": "out output",
},
},
)
)
rows = sorted(fake_session.rows.values(), key=lambda row: row.position)
assert len(rows) == 2
assert rows[0].tool == "shell_run"
assert rows[0].tool_input == '{"script": "lookup find"}'
assert rows[0].observation == "find output"
assert rows[1].tool == "shell_run"
assert rows[1].tool_input == '{"script": "lookup out"}'
assert rows[1].observation == "out output"
def test_repeated_tool_calls_with_placeholder_call_id_and_reused_index_create_distinct_rows(monkeypatch):
fake_session = _FakeDbSession()
monkeypatch.setattr(app_runner_module.db, "session", fake_session)
qm = _FakeQueueManager()
recorder = app_runner_module._AgentProcessRecorder(
dify_context=_dify_ctx(),
message_id="msg-1",
queue_manager=qm, # type: ignore[arg-type]
)
for script, output in (("lookup find", "find output"), ("lookup out", "out output")):
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_call",
"index": 0,
"part": {
"part_kind": "tool-call",
"tool_name": "shell_run",
"tool_call_id": "None",
"args": {"script": script},
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_result",
"part": {
"part_kind": "tool-return",
"tool_name": "shell_run",
"tool_call_id": "None",
"content": output,
},
},
)
)
rows = sorted(fake_session.rows.values(), key=lambda row: row.position)
assert len(rows) == 2
assert rows[0].tool == "shell_run"
assert rows[0].tool_input == '{"script": "lookup find"}'
assert rows[0].observation == "find output"
assert rows[1].tool == "shell_run"
assert rows[1].tool_input == '{"script": "lookup out"}'
assert rows[1].observation == "out output"
def test_prior_session_snapshot_is_threaded_into_request():
prior = CompositorSessionSnapshot(layers=[])
client = FakeAgentBackendRunClient()
@@ -1214,19 +1088,6 @@ def test_failed_run_raises_agent_backend_error():
assert store.saved == []
def test_agent_backend_failure_to_exception_maps_rate_limit_reason():
err = app_runner_module._agent_backend_failure_to_exception(
AgentBackendRunFailedInternalEvent(
run_id="run-1",
error="quota exceeded",
reason="InvokeRateLimitError",
)
)
assert isinstance(err, InvokeRateLimitError)
assert str(err) == "quota exceeded"
def test_stopped_task_cancels_agent_backend_run_and_skips_session_save():
client = _RecordingFakeAgentBackendRunClient()
store = _FakeSessionStore()
@@ -116,12 +116,10 @@ class TestAgentChatAppGeneratorGenerate:
)
thread_obj = mocker.MagicMock()
thread_constructor = mocker.patch(
mocker.patch(
"core.app.apps.agent_chat.app_generator.threading.Thread",
return_value=thread_obj,
)
session = mocker.MagicMock()
mocker.patch("core.app.apps.agent_chat.app_generator.db.session", return_value=session)
mocker.patch(
"core.app.apps.agent_chat.app_generator.AgentChatAppGenerateResponseConverter.convert",
@@ -146,7 +144,6 @@ class TestAgentChatAppGeneratorGenerate:
assert result == {"result": "ok"}
assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1"
assert thread_constructor.call_args.kwargs["kwargs"]["session"] is session
thread_obj.start.assert_called_once()
def test_generate_without_file_config(self, generator, mocker: MockerFixture):
@@ -3,11 +3,10 @@ from unittest.mock import Mock
import pytest
from core.app.apps.base_app_generate_response_converter import AppGenerateResponseConverter
from core.app.entities.queue_entities import QueueErrorEvent
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
from core.errors.error import QuotaExceededError
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError, InvokeRateLimitError
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError
from models.enums import MessageStatus
@@ -69,11 +68,6 @@ class TestBasedGenerateTaskPipeline:
assert error_response.task_id == "task-1"
assert ping_response.task_id == "task-1"
def test_stream_converter_maps_invoke_rate_limit_error(self):
data = AppGenerateResponseConverter._error_to_stream_response(InvokeRateLimitError("quota exceeded"))
assert data == {"code": "rate_limit_error", "status": 429, "message": "quota exceeded"}
def test_handle_output_moderation_when_flagged(self, pipeline):
handler = Mock()
handler.moderation_completion.return_value = ("filtered", True)
@@ -1,32 +0,0 @@
import json
from pytest_mock import MockerFixture
from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType
def test_model_provider_credentials_cache_get_returns_decoded_dict(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.model_provider_cache.redis_client")
cache = ProviderCredentialsCache(
tenant_id="tenant",
identity_id="identity",
cache_type=ProviderCredentialsCacheType.PROVIDER,
)
payload = {"api_key": "secret"}
redis_client_mock.get.return_value = json.dumps(payload).encode("utf-8")
assert cache.get() == payload
def test_model_provider_credentials_cache_get_returns_none_for_invalid_utf8(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.model_provider_cache.redis_client")
cache = ProviderCredentialsCache(
tenant_id="tenant",
identity_id="identity",
cache_type=ProviderCredentialsCacheType.PROVIDER,
)
redis_client_mock.get.return_value = b"\xff"
assert cache.get() is None
@@ -1,24 +0,0 @@
import json
from pytest_mock import MockerFixture
from core.helper.provider_cache import ToolProviderCredentialsCache
def test_provider_credentials_cache_get_returns_decoded_dict(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.provider_cache.redis_client")
cache = ToolProviderCredentialsCache(tenant_id="tenant", provider="provider", credential_id="credential")
payload = {"api_key": "secret"}
redis_client_mock.get.return_value = json.dumps(payload).encode("utf-8")
assert cache.get() == payload
def test_provider_credentials_cache_get_returns_none_for_invalid_utf8(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.provider_cache.redis_client")
cache = ToolProviderCredentialsCache(tenant_id="tenant", provider="provider", credential_id="credential")
redis_client_mock.get.return_value = b"\xff"
assert cache.get() is None
@@ -38,21 +38,6 @@ def test_tool_parameter_cache_get_returns_none_for_invalid_json(mocker: MockerFi
assert cache.get() is None
def test_tool_parameter_cache_get_returns_none_for_invalid_utf8(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.tool_parameter_cache.redis_client")
cache = ToolParameterCache(
tenant_id="tenant",
provider="provider",
tool_name="tool",
cache_type=ToolParameterCacheType.PARAMETER,
identity_id="identity",
)
redis_client_mock.get.return_value = b"\xff"
assert cache.get() is None
def test_tool_parameter_cache_get_returns_none_when_key_is_missing(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.tool_parameter_cache.redis_client")
cache = ToolParameterCache(
@@ -22,7 +22,6 @@ def test_request_download_file_accepts_tool_file_reference() -> None:
assert payload.file.transfer_method == "tool_file"
assert payload.file.reference == reference
assert payload.for_external is True
def test_request_download_file_accepts_remote_url() -> None:
@@ -43,25 +42,6 @@ def test_request_download_file_accepts_remote_url() -> None:
assert payload.file.url == "https://example.com/report.pdf"
def test_request_download_file_accepts_internal_download_request() -> None:
reference = build_file_reference(record_id="tool-file-1")
payload = RequestRequestDownloadFile.model_validate(
{
"tenant_id": "tenant-1",
"user_id": "user-1",
"user_from": "account",
"invoke_from": "debugger",
"file": {
"transfer_method": "tool_file",
"reference": reference,
},
"for_external": False,
}
)
assert payload.for_external is False
def test_request_download_file_rejects_remote_url_without_url() -> None:
with pytest.raises(ValidationError, match="url is required"):
_ = RequestRequestDownloadFile.model_validate(
@@ -60,37 +60,6 @@ def test_create_file_by_raw_stores_file_and_persists_record() -> None:
session.refresh.assert_called_once_with(file_model)
def test_create_file_by_raw_prefers_filename_extension_over_mimetype() -> None:
manager = ToolFileManager()
session = Mock()
session.refresh.side_effect = lambda model: setattr(model, "id", "tf-docx")
def tool_file_factory(**kwargs):
return SimpleNamespace(**kwargs)
with (
patch("core.tools.tool_file_manager.storage") as storage,
patch("core.tools.tool_file_manager.ToolFile", side_effect=tool_file_factory),
patch("core.tools.tool_file_manager.uuid4", return_value=SimpleNamespace(hex="abc")),
_patch_session_factory(session),
):
file_model = manager.create_file_by_raw(
user_id="u1",
tenant_id="t1",
conversation_id="c1",
file_binary=b"docx",
mimetype="application/octet-stream",
filename="report.docx",
)
assert file_model.name == "report.docx"
assert file_model.file_key == "tools/t1/abc.docx"
storage.save.assert_called_once_with("tools/t1/abc.docx", b"docx")
session.add.assert_called_once_with(file_model)
session.commit.assert_called_once()
session.refresh.assert_called_once_with(file_model)
def test_create_file_by_url_downloads_and_persists_record() -> None:
manager = ToolFileManager()
response = Mock()
@@ -119,32 +88,6 @@ def test_create_file_by_url_downloads_and_persists_record() -> None:
session.refresh.assert_called_once_with(file_model)
def test_create_file_by_url_prefers_url_extension_over_mimetype() -> None:
manager = ToolFileManager()
response = Mock()
response.content = b"docx"
response.headers = {"Content-Type": "application/octet-stream"}
response.raise_for_status.return_value = None
session = Mock()
def tool_file_factory(**kwargs):
return SimpleNamespace(**kwargs)
session.refresh.side_effect = lambda model: setattr(model, "id", "tf-docx")
with (
patch("core.tools.tool_file_manager.storage") as storage,
patch("core.tools.tool_file_manager.ToolFile", side_effect=tool_file_factory),
patch("core.tools.tool_file_manager.uuid4", return_value=SimpleNamespace(hex="urlabc")),
_patch_session_factory(session),
patch("core.tools.tool_file_manager.remote_fetcher.make_request", return_value=response),
):
file_model = manager.create_file_by_url("u1", "t1", "https://example.com/report.docx?download=1", "c1")
assert file_model.file_key == "tools/t1/urlabc.docx"
assert file_model.name == "urlabc.docx"
storage.save.assert_called_once_with("tools/t1/urlabc.docx", b"docx")
def test_create_file_by_url_raises_on_timeout() -> None:
manager = ToolFileManager()
@@ -7,10 +7,9 @@ from core.tools.entities.tool_entities import ToolInvokeMessage
class _FakeToolFile:
def __init__(self, mimetype: str, name: str | None):
def __init__(self, mimetype: str):
self.id = "fake-tool-file-id"
self.mimetype = mimetype
self.name = name or "fake-tool-file.bin"
class _FakeToolFileManager:
@@ -39,7 +38,7 @@ class _FakeToolFileManager:
"mimetype": mimetype,
"filename": filename,
}
return _FakeToolFile(mimetype, filename)
return _FakeToolFile(mimetype)
@pytest.fixture(autouse=True)
@@ -90,29 +89,6 @@ def test_transform_tool_invoke_messages_mimetype_key_present_but_none():
assert o.meta["tool_file_id"] == "fake-tool-file-id"
def test_transform_tool_invoke_messages_prefers_filename_extension_over_mimetype():
msg = ToolInvokeMessage(
type=ToolInvokeMessage.MessageType.BLOB,
message=ToolInvokeMessage.BlobMessage(blob=b"docx"),
meta={"mime_type": "application/octet-stream", "filename": "report.docx"},
)
out = list(
mt.ToolFileMessageTransformer.transform_tool_invoke_messages(
messages=_gen([msg]),
user_id="u1",
tenant_id="t1",
conversation_id="c1",
)
)
assert _FakeToolFileManager.last_call is not None
assert _FakeToolFileManager.last_call["filename"] == "report.docx"
assert len(out) == 1
assert isinstance(out[0].message, ToolInvokeMessage.TextMessage)
assert out[0].message.text.endswith(".docx")
def test_transform_tool_invoke_messages_parses_existing_tool_file_link_meta():
msg = ToolInvokeMessage(
type=ToolInvokeMessage.MessageType.IMAGE_LINK,
@@ -1,12 +1,10 @@
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import cast
from unittest.mock import MagicMock, patch
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.layers.ask_human import AskHumanToolResult
from dify_agent.protocol import PydanticAIStreamRunEvent, RunStartedEvent, RunSucceededEvent, RunSucceededEventData
from pydantic_ai.messages import PartDeltaEvent, TextPartDelta
from dify_agent.protocol import RunStartedEvent, RunSucceededEvent, RunSucceededEventData
from clients.agent_backend import (
AgentBackendRunEventAdapter,
@@ -192,30 +190,6 @@ class FileOutputBackendClient(FakeAgentBackendRunClient):
)
class AgentMessageDeltaBackendClient(FakeAgentBackendRunClient):
def _events(self, run_id: str):
created_at = datetime(2026, 1, 1, tzinfo=UTC)
return (
RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at),
PydanticAIStreamRunEvent(
id="2-0",
run_id=run_id,
created_at=created_at,
data=PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="hello ")),
agent_message_delta="hello ",
),
RunSucceededEvent(
id="3-0",
run_id=run_id,
created_at=created_at,
data=RunSucceededEventData(
output={"text": "hello agent"},
session_snapshot=CompositorSessionSnapshot(layers=[]),
),
),
)
def _node(
*,
scenario: FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS,
@@ -303,19 +277,6 @@ def test_agent_node_run_maps_successful_agent_backend_run_to_node_result():
assert layers["llm"]["config"]["credentials"] == "[REDACTED]"
def test_agent_node_run_ignores_agent_message_delta_until_terminal_result():
events = list(_node(agent_backend_client=AgentMessageDeltaBackendClient())._run())
assert len(events) == 1
result = cast(StreamCompletedEvent, events[0]).node_run_result
assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED
assert result.outputs == {"text": "hello agent"}
agent_backend = result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"]
assert agent_backend["status"] == "succeeded"
assert agent_backend["agent_message_delta_count"] == 1
assert agent_backend["agent_message_delta_length"] == len("hello ")
def test_agent_node_run_normalizes_declared_file_output_with_canonical_mapping():
tool_reference = build_file_reference(record_id="tool-file-1")
with patch(
@@ -57,17 +57,6 @@ def test_reback_resolves_tenant_tool_file_to_file():
assert file.extension == ".png"
def test_reback_prefers_filename_extension_over_mimetype():
tf = _seed(mimetype="application/octet-stream", name="report.docx", size=99)
file = reback_tool_file_output(tenant_id=TENANT, tool_file_id=tf)
assert file is not None
assert file.filename == "report.docx"
assert file.mime_type == "application/octet-stream"
assert file.extension == ".docx"
assert file.type == FileType.CUSTOM
def test_reback_other_tenant_returns_none():
tf = _seed()
assert reback_tool_file_output(tenant_id="33333333-3333-3333-3333-333333333333", tool_file_id=tf) is None
@@ -165,20 +165,6 @@ def test_build_from_mapping_accepts_opaque_related_id_for_tool_file(mock_tool_fi
assert file.storage_key == "tool_file.pdf"
def test_build_from_mapping_prefers_tool_filename_extension_over_mimetype(mock_tool_file):
mock_tool_file.name = "report.docx"
mock_tool_file.file_key = "tools/test_tenant_id/file.bin"
mock_tool_file.mimetype = "application/octet-stream"
mapping = tool_file_mapping(file_type="document")
file = build_from_mapping(mapping=mapping, tenant_id=TEST_TENANT_ID)
assert file.extension == ".docx"
assert file.filename == "report.docx"
assert file.mime_type == "application/octet-stream"
assert file.storage_key == "tools/test_tenant_id/file.bin"
@pytest.mark.parametrize(
("file_type", "should_pass", "expected_error"),
[
@@ -227,25 +213,6 @@ def test_build_from_remote_url(mock_http_head):
assert file.size == 2048
def test_build_from_remote_url_prefers_filename_extension_over_mimetype():
mapping = {
"transfer_method": "remote_url",
"url": TEST_REMOTE_URL,
"type": "document",
}
with patch(
"factories.file_factory.builders.get_remote_file_info",
return_value=("application/octet-stream", "report.docx", 99),
):
file = build_from_mapping(mapping=mapping, tenant_id=TEST_TENANT_ID)
assert file.filename == "report.docx"
assert file.extension == ".docx"
assert file.mime_type == "application/octet-stream"
assert file.size == 99
@pytest.mark.parametrize(
("file_type", "should_pass", "expected_error"),
[
@@ -1,6 +1,4 @@
from decimal import Decimal
from fields.message_fields import ExploreMessageListItem, MessageListItem, WebMessageListItem
from fields.message_fields import ExploreMessageListItem, MessageListItem
def _base_kwargs():
@@ -36,33 +34,3 @@ class TestExploreMessageListItem:
# Guard the public service-API contract: the base item must not leak metadata.
payload = MessageListItem(**_base_kwargs()).model_dump(mode="json")
assert "metadata" not in payload
def test_message_list_item_exposes_usage_fields(self):
payload = MessageListItem(
**_base_kwargs(),
message_tokens=7,
answer_tokens=11,
provider_response_latency=1.25,
total_price=Decimal("0.0001234"),
currency="USD",
).model_dump(mode="json")
assert payload["message_tokens"] == 7
assert payload["answer_tokens"] == 11
assert payload["total_tokens"] == 18
assert payload["provider_response_latency"] == 1.25
assert payload["total_price"] == "0.0001234"
assert payload["currency"] == "USD"
def test_web_message_list_item_exposes_usage_and_metadata(self):
payload = WebMessageListItem(
**_base_kwargs(),
metadata={"usage": {"total_tokens": 18}},
message_tokens=7,
answer_tokens=11,
).model_dump(mode="json")
assert payload["metadata"] == {"usage": {"total_tokens": 18}}
assert payload["message_tokens"] == 7
assert payload["answer_tokens"] == 11
assert payload["total_tokens"] == 18
@@ -247,6 +247,32 @@ def test_generate_presigned_url(monkeypatch: pytest.MonkeyPatch):
assert url == "http://signed-url"
def test_generate_presigned_url_with_download_headers(monkeypatch: pytest.MonkeyPatch):
_configure_storage(monkeypatch)
client, _ = _mock_client(monkeypatch)
client.generate_presigned_url.return_value = "http://signed-url"
storage = ArchiveStorage(bucket=BUCKET_NAME)
url = storage.generate_presigned_url(
"key",
expires_in=123,
filename="workflow-run-logs-2025-03.zip",
content_type="application/zip",
)
client.generate_presigned_url.assert_called_once_with(
ClientMethod="get_object",
Params={
"Bucket": "archive-bucket",
"Key": "key",
"ResponseContentDisposition": "attachment; filename*=UTF-8''workflow-run-logs-2025-03.zip",
"ResponseContentType": "application/zip",
},
ExpiresIn=123,
)
assert url == "http://signed-url"
def test_generate_presigned_url_error(monkeypatch: pytest.MonkeyPatch):
_configure_storage(monkeypatch)
client, _ = _mock_client(monkeypatch)
+30 -1
View File
@@ -3,6 +3,7 @@ from unittest.mock import MagicMock
import pytest
from flask import Request
from werkzeug.exceptions import Unauthorized
from werkzeug.wrappers import Response
from constants import COOKIE_NAME_ACCESS_TOKEN, COOKIE_NAME_WEBAPP_ACCESS_TOKEN
@@ -11,10 +12,17 @@ from libs.token import extract_access_token, extract_webapp_access_token, set_cs
class MockRequest:
def __init__(self, headers: dict[str, str], cookies: dict[str, str], args: dict[str, str]):
def __init__(
self,
headers: dict[str, str],
cookies: dict[str, str],
args: dict[str, str],
path: str = "/console/api/test",
):
self.headers: dict[str, str] = headers
self.cookies: dict[str, str] = cookies
self.args: dict[str, str] = args
self.path = path
def test_extract_access_token():
@@ -63,3 +71,24 @@ def test_set_csrf_cookie_includes_domain_when_configured(monkeypatch: pytest.Mon
assert any("csrf_token=abc123" in c for c in cookies)
assert any("Domain=example.com" in c for c in cookies)
assert all("__Host-" not in c for c in cookies)
def test_workflow_run_archive_download_file_bypasses_csrf():
request = cast(
Request,
MockRequest(
headers={},
cookies={},
args={},
path="/console/api/workflow-run-archives/downloads/5923ce20291444af45f0580fb49f1cc9/file",
),
)
token.check_csrf_token(request, "account-1")
def test_non_whitelisted_path_requires_csrf():
request = cast(Request, MockRequest(headers={}, cookies={}, args={}, path="/console/api/test"))
with pytest.raises(Unauthorized):
token.check_csrf_token(request, "account-1")
@@ -0,0 +1,197 @@
import datetime
import json
from collections.abc import Iterator
from typing import cast
from unittest.mock import MagicMock
from services.retention.workflow_run.archive_bundle_index import (
ARCHIVE_BUNDLE_ROOT_PREFIX,
ArchiveBundleManifest,
WorkflowRunArchiveBundleIndexBackfill,
calculate_archive_bundle_index_values,
decode_archive_bundle_manifest,
upsert_archive_bundle_index_from_manifest,
)
from services.retention.workflow_run.constants import ARCHIVE_BUNDLE_FORMAT, ARCHIVE_BUNDLE_SCHEMA_VERSION
TENANT_ID = "1251fe32-c0c7-4fe2-a7bd-a8105267faf5"
BUNDLE_ID = "bundle-a"
OBJECT_PREFIX = (
f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix=1/tenant_id={TENANT_ID}/"
f"year=2025/month=03/shard=00-of-01/bundle={BUNDLE_ID}"
)
MANIFEST_KEY = f"{OBJECT_PREFIX}/manifest.json"
class FakeArchiveStorage:
listed_prefixes: list[str]
objects: dict[str, bytes]
def __init__(self, objects: dict[str, bytes]) -> None:
self.objects = objects
self.listed_prefixes = []
def list_objects(self, prefix: str) -> list[str]:
self.listed_prefixes.append(prefix)
return sorted(key for key in self.objects if key.startswith(prefix))
def get_object(self, key: str) -> bytes:
return self.objects[key]
class FakeSessionContext:
session: MagicMock
def __init__(self, session: MagicMock) -> None:
self.session = session
def __enter__(self) -> MagicMock:
return self.session
def __exit__(self, exc_type: object, exc: object, traceback: object) -> None:
return None
class FakeSessionFactory:
session: MagicMock
def __init__(self, session: MagicMock) -> None:
self.session = session
def __call__(self) -> FakeSessionContext:
return FakeSessionContext(self.session)
class FailingSessionFactory:
def __call__(self) -> Iterator[MagicMock]:
raise AssertionError("dry-run should not open a database session")
def _manifest(*, object_prefix: str = OBJECT_PREFIX, month: int = 3) -> ArchiveBundleManifest:
return ArchiveBundleManifest(
schema_version=ARCHIVE_BUNDLE_SCHEMA_VERSION,
archive_format=ARCHIVE_BUNDLE_FORMAT,
tenant_id=TENANT_ID,
tenant_prefix="1",
year=2025,
month=month,
shard="00-of-01",
bundle_id=BUNDLE_ID,
object_prefix=object_prefix,
workflow_run_count=2,
workflow_node_execution_count=3,
min_created_at="2025-03-01T00:00:00+00:00",
max_created_at="2025-03-02T00:00:00+00:00",
min_run_id="run-a",
max_run_id="run-b",
archived_at="2026-06-25T08:00:00+00:00",
tables={
"workflow_runs": {
"row_count": 2,
"checksum": "checksum-a",
"size_bytes": 100,
"object_key": f"{object_prefix}/workflow_runs.parquet",
},
"workflow_node_executions": {
"row_count": 3,
"checksum": "checksum-b",
"size_bytes": 200,
"object_key": f"{object_prefix}/workflow_node_executions.parquet",
},
},
run_ids=["run-a", "run-b"],
)
def _manifest_bytes(manifest: ArchiveBundleManifest | None = None) -> bytes:
return json.dumps(manifest or _manifest()).encode("utf-8")
def test_decode_and_calculate_archive_bundle_index_values() -> None:
data = _manifest_bytes()
manifest = decode_archive_bundle_manifest(data)
values = calculate_archive_bundle_index_values(manifest, len(data))
assert manifest["tenant_id"] == TENANT_ID
assert values.row_count == 5
assert values.archive_bytes == len(data) + 300
assert values.archived_at == datetime.datetime(2026, 6, 25, 8, 0)
def test_upsert_archive_bundle_index_inserts_new_bundle() -> None:
session = MagicMock()
session.scalar.return_value = None
data = _manifest_bytes()
bundle = upsert_archive_bundle_index_from_manifest(session, decode_archive_bundle_manifest(data), len(data))
assert bundle.tenant_id == TENANT_ID
assert bundle.year == 2025
assert bundle.month == 3
assert bundle.workflow_run_count == 2
assert bundle.row_count == 5
assert bundle.archive_bytes == len(data) + 300
session.add.assert_called_once_with(bundle)
def test_upsert_archive_bundle_index_updates_existing_bundle() -> None:
existing = MagicMock()
session = MagicMock()
session.scalar.return_value = existing
data = _manifest_bytes()
bundle = upsert_archive_bundle_index_from_manifest(session, decode_archive_bundle_manifest(data), len(data))
assert bundle is existing
assert existing.workflow_run_count == 2
assert existing.row_count == 5
assert existing.archive_bytes == len(data) + 300
assert existing.archived_at == datetime.datetime(2026, 6, 25, 8, 0)
session.add.assert_not_called()
def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index() -> None:
storage = FakeArchiveStorage({MANIFEST_KEY: _manifest_bytes()})
session = MagicMock()
session.scalar.return_value = None
backfill = WorkflowRunArchiveBundleIndexBackfill(
storage=cast(MagicMock, storage),
session_factory=cast(MagicMock, FakeSessionFactory(session)),
)
summary = backfill.run(tenant_ids=[TENANT_ID], year=2025, month=3)
assert storage.listed_prefixes == [
f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix=1/tenant_id={TENANT_ID}/year=2025/month=03/"
]
assert summary.manifests_found == 1
assert summary.bundles_processed == 1
assert summary.bundles_upserted == 1
assert summary.bundles_failed == 0
session.add.assert_called_once()
session.commit.assert_called_once()
def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None:
other_month_prefix = OBJECT_PREFIX.replace("month=03", "month=04")
storage = FakeArchiveStorage(
{
MANIFEST_KEY: _manifest_bytes(),
f"{other_month_prefix}/manifest.json": _manifest_bytes(
_manifest(object_prefix=other_month_prefix, month=4)
),
}
)
backfill = WorkflowRunArchiveBundleIndexBackfill(
storage=cast(MagicMock, storage),
session_factory=cast(MagicMock, FailingSessionFactory()),
)
summary = backfill.run(tenant_prefixes=["1"], year=2025, month=3, dry_run=True)
assert storage.listed_prefixes == [f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix=1/"]
assert summary.manifests_found == 1
assert summary.bundles_processed == 1
assert summary.bundles_upserted == 0
assert summary.archive_bytes > 0
@@ -0,0 +1,273 @@
import datetime
import hashlib
import io
import json
import zipfile
from types import SimpleNamespace
from typing import cast
import pyarrow as pa
import pyarrow.parquet as pq
from sqlalchemy.orm import Session, sessionmaker
from libs.archive_storage import ArchiveStorage
from models.workflow import WorkflowRunArchiveBundle
from services.retention.workflow_run.archive_bundle_index import ARCHIVE_BUNDLE_ROOT_PREFIX, ArchiveBundleManifest
from services.retention.workflow_run.archive_download_preparation import (
WorkflowRunArchiveDownloadPreparer,
build_archive_download_storage_key,
)
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
WorkflowRunArchiveDownloadTask,
WorkflowRunArchiveDownloadTaskCache,
build_archive_download_id,
build_pending_archive_download_task,
)
from services.retention.workflow_run.constants import ARCHIVE_BUNDLE_FORMAT, ARCHIVE_BUNDLE_SCHEMA_VERSION
TENANT_ID = "1251fe32-c0c7-4fe2-a7bd-a8105267faf5"
BUNDLE_ID = "bundle-a"
SHARD = "00-of-01"
OBJECT_PREFIX = (
f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix=1/tenant_id={TENANT_ID}/"
f"year=2025/month=03/shard={SHARD}/bundle={BUNDLE_ID}"
)
MANIFEST_KEY = f"{OBJECT_PREFIX}/manifest.json"
class FakeArchiveStorage:
objects: dict[str, bytes]
put_objects: dict[str, bytes]
def __init__(self, objects: dict[str, bytes]) -> None:
self.objects = dict(objects)
self.put_objects = {}
def get_object(self, key: str) -> bytes:
return self.objects[key]
def put_object(self, key: str, data: bytes) -> str:
self.put_objects[key] = data
return hashlib.md5(data).hexdigest()
class FakeTaskCache:
task: WorkflowRunArchiveDownloadTask | None
saved_tasks: list[WorkflowRunArchiveDownloadTask]
def __init__(self, task: WorkflowRunArchiveDownloadTask | None) -> None:
self.task = task
self.saved_tasks = []
def get(self, *, tenant_id: str, download_id: str) -> WorkflowRunArchiveDownloadTask | None:
if self.task and self.task.tenant_id == tenant_id and self.task.download_id == download_id:
return self.task
return None
def save(self, task: WorkflowRunArchiveDownloadTask) -> None:
self.task = task
self.saved_tasks.append(task)
class FakeSessionContext:
bundles: list[WorkflowRunArchiveBundle]
def __init__(self, bundles: list[WorkflowRunArchiveBundle]) -> None:
self.bundles = bundles
def __enter__(self) -> "FakeSessionContext":
return self
def __exit__(self, exc_type: object, exc: object, traceback: object) -> None:
return None
def scalars(self, stmt: object) -> list[WorkflowRunArchiveBundle]:
return self.bundles
class FakeSessionFactory:
bundles: list[WorkflowRunArchiveBundle]
def __init__(self, bundles: list[WorkflowRunArchiveBundle]) -> None:
self.bundles = bundles
def __call__(self) -> FakeSessionContext:
return FakeSessionContext(self.bundles)
def _object_prefix(bundle_id: str = BUNDLE_ID) -> str:
return (
f"{ARCHIVE_BUNDLE_ROOT_PREFIX}tenant_prefix=1/tenant_id={TENANT_ID}/"
f"year=2025/month=03/shard={SHARD}/bundle={bundle_id}"
)
def _bundle(bundle_id: str = BUNDLE_ID) -> WorkflowRunArchiveBundle:
return cast(WorkflowRunArchiveBundle, SimpleNamespace(shard=SHARD, bundle_id=bundle_id))
def _task(bundle_refs: list[tuple[str, str]] | None = None) -> WorkflowRunArchiveDownloadTask:
refs = bundle_refs or [(SHARD, BUNDLE_ID)]
return build_pending_archive_download_task(
tenant_id=TENANT_ID,
requested_by="account-1",
year=2025,
month=3,
bundle_ids=[bundle_id for _, bundle_id in refs],
bundle_refs=refs,
archive_bytes=1024,
download_id=build_archive_download_id(
tenant_id=TENANT_ID,
year=2025,
month=3,
bundle_refs=refs,
),
now=datetime.datetime(2026, 6, 25, 8, 0, tzinfo=datetime.UTC),
)
def _manifest_bytes(table_payloads: dict[str, bytes], *, bundle_id: str = BUNDLE_ID) -> bytes:
object_prefix = _object_prefix(bundle_id)
manifest = ArchiveBundleManifest(
schema_version=ARCHIVE_BUNDLE_SCHEMA_VERSION,
archive_format=ARCHIVE_BUNDLE_FORMAT,
tenant_id=TENANT_ID,
tenant_prefix="1",
year=2025,
month=3,
shard=SHARD,
bundle_id=bundle_id,
object_prefix=object_prefix,
workflow_run_count=2,
workflow_node_execution_count=0,
min_created_at="2025-03-01T00:00:00+00:00",
max_created_at="2025-03-02T00:00:00+00:00",
min_run_id="run-a",
max_run_id="run-b",
archived_at="2026-06-25T08:00:00+00:00",
tables={
table_name: {
"row_count": 1,
"checksum": hashlib.md5(payload).hexdigest(),
"size_bytes": len(payload),
"object_key": f"{object_prefix}/{table_name}.parquet",
}
for table_name, payload in table_payloads.items()
},
run_ids=["run-a", "run-b"],
)
return json.dumps(manifest).encode("utf-8")
def _preparer(
*,
task: WorkflowRunArchiveDownloadTask,
storage: FakeArchiveStorage | None = None,
archive_storage: FakeArchiveStorage | None = None,
download_storage: FakeArchiveStorage | None = None,
cache: FakeTaskCache,
bundles: list[WorkflowRunArchiveBundle] | None = None,
) -> WorkflowRunArchiveDownloadPreparer:
source_storage = archive_storage or storage
target_storage = download_storage or storage
assert source_storage is not None
assert target_storage is not None
return WorkflowRunArchiveDownloadPreparer(
archive_storage=cast(ArchiveStorage, source_storage),
download_storage=cast(ArchiveStorage, target_storage),
cache=cast(WorkflowRunArchiveDownloadTaskCache, cache),
session_factory=cast(sessionmaker[Session], FakeSessionFactory(bundles or [_bundle()])),
)
def _parquet_bytes(records: list[dict[str, object]]) -> bytes:
buffer = io.BytesIO()
pq.write_table(pa.Table.from_pylist(records), buffer)
return buffer.getvalue()
def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() -> None:
bundle_refs = [(SHARD, "bundle-a"), (SHARD, "bundle-b")]
task = _task(bundle_refs)
first_bundle_payloads = {
"workflow_app_logs": _parquet_bytes([{"id": "log-a", "workflow_run_id": "run-a"}]),
"workflow_runs": _parquet_bytes([{"id": "run-a", "status": "succeeded"}]),
}
second_bundle_payloads = {
"workflow_app_logs": _parquet_bytes([{"id": "log-b", "workflow_run_id": "run-b"}]),
"workflow_runs": _parquet_bytes([{"id": "run-b", "status": "failed"}]),
}
archive_storage = FakeArchiveStorage(
{
f"{_object_prefix('bundle-a')}/manifest.json": _manifest_bytes(
first_bundle_payloads,
bundle_id="bundle-a",
),
**{
f"{_object_prefix('bundle-a')}/{table}.parquet": payload
for table, payload in first_bundle_payloads.items()
},
f"{_object_prefix('bundle-b')}/manifest.json": _manifest_bytes(
second_bundle_payloads,
bundle_id="bundle-b",
),
**{
f"{_object_prefix('bundle-b')}/{table}.parquet": payload
for table, payload in second_bundle_payloads.items()
},
}
)
download_storage = FakeArchiveStorage({})
cache = FakeTaskCache(task)
preparer = _preparer(
task=task,
archive_storage=archive_storage,
download_storage=download_storage,
cache=cache,
bundles=[_bundle("bundle-a"), _bundle("bundle-b")],
)
result = preparer.prepare(tenant_id=TENANT_ID, download_id=task.download_id)
assert result is not None
assert result.status == WorkflowRunArchiveDownloadStatus.READY
assert result.storage_key == build_archive_download_storage_key(task)
assert result.file_name == "workflow-run-logs-2025-03.zip"
assert cache.saved_tasks[0].status == WorkflowRunArchiveDownloadStatus.PROCESSING
assert cache.saved_tasks[-1].status == WorkflowRunArchiveDownloadStatus.READY
assert archive_storage.put_objects == {}
archive_payload = download_storage.put_objects[result.storage_key]
with zipfile.ZipFile(io.BytesIO(archive_payload)) as archive:
names = set(archive.namelist())
assert names == {
"workflow-run-logs-2025-03/workflow_app_logs.csv",
"workflow-run-logs-2025-03/workflow_runs.csv",
}
workflow_runs_csv = archive.read("workflow-run-logs-2025-03/workflow_runs.csv").decode("utf-8")
assert workflow_runs_csv.count('"id","status"') == 1
assert '"run-a","succeeded"' in workflow_runs_csv
assert '"run-b","failed"' in workflow_runs_csv
def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch() -> None:
task = _task()
table_payloads = {"workflow_runs": _parquet_bytes([{"id": "run-a", "status": "succeeded"}])}
manifest_data = json.loads(_manifest_bytes(table_payloads).decode("utf-8"))
manifest_data["tables"]["workflow_runs"]["checksum"] = "bad-checksum"
storage = FakeArchiveStorage(
{
MANIFEST_KEY: json.dumps(manifest_data).encode("utf-8"),
f"{OBJECT_PREFIX}/workflow_runs.parquet": table_payloads["workflow_runs"],
}
)
cache = FakeTaskCache(task)
preparer = _preparer(task=task, storage=storage, cache=cache)
result = preparer.prepare(tenant_id=TENANT_ID, download_id=task.download_id)
assert result is not None
assert result.status == WorkflowRunArchiveDownloadStatus.FAILED
assert "checksum mismatch" in (result.error or "")
assert storage.put_objects == {}
@@ -0,0 +1,196 @@
import datetime
import pytest
from services.retention.workflow_run.archive_download_task_cache import (
ARCHIVE_DOWNLOAD_FORMAT_VERSION,
WorkflowRunArchiveDownloadStatus,
WorkflowRunArchiveDownloadTask,
WorkflowRunArchiveDownloadTaskCache,
build_archive_download_id,
build_pending_archive_download_task,
)
class FakeRedis:
store: dict[str, tuple[int | datetime.timedelta, str]]
def __init__(self) -> None:
self.store = {}
def get(self, name: str | bytes) -> bytes | str | None:
key = name.decode("utf-8") if isinstance(name, bytes) else name
item = self.store.get(key)
return item[1] if item else None
def setex(self, name: str | bytes, time: int | datetime.timedelta, value: str) -> object:
key = name.decode("utf-8") if isinstance(name, bytes) else name
self.store[key] = (time, value)
return True
def set(
self,
name: str | bytes,
value: str,
ex: int | None = None,
nx: bool = False,
) -> object:
key = name.decode("utf-8") if isinstance(name, bytes) else name
if nx and key in self.store:
return None
self.store[key] = (ex or 0, value)
return True
def delete(self, *names: str | bytes) -> object:
deleted = 0
for name in names:
key = name.decode("utf-8") if isinstance(name, bytes) else name
if self.store.pop(key, None) is not None:
deleted += 1
return deleted
def test_build_pending_archive_download_task_sets_ephemeral_payload() -> None:
now = datetime.datetime(2026, 6, 25, 8, 0, tzinfo=datetime.UTC)
task = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a", "bundle-b"],
archive_bytes=1024,
ttl_seconds=3600,
download_id="download-1",
now=now,
)
assert task.status == WorkflowRunArchiveDownloadStatus.PENDING
assert task.bundle_count == 2
assert task.expires_at == now + datetime.timedelta(seconds=3600)
def test_build_archive_download_id_is_stable_for_same_bundle_set() -> None:
first = build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("01-of-02", "bundle-b"), ("00-of-02", "bundle-a")],
)
second = build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("00-of-02", "bundle-a"), ("01-of-02", "bundle-b")],
)
assert first == second
assert len(first) == 32
def test_build_archive_download_id_changes_when_content_or_format_changes() -> None:
base = build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("00-of-01", "bundle-a")],
)
changed_bundle = build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("00-of-01", "bundle-b")],
)
changed_format = build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("00-of-01", "bundle-a")],
download_format_version=f"{ARCHIVE_DOWNLOAD_FORMAT_VERSION}-next",
)
assert base != changed_bundle
assert base != changed_format
def test_build_archive_download_id_rejects_empty_bundle_refs() -> None:
with pytest.raises(ValueError, match="bundle_refs must not be empty"):
build_archive_download_id(tenant_id="tenant-1", year=2025, month=3, bundle_refs=[])
def test_archive_download_task_cache_round_trips_with_ttl() -> None:
redis = FakeRedis()
cache = WorkflowRunArchiveDownloadTaskCache(redis)
task = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a"],
archive_bytes=1024,
ttl_seconds=3600,
download_id="download-1",
)
cache.save(task)
restored = cache.get(tenant_id="tenant-1", download_id="download-1")
assert restored == task
ttl, _ = redis.store["workflow_run_archive_download:tenant-1:download-1"]
assert isinstance(ttl, int)
assert 0 < ttl <= 3600
def test_archive_download_task_cache_create_if_absent_is_idempotent() -> None:
redis = FakeRedis()
cache = WorkflowRunArchiveDownloadTaskCache(redis)
first = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a"],
archive_bytes=1024,
ttl_seconds=3600,
download_id="download-1",
)
second = first.model_copy(update={"archive_bytes": 2048})
assert cache.create_if_absent(first) is True
assert cache.create_if_absent(second) is False
restored = cache.get(tenant_id="tenant-1", download_id="download-1")
assert restored == first
def test_archive_download_task_cache_delete_removes_entry() -> None:
redis = FakeRedis()
cache = WorkflowRunArchiveDownloadTaskCache(redis)
task = WorkflowRunArchiveDownloadTask(
download_id="download-1",
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=[],
bundle_count=0,
archive_bytes=0,
status=WorkflowRunArchiveDownloadStatus.FAILED,
error="failed",
created_at=datetime.datetime.now(datetime.UTC),
updated_at=datetime.datetime.now(datetime.UTC),
expires_at=datetime.datetime.now(datetime.UTC) + datetime.timedelta(seconds=3600),
)
cache.save(task)
cache.delete(tenant_id="tenant-1", download_id="download-1")
assert cache.get(tenant_id="tenant-1", download_id="download-1") is None
def test_archive_download_task_cache_ignores_malformed_json() -> None:
redis = FakeRedis()
cache = WorkflowRunArchiveDownloadTaskCache(redis)
redis.setex("workflow_run_archive_download:tenant-1:download-1", 3600, "{")
assert cache.get(tenant_id="tenant-1", download_id="download-1") is None
@@ -0,0 +1,326 @@
import datetime
from types import SimpleNamespace
from typing import cast
from unittest.mock import MagicMock
import pytest
from models.workflow import WorkflowRunArchiveBundle
from services.retention.workflow_run.archive_download_task_cache import (
WorkflowRunArchiveDownloadStatus,
WorkflowRunArchiveDownloadTask,
WorkflowRunArchiveDownloadTaskCache,
build_archive_download_id,
build_pending_archive_download_task,
)
from services.retention.workflow_run.archive_log_service import (
ArchiveDownloadTaskDispatcher,
WorkflowRunArchiveDownloadNotReadyError,
WorkflowRunArchiveNotFoundError,
create_workflow_run_archive_download_task,
get_ready_workflow_run_archive_download_task,
list_workflow_run_archives,
)
class FakeTaskCache:
created_task: WorkflowRunArchiveDownloadTask | None
saved_task: WorkflowRunArchiveDownloadTask | None
existing_task: WorkflowRunArchiveDownloadTask | None
tasks_by_download_id: dict[str, WorkflowRunArchiveDownloadTask]
create_result: bool
def __init__(
self,
*,
create_result: bool = True,
existing_task: WorkflowRunArchiveDownloadTask | None = None,
tasks_by_download_id: dict[str, WorkflowRunArchiveDownloadTask] | None = None,
) -> None:
self.created_task = None
self.saved_task = None
self.existing_task = existing_task
self.tasks_by_download_id = tasks_by_download_id or {}
self.create_result = create_result
def create_if_absent(self, task: WorkflowRunArchiveDownloadTask) -> bool:
self.created_task = task
return self.create_result
def get(self, *, tenant_id: str, download_id: str) -> WorkflowRunArchiveDownloadTask | None:
if self.tasks_by_download_id:
return self.tasks_by_download_id.get(download_id)
return self.existing_task
def save(self, task: WorkflowRunArchiveDownloadTask) -> None:
self.saved_task = task
def _bundle(
*,
shard: str,
bundle_id: str,
archive_bytes: int,
year: int = 2025,
month: int = 3,
workflow_run_count: int = 1,
row_count: int = 9,
archived_at: datetime.datetime | None = None,
) -> WorkflowRunArchiveBundle:
return cast(
WorkflowRunArchiveBundle,
SimpleNamespace(
year=year,
month=month,
shard=shard,
bundle_id=bundle_id,
workflow_run_count=workflow_run_count,
row_count=row_count,
archive_bytes=archive_bytes,
archived_at=archived_at or datetime.datetime(2026, 6, 25, 8, 0),
),
)
def _fake_dispatcher(dispatched_tasks: list[WorkflowRunArchiveDownloadTask]) -> ArchiveDownloadTaskDispatcher:
def dispatch(
task: WorkflowRunArchiveDownloadTask,
cache: WorkflowRunArchiveDownloadTaskCache,
) -> WorkflowRunArchiveDownloadTask:
dispatched_tasks.append(task)
return task.model_copy(update={"celery_task_id": "celery-task-1"})
return dispatch
def test_list_workflow_run_archives_aggregates_month_rows() -> None:
latest = datetime.datetime(2026, 6, 25, 8, 0)
previous = datetime.datetime(2026, 6, 24, 8, 0)
session = MagicMock()
march_download_id = build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("00-of-01", "bundle-a"), ("00-of-01", "bundle-b")],
)
ready_task = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a", "bundle-b"],
bundle_refs=[("00-of-01", "bundle-a"), ("00-of-01", "bundle-b")],
archive_bytes=4096,
download_id=march_download_id,
).model_copy(
update={
"status": WorkflowRunArchiveDownloadStatus.READY,
"file_name": "workflow-run-logs-2025-03.zip",
"storage_key": "workflow-run-archive-downloads/tenant-1/2025/03/download.zip",
"file_size_bytes": 8192,
}
)
cache = FakeTaskCache(tasks_by_download_id={march_download_id: ready_task})
session.scalars.return_value = [
_bundle(
year=2025,
month=3,
shard="00-of-01",
bundle_id="bundle-a",
workflow_run_count=40,
row_count=360,
archive_bytes=1024,
archived_at=previous,
),
_bundle(
year=2025,
month=3,
shard="00-of-01",
bundle_id="bundle-b",
workflow_run_count=60,
row_count=540,
archive_bytes=3072,
archived_at=latest,
),
_bundle(
year=2025,
month=2,
shard="00-of-01",
bundle_id="bundle-c",
workflow_run_count=20,
row_count=180,
archive_bytes=1024,
archived_at=previous,
),
]
result = list_workflow_run_archives(session, "tenant-1", cache=cast(WorkflowRunArchiveDownloadTaskCache, cache))
assert result.summary.archived_month_count == 2
assert result.summary.workflow_run_count == 120
assert result.summary.archive_bytes == 5120
assert result.summary.latest_archived_at == latest
assert result.months[0].year == 2025
assert result.months[0].month == 3
assert result.months[0].bundle_count == 2
assert result.months[0].workflow_run_count == 100
assert result.months[0].row_count == 900
assert result.months[0].download_task == ready_task
assert result.months[1].download_task is None
def test_create_workflow_run_archive_download_task_creates_stable_pending_task() -> None:
session = MagicMock()
session.scalars.return_value = [
_bundle(shard="01-of-02", bundle_id="bundle-b", archive_bytes=2048),
_bundle(shard="00-of-02", bundle_id="bundle-a", archive_bytes=1024),
]
cache = FakeTaskCache()
dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = []
task = create_workflow_run_archive_download_task(
session,
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
cache=cast(WorkflowRunArchiveDownloadTaskCache, cache),
dispatcher=_fake_dispatcher(dispatched_tasks),
)
assert task.download_id == build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("01-of-02", "bundle-b"), ("00-of-02", "bundle-a")],
)
assert task.requested_by == "account-1"
assert task.bundle_ids == ["bundle-b", "bundle-a"]
assert [(ref.shard, ref.bundle_id) for ref in task.bundle_refs] == [
("01-of-02", "bundle-b"),
("00-of-02", "bundle-a"),
]
assert task.archive_bytes == 3072
assert cache.created_task == dispatched_tasks[0]
assert task.celery_task_id == "celery-task-1"
def test_create_workflow_run_archive_download_task_returns_existing_task_when_cache_key_exists() -> None:
session = MagicMock()
session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)]
existing_task = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a"],
archive_bytes=1024,
download_id="existing-download",
).model_copy(update={"celery_task_id": "celery-task-1"})
cache = FakeTaskCache(create_result=False, existing_task=existing_task)
dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = []
task = create_workflow_run_archive_download_task(
session,
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
cache=cast(WorkflowRunArchiveDownloadTaskCache, cache),
dispatcher=_fake_dispatcher(dispatched_tasks),
)
assert task == existing_task
assert cache.saved_task is None
assert dispatched_tasks == []
def test_create_workflow_run_archive_download_task_retries_failed_cached_task() -> None:
session = MagicMock()
session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)]
existing_task = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a"],
bundle_refs=[("00-of-01", "bundle-a")],
archive_bytes=1024,
download_id=build_archive_download_id(
tenant_id="tenant-1",
year=2025,
month=3,
bundle_refs=[("00-of-01", "bundle-a")],
),
).model_copy(update={"status": WorkflowRunArchiveDownloadStatus.FAILED, "error": "failed"})
cache = FakeTaskCache(create_result=False, existing_task=existing_task)
dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = []
task = create_workflow_run_archive_download_task(
session,
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
cache=cast(WorkflowRunArchiveDownloadTaskCache, cache),
dispatcher=_fake_dispatcher(dispatched_tasks),
)
assert task.status == WorkflowRunArchiveDownloadStatus.PENDING
assert task.error is None
assert task.celery_task_id == "celery-task-1"
assert cache.saved_task == dispatched_tasks[0]
def test_create_workflow_run_archive_download_task_rejects_missing_month() -> None:
session = MagicMock()
session.scalars.return_value = []
with pytest.raises(WorkflowRunArchiveNotFoundError):
create_workflow_run_archive_download_task(
session,
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
cache=cast(WorkflowRunArchiveDownloadTaskCache, FakeTaskCache()),
dispatcher=_fake_dispatcher([]),
)
def test_get_ready_workflow_run_archive_download_task_requires_ready_file() -> None:
pending_task = build_pending_archive_download_task(
tenant_id="tenant-1",
requested_by="account-1",
year=2025,
month=3,
bundle_ids=["bundle-a"],
archive_bytes=1024,
download_id="download-1",
)
cache = FakeTaskCache(existing_task=pending_task)
with pytest.raises(WorkflowRunArchiveDownloadNotReadyError):
get_ready_workflow_run_archive_download_task(
tenant_id="tenant-1",
download_id="download-1",
cache=cast(WorkflowRunArchiveDownloadTaskCache, cache),
)
ready_task = pending_task.model_copy(
update={
"status": WorkflowRunArchiveDownloadStatus.READY,
"storage_key": "downloads/download-1.zip",
"file_name": "workflow-run-logs-2025-03.zip",
}
)
cache = FakeTaskCache(existing_task=ready_task)
assert (
get_ready_workflow_run_archive_download_task(
tenant_id="tenant-1",
download_id="download-1",
cache=cast(WorkflowRunArchiveDownloadTaskCache, cache),
)
== ready_task
)
@@ -59,31 +59,6 @@ def test_request_download_url_builds_file_under_bound_scope(
assert result.download_url == "https://files.example.com/x"
def test_request_download_url_supports_internal_download_urls() -> None:
fake_file = MagicMock(filename="report.pdf", mime_type="application/pdf", size=123)
service = FileRequestService(access_controller=MagicMock())
with (
patch("services.file_request_service.bind_file_access_scope", return_value=nullcontext()),
patch.object(service, "_build_file", return_value=fake_file),
patch(
"services.file_request_service.file_helpers.resolve_file_url",
return_value="http://internal-files/report.pdf",
) as resolve_file_url,
):
result = service.request_download_url(
tenant_id="tenant-1",
user_id="user-1",
user_from="account",
invoke_from="debugger",
file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:tool-file-1"},
for_external=False,
)
resolve_file_url.assert_called_once_with(fake_file, for_external=False)
assert result.download_url == "http://internal-files/report.pdf"
def test_request_download_url_rejects_unsupported_files() -> None:
service = FileRequestService(access_controller=MagicMock())
@@ -10,12 +10,9 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy import create_engine, select
from sqlalchemy.orm import sessionmaker
import services.summary_index_service as summary_module
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from models.dataset import DocumentSegmentSummary
from models.enums import SegmentStatus, SummaryStatus
from services.summary_index_service import SummaryIndexService
@@ -656,48 +653,32 @@ def test_generate_summaries_for_document_applies_segment_ids_and_only_parent_chu
session.scalars.assert_called()
def test_disable_summaries_for_segments_updates_sqlite_records() -> None:
dataset = SimpleNamespace(id="dataset-1", indexing_technique=IndexTechniqueType.ECONOMY)
engine = create_engine("sqlite+pysqlite:///:memory:")
DocumentSegmentSummary.__table__.create(engine)
summary_rows = [
{
"id": "sum-1",
"dataset_id": dataset.id,
"document_id": "doc-1",
"chunk_id": "seg-1",
"summary_content": "s",
"summary_index_node_id": "n1",
"status": SummaryStatus.COMPLETED,
"enabled": True,
},
{
"id": "sum-2",
"dataset_id": dataset.id,
"document_id": "doc-1",
"chunk_id": "seg-1",
"summary_content": "s",
"summary_index_node_id": None,
"status": SummaryStatus.COMPLETED,
"enabled": True,
},
]
with engine.begin() as connection:
connection.execute(DocumentSegmentSummary.__table__.insert(), summary_rows)
def test_disable_summaries_for_segments_handles_vector_delete_error(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
summary1 = _summary_record(summary_content="s", node_id="n1")
summary2 = _summary_record(summary_content="s", node_id=None)
session_maker = sessionmaker(bind=engine, expire_on_commit=False)
summary_module.session_factory.configure(engine, expire_on_commit=False)
session = MagicMock()
session.scalars.return_value.all.return_value = [summary1, summary2]
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
monkeypatch.setattr(
summary_module,
"Vector",
MagicMock(return_value=MagicMock(delete_by_ids=MagicMock(side_effect=RuntimeError("boom")))),
)
monkeypatch.setitem(
sys.modules, "libs.datetime_utils", SimpleNamespace(naive_utc_now=MagicMock(return_value=datetime(2024, 1, 1)))
)
SummaryIndexService.disable_summaries_for_segments(dataset, segment_ids=["seg-1"], disabled_by="u")
with session_maker() as session:
summaries = session.scalars(select(DocumentSegmentSummary).order_by(DocumentSegmentSummary.id)).all()
assert [(summary.id, summary.enabled, summary.disabled_by) for summary in summaries] == [
("sum-1", False, "u"),
("sum-2", False, "u"),
]
assert all(summary.disabled_at is not None for summary in summaries)
assert summary1.enabled is False
assert summary1.disabled_by == "u"
session.commit.assert_called_once()
def test_disable_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.MonkeyPatch) -> None:
Generated
+17 -24
View File
@@ -322,15 +322,16 @@ wheels = [
[[package]]
name = "anyio"
version = "4.14.1"
version = "4.11.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "idna" },
{ name = "sniffio" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/3b/72/5562aabb8dd7181e8e860622a38bea08d17842b99ecd4c91f84ac95251b0/anyio-4.14.1.tar.gz", hash = "sha256:8d648a3544c1a700e3ff78615cd679e4c5c3f149904287e73687b2596963629e", size = 254831, upload-time = "2026-06-24T20:56:06.017Z" }
sdist = { url = "https://files.pythonhosted.org/packages/c6/78/7d432127c41b50bccba979505f272c16cbcadcc33645d5fa3a738110ae75/anyio-4.11.0.tar.gz", hash = "sha256:82a8d0b81e318cc5ce71a5f1f8b5c4e63619620b63141ef8c995fa0db95a57c4", size = 219094, upload-time = "2025-09-23T09:19:12.58Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/b0/7b/90df4a0a816d98d6ea26f559d87836d494a2cf1fcf063be67df50a7bcc30/anyio-4.14.1-py3-none-any.whl", hash = "sha256:4e5533c5b8ff0a24f5d7a176cbe6877129cd183893f66b537f8f227d10527d72", size = 124875, upload-time = "2026-06-24T20:56:04.413Z" },
{ url = "https://files.pythonhosted.org/packages/15/b3/9b1a8074496371342ec1e796a96f99c82c945a339cd81a8e73de28b4cf9e/anyio-4.11.0-py3-none-any.whl", hash = "sha256:0287e96f4d26d4149305414d4e3bc32f0dcd0862365a4bddea19d7a1ec38c4fc", size = 109097, upload-time = "2025-09-23T09:19:10.601Z" },
]
[[package]]
@@ -1280,12 +1281,10 @@ wheels = [
[[package]]
name = "dify-agent"
version = "1.16.0rc1"
version = "0.1.0"
source = { editable = "../dify-agent" }
dependencies = [
{ name = "anyio" },
{ name = "httpx" },
{ name = "httpx2" },
{ name = "pydantic" },
{ name = "pydantic-ai-slim" },
{ name = "typer" },
@@ -1294,14 +1293,10 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "aiosqlite", marker = "extra == 'shellctl-server'", specifier = ">=0.21.0,<1.0.0" },
{ name = "anyio", specifier = ">=4.12.1,<5.0.0" },
{ name = "fastapi", marker = "extra == 'server'", specifier = "==0.136.0" },
{ name = "fastapi", marker = "extra == 'shellctl-server'", specifier = "==0.136.0" },
{ name = "graphon", marker = "extra == 'server'", specifier = "==0.5.2" },
{ name = "grpclib", extras = ["protobuf"], marker = "extra == 'grpc'", specifier = ">=0.4.9,<0.5.0" },
{ name = "httpx", specifier = "==0.28.1" },
{ name = "httpx2", specifier = ">=2.5.0,<3.0.0" },
{ name = "jsonschema", marker = "extra == 'server'", specifier = ">=4.23.0,<5.0.0" },
{ name = "jwcrypto", marker = "extra == 'server'", specifier = ">=1.5.6,<2" },
{ name = "logfire", extras = ["fastapi", "httpx", "redis"], marker = "extra == 'server'", specifier = ">=4.37.0,<5.0.0" },
@@ -1311,13 +1306,12 @@ requires-dist = [
{ name = "pydantic-ai-slim", extras = ["anthropic", "google", "openai"], marker = "extra == 'server'", specifier = ">=1.85.1,<2.0.0" },
{ name = "pydantic-settings", marker = "extra == 'server'", specifier = ">=2.12.0,<3.0.0" },
{ name = "redis", marker = "extra == 'server'", specifier = ">=7.4.0,<8.0.0" },
{ name = "sqlmodel", marker = "extra == 'shellctl-server'", specifier = ">=0.0.24,<0.1.0" },
{ name = "shell-session-manager", marker = "extra == 'server'", specifier = "==2.4.0" },
{ name = "typer", specifier = ">=0.16.1,<0.17" },
{ name = "typing-extensions", specifier = ">=4.12.2,<5.0.0" },
{ name = "uvicorn", extras = ["standard"], marker = "extra == 'server'", specifier = "==0.46.0" },
{ name = "uvicorn", extras = ["standard"], marker = "extra == 'shellctl-server'", specifier = "==0.46.0" },
]
provides-extras = ["grpc", "server", "shellctl-server"]
provides-extras = ["grpc", "server"]
[package.metadata.requires-dev]
dev = [
@@ -1338,7 +1332,7 @@ docs = [
[[package]]
name = "dify-api"
version = "1.16.0rc1"
version = "1.15.0"
source = { virtual = "." }
dependencies = [
{ name = "aliyun-log-python-sdk" },
@@ -3289,15 +3283,15 @@ wheels = [
[[package]]
name = "httpcore2"
version = "2.5.0"
version = "2.3.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "h11" },
{ name = "truststore" },
]
sdist = { url = "https://files.pythonhosted.org/packages/47/06/5c12df521b5322fb1114a83d46911b2fbcb8855ddb3a635f11c01a214af5/httpcore2-2.5.0.tar.gz", hash = "sha256:88aa170137c17328d5ac44234f9fd10706466d5fb347f3edac4d39b91137b09d", size = 64808, upload-time = "2026-06-25T14:16:56.472Z" }
sdist = { url = "https://files.pythonhosted.org/packages/e6/34/18f1c596e677962f040284246f393b10a1f8ce440b3a7e69c637d0f1c7ad/httpcore2-2.3.0.tar.gz", hash = "sha256:07327e251560960eea8e969d92d4c6a325feb13cca39e25340731336c3baf924", size = 64300, upload-time = "2026-06-01T13:15:02.998Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c9/a1/7564199d1a8728fe737b0a72e5b3f8d92dfe085a74ddf7cdd83bce5f206d/httpcore2-2.5.0-py3-none-any.whl", hash = "sha256:5ce35188de461d31e8d000bfb8ef8bf22c6c16587a211e5571deaa5e9bdf842a", size = 80330, upload-time = "2026-06-25T14:16:53.634Z" },
{ url = "https://files.pythonhosted.org/packages/c2/dd/3357218c69360d1cecc196c230c9a1d5c9afd5dba362056e23e60a5e64e5/httpcore2-2.3.0-py3-none-any.whl", hash = "sha256:477e9e334f74e5240dcac002e890580f36a57d40ff0fb14cc9655731d23b8415", size = 80024, upload-time = "2026-06-01T13:15:00.001Z" },
]
[[package]]
@@ -3361,18 +3355,17 @@ wheels = [
[[package]]
name = "httpx2"
version = "2.5.0"
version = "2.3.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "anyio" },
{ name = "httpcore2" },
{ name = "idna" },
{ name = "truststore" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d0/e2/b5dedc0cf35aa65de5f541ccd30d2bc1fd7f1d43c9ab09f8ed9a7342317b/httpx2-2.5.0.tar.gz", hash = "sha256:e2df9cb4611021527ff8a675b1c320b610a2ec397acc8d6fe6e91df2d9b33c29", size = 83121, upload-time = "2026-06-25T14:16:57.491Z" }
sdist = { url = "https://files.pythonhosted.org/packages/9f/9a/cca0b9145f13d8ae34b885ae28d403a1469a433abc78e0f94f4ce94e650b/httpx2-2.3.0.tar.gz", hash = "sha256:227e7c41d95a76d4077a52640564132777215fc3394e07b66a3116c33d668fa9", size = 81115, upload-time = "2026-06-01T13:15:04.324Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/31/22/859d8252dad9bc9adee34b52e62cde621ece07b042ccb2ab4da1be46695f/httpx2-2.5.0-py3-none-any.whl", hash = "sha256:3d2d4d9cf4b61f1a1f46a95947cfdb47e80cb56a2f91c6256ac8f58e4891df41", size = 76652, upload-time = "2026-06-25T14:16:55.23Z" },
{ url = "https://files.pythonhosted.org/packages/87/ce/ae2911859847f9ba1d6b23027e53481cbeb50b93234f355a968d300ca2cb/httpx2-2.3.0-py3-none-any.whl", hash = "sha256:6f393663bdf6dbe7fe90118e3eb5b2bd024a675cae0390ac08cec9198812d8b7", size = 74538, upload-time = "2026-06-01T13:15:01.566Z" },
]
[[package]]
@@ -3430,11 +3423,11 @@ wheels = [
[[package]]
name = "idna"
version = "3.18"
version = "3.11"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/cd/63/9496c57188a2ee585e0f1db071d75089a11e98aa86eb99d9d7618fc1edce/idna-3.18.tar.gz", hash = "sha256:ffb385a7e039654cef1ab9ef32c6fafe283c0c0467bba1d9029738ce4a14a848", size = 196711, upload-time = "2026-06-02T14:34:07.794Z" }
sdist = { url = "https://files.pythonhosted.org/packages/6f/6d/0703ccc57f3a7233505399edb88de3cbd678da106337b9fcde432b65ed60/idna-3.11.tar.gz", hash = "sha256:795dafcc9c04ed0c1fb032c2aa73654d8e8c5023a7df64a53f39190ada629902", size = 194582, upload-time = "2025-10-12T14:55:20.501Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1e/5e/d4e9f1a599fb8e573b7b87160658329fbf28d19eac2718f51fc3def3aa5a/idna-3.18-py3-none-any.whl", hash = "sha256:7f952cbe720b688055e3f87de14f5c3e5fdaa8bc3928985c4077ca689de849a2", size = 65455, upload-time = "2026-06-02T14:34:06.319Z" },
{ url = "https://files.pythonhosted.org/packages/0e/61/66938bbb5fc52dbdf84594873d5b51fb1f7c7794e9c0f5bd885f30bc507b/idna-3.11-py3-none-any.whl", hash = "sha256:771a87f49d9defaf64091e6e6fe9c18d4833f140bd19464795bc32d966ca37ea", size = 71008, upload-time = "2025-10-12T14:55:18.883Z" },
]
[[package]]
+4 -12
View File
@@ -20,12 +20,6 @@ DIFY_AGENT_PLUGIN_DAEMON_URL=http://localhost:5002
# API key sent to the Dify plugin daemon.
DIFY_AGENT_PLUGIN_DAEMON_API_KEY=
# Dify API inner endpoints
# Base URL for Dify API inner endpoints used by Agent Stub config/file/drive requests.
DIFY_AGENT_INNER_API_URL=http://localhost:5001
# Must match API/worker INNER_API_KEY_FOR_PLUGIN, not the generic INNER_API_KEY.
DIFY_AGENT_INNER_API_KEY=
# Shell layer
# Base URL for the shellctl server used by the dify.shell layer. Leave empty to disable shell layer use.
DIFY_AGENT_SHELLCTL_ENTRYPOINT=
@@ -36,14 +30,12 @@ DIFY_AGENT_SHELLCTL_AUTH_TOKEN=
# Public Agent Stub URL reachable from shellctl-managed remote machines.
# Use http(s)://.../agent-stub for HTTP or grpc://host:port for gRPC.
# Leave empty to avoid injecting DIFY_AGENT_STUB_* into shell.run jobs.
DIFY_AGENT_STUB_API_BASE_URL=http://localhost:5050/agent-stub
# Optional bind override used only when DIFY_AGENT_STUB_API_BASE_URL uses grpc://.
DIFY_AGENT_STUB_URL=
# Optional bind override used only when DIFY_AGENT_STUB_URL uses grpc://.
DIFY_AGENT_STUB_GRPC_BIND_ADDRESS=
# Server-wide root secret used to derive Agent Stub JWE keys.
# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens.
# Replace this development default in production.
# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))'
DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY
# Required when DIFY_AGENT_STUB_URL is set; must be unpadded base64url for 32 bytes.
DIFY_AGENT_SERVER_SECRET_KEY=
# Shared plugin-daemon HTTP client timeouts and limits.
# Plugin-daemon HTTP connect timeout in seconds.
+2 -2
View File
@@ -6,8 +6,8 @@
# cd /app/api && .venv/bin/uvicorn dify_agent.server.app:app --host 0.0.0.0 --port 5050
#
# Unlike the dify-api image (which only installs the base `dify-agent`
# dependency), this image installs the `[server]` extra, so jwcrypto, fastapi,
# uvicorn, etc. are present and the server can
# dependency), this image installs the `[server]` extra, so jwcrypto,
# shell-session-manager, fastapi, uvicorn, etc. are present and the server can
# actually start. dify-api is intentionally left lean.
# base image
+13 -10
View File
@@ -12,15 +12,19 @@ FROM python:3.12-slim-bookworm AS base
ARG NODE_VERSION=22.22.1
ARG PNPM_VERSION=11.9.0
ARG UV_VERSION=0.8.9
ARG DIFY_AGENT_TOOL_SPEC=.[grpc,shellctl-server]
ARG DIFY_AGENT_TOOL_SPEC=.[grpc]
ARG SHELL_SESSION_MANAGER_TOOL_SPEC=shell-session-manager==2.4.0
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1
PIP_NO_CACHE_DIR=1 \
DIFY_AGENT_STUB_DRIVE_BASE=/mnt/drive \
UV_TOOL_DIR=/opt/dify-agent-tools/envs \
UV_TOOL_BIN_DIR=/opt/dify-agent-tools/bin
ENV PATH="${UV_TOOL_BIN_DIR}:${PATH}"
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
bash \
ca-certificates \
curl \
file \
@@ -58,25 +62,24 @@ WORKDIR /opt/dify-agent
FROM base AS tools
ARG DIFY_AGENT_TOOL_SPEC
ARG SHELL_SESSION_MANAGER_TOOL_SPEC
COPY pyproject.toml uv.lock README.md ./
COPY src ./src
RUN uv export --frozen --no-dev --all-extras --no-emit-project --no-hashes \
> /tmp/dify-agent-constraints.txt \
&& UV_TOOL_DIR=/opt/dify-agent-tools/envs \
UV_TOOL_BIN_DIR=/opt/dify-agent-tools/bin \
uv tool install --force --python /usr/local/bin/python --no-python-downloads \
&& uv tool install --force --python /usr/local/bin/python --no-python-downloads \
--constraints /tmp/dify-agent-constraints.txt --link-mode=copy "${DIFY_AGENT_TOOL_SPEC}" \
&& uv tool install --force --python /usr/local/bin/python --no-python-downloads \
--constraints /tmp/dify-agent-constraints.txt --link-mode=copy "${SHELL_SESSION_MANAGER_TOOL_SPEC}" \
&& rm -f /tmp/dify-agent-constraints.txt
FROM base AS production
ENV PATH="/opt/dify-agent-tools/bin:${PATH}"
COPY --from=tools /opt/dify-agent-tools/envs /opt/dify-agent-tools/envs
COPY --from=tools /opt/dify-agent-tools/bin /opt/dify-agent-tools/bin
COPY --from=tools ${UV_TOOL_DIR} ${UV_TOOL_DIR}
COPY --from=tools ${UV_TOOL_BIN_DIR} ${UV_TOOL_BIN_DIR}
RUN useradd --create-home --shell /bin/sh dify \
&& mkdir -p /mnt/drive \
@@ -75,9 +75,8 @@ See `.example.env` for the full server settings template.
If you plan to run `dify.shell`, also configure `DIFY_AGENT_SHELLCTL_ENTRYPOINT`
and, when shell jobs need to call back with the `dify-agent` command, set
`DIFY_AGENT_STUB_API_BASE_URL`. The supplied default configs include a
development `DIFY_AGENT_SERVER_SECRET_KEY`, but production deployments should
override it with a unique 32-byte base64url value as documented in `.example.env`.
`DIFY_AGENT_STUB_API_BASE_URL` plus a 32-byte base64url
`DIFY_AGENT_SERVER_SECRET_KEY` as documented in `.example.env`.
## Start the Dify Agent server
+3 -5
View File
@@ -42,7 +42,7 @@ also reads `.env` and `dify-agent/.env` when present.
| `DIFY_AGENT_SHELLCTL_AUTH_TOKEN` | empty | Optional bearer token sent to the shellctl server. |
| `DIFY_AGENT_STUB_API_BASE_URL` | empty | Public Agent Stub API base URL reachable from shellctl-managed remote machines. HTTP may be the service root or `/agent-stub`; gRPC must be `grpc://host:port`. Enables `DIFY_AGENT_STUB_*` env injection for user `shell.run` jobs. |
| `DIFY_AGENT_STUB_GRPC_BIND_ADDRESS` | empty | Optional `host:port` bind override used only when `DIFY_AGENT_STUB_API_BASE_URL` uses `grpc://`. |
| `DIFY_AGENT_SERVER_SECRET_KEY` | empty | Security-sensitive server-wide root secret used to derive the JWE encryption key for Agent Stub bearer tokens; required when `DIFY_AGENT_STUB_API_BASE_URL` is set. The supplied default config uses a development value; set a unique unpadded base64url 32-byte secret in production. |
| `DIFY_AGENT_SERVER_SECRET_KEY` | empty | Server-wide root secret used to derive Agent Stub JWE keys; required when `DIFY_AGENT_STUB_API_BASE_URL` is set and must be unpadded base64url for 32 bytes. |
| `DIFY_AGENT_PLUGIN_DAEMON_CONNECT_TIMEOUT` | `10` | Plugin-daemon HTTP connect timeout in seconds. |
| `DIFY_AGENT_PLUGIN_DAEMON_READ_TIMEOUT` | `600` | Plugin-daemon HTTP read timeout in seconds. |
| `DIFY_AGENT_PLUGIN_DAEMON_WRITE_TIMEOUT` | `30` | Plugin-daemon HTTP write timeout in seconds. |
@@ -64,11 +64,9 @@ DIFY_AGENT_INNER_API_URL=http://localhost:5001
DIFY_AGENT_INNER_API_KEY=replace-with-dify-inner-api-key-for-plugin
DIFY_AGENT_SHELLCTL_ENTRYPOINT=http://127.0.0.1:5004
DIFY_AGENT_SHELLCTL_AUTH_TOKEN=replace-with-shellctl-token
# Generate with: python -c 'import base64, secrets; print(base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode())'
DIFY_AGENT_STUB_API_BASE_URL=https://agent.example.com/agent-stub
# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens.
# Replace this development default in production.
# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))'
DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY
DIFY_AGENT_SERVER_SECRET_KEY=replace-with-base64url-32-byte-secret
```
Run records and event streams use the same retention. Status writes refresh the
@@ -56,22 +56,18 @@ with `dify-agent ...`, also enable the Agent Stub:
```env
DIFY_AGENT_STUB_API_BASE_URL=https://agent.example.com/agent-stub
# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens.
# Replace this development default in production.
# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))'
DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY
DIFY_AGENT_SERVER_SECRET_KEY=replace-with-base64url-32-byte-secret
```
HTTP `DIFY_AGENT_STUB_API_BASE_URL` may be either the service root or the
explicit `/agent-stub` API root; the server normalizes the service root to
`/agent-stub`. Other HTTP paths are rejected at startup.
The supplied Docker and `.example.env` configs use a development
`DIFY_AGENT_SERVER_SECRET_KEY`. Override it in production with unpadded base64url
text for exactly 32 decoded bytes. One way to generate it is:
`DIFY_AGENT_SERVER_SECRET_KEY` must be unpadded base64url text for exactly 32
decoded bytes. One way to generate it is:
```bash
python -c 'import secrets; print(secrets.token_urlsafe(32))'
python -c 'import base64, secrets; print(base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode())'
```
## Client request shape
@@ -234,11 +230,12 @@ The provided `docker/local-sandbox/Dockerfile` installs:
- `tmux`, required by `shellctl` to manage shell jobs;
- common shell workspace tools: `git`, `openssh-client`, `jq`, `ripgrep`,
`unzip`, `zip`, `file`, `procps`, and `less`;
- `dify-agent[grpc,shellctl-server]` as a standalone uv tool, which provides
both the Agent Stub client CLI and the built-in `shellctl` CLI/server;
- `shell-session-manager==2.3.1` as a standalone uv tool, which provides the
`shellctl` CLI/server;
- `uv`, so uv shebang scripts with PEP 723 metadata can run inside the shell
workspace and Python CLI tools can be installed with isolated tool
environments;
- `node==22.22.1` and `pnpm==11.9.0`, so JavaScript and TypeScript tooling can
run inside the shell workspace without per-job installation;
- the `dify-agent[grpc]` Agent Stub client CLI as a standalone uv tool;
- a non-root default user named `dify`.
@@ -36,7 +36,6 @@ message FileMapping {
message FileDownloadRequest {
FileMapping file = 1;
optional bool for_external = 2;
}
message FileDownloadResponse {
+3 -13
View File
@@ -1,13 +1,11 @@
[project]
name = "dify-agent"
version = "1.16.0-rc1"
version = "0.1.0"
description = "Add your description here"
readme = "README.md"
requires-python = ">=3.12,<4.0"
dependencies = [
"anyio>=4.12.1,<5.0.0",
"httpx==0.28.1",
"httpx2>=2.5.0,<3.0.0",
"pydantic>=2.12.5,<2.13",
"pydantic-ai-slim>=1.102.0,<2.0.0",
"typer>=0.16.1,<0.17",
@@ -17,9 +15,6 @@ dependencies = [
[project.scripts]
dify-agent = "dify_agent.agent_stub.cli.main:main"
dify-agent-stub-server = "dify_agent.agent_stub.server.cli:main"
shellctl = "shellctl.cli:main"
shellctl-sanitize-pty = "shellctl_runtime.sanitize:main"
shellctl-runner-exit = "shellctl_runtime.runner_exit:main"
[project.optional-dependencies]
grpc = ["grpclib[protobuf]>=0.4.9,<0.5.0", "protobuf>=6.33.5,<7.0.0"]
@@ -32,18 +27,13 @@ server = [
"pydantic-ai-slim[anthropic,google,openai]>=1.85.1,<2.0.0",
"pydantic-settings>=2.12.0,<3.0.0",
"redis>=7.4.0,<8.0.0",
"uvicorn[standard]==0.46.0",
]
shellctl-server = [
"aiosqlite>=0.21.0,<1.0.0",
"fastapi==0.136.0",
"sqlmodel>=0.0.24,<0.1.0",
"shell-session-manager==2.4.0",
"uvicorn[standard]==0.46.0",
]
[tool.setuptools.packages.find]
where = ["src"]
include = ["agenton*", "agenton_collections*", "dify_agent*", "shellctl*", "shellctl_runtime*"]
include = ["agenton*", "agenton_collections*", "dify_agent*"]
[tool.pyright]
include = ["src", "examples", "tests"]
@@ -205,7 +205,6 @@ class DifyStreamedResponse(StreamedResponse):
@override
async def _get_event_iterator(self) -> AsyncIterator[ModelResponseStreamEvent]:
chunk_sequence = 0
async for chunk in self.chunks:
if chunk.delta.usage is not None:
self._usage: RequestUsage = _map_usage(chunk.delta.usage)
@@ -217,10 +216,8 @@ class DifyStreamedResponse(StreamedResponse):
chunk,
self.provider_name_value,
self._embedded_thinking_parser,
chunk_sequence,
):
yield event
chunk_sequence += 1
for event in self._embedded_thinking_parser.flush(self._parts_manager, self.provider_name_value):
yield event
@@ -554,21 +551,11 @@ def _normalize_finish_reason(finish_reason: str) -> FinishReason:
return "error"
def _normalize_tool_call_id(tool_call_id: str | None) -> str | None:
if tool_call_id is None:
return None
normalized = tool_call_id.strip()
if not normalized or normalized.lower() in {"none", "null"}:
return None
return normalized
def _chunk_to_stream_events(
parts_manager: ModelResponsePartsManager,
chunk: LLMResultChunk,
provider_name: str,
embedded_thinking_parser: "_EmbeddedThinkingParser",
chunk_sequence: int,
) -> list[ModelResponseStreamEvent]:
events: list[ModelResponseStreamEvent] = []
message = chunk.delta.message
@@ -584,14 +571,13 @@ def _chunk_to_stream_events(
events.append(parts_manager.handle_part(vendor_part_id=None, part=part))
for index, tool_call in enumerate(message.tool_calls):
tool_call_id = _normalize_tool_call_id(tool_call.id)
vendor_id = tool_call_id or f"chunk-{chunk_sequence}-tool-{index}"
vendor_id = tool_call.id or f"chunk-{chunk.delta.index}-tool-{index}"
events.append(
parts_manager.handle_tool_call_part(
vendor_part_id=vendor_id,
tool_name=tool_call.function.name,
args=tool_call.function.arguments,
tool_call_id=tool_call_id or vendor_id,
tool_call_id=tool_call.id,
provider_name=provider_name,
)
)
@@ -1,6 +1,6 @@
"""Shellctl-backed shell provider adapter for dify-agent.
The built-in shellctl SDK owns the HTTP timeout policy for long-polling
The shell-session-manager SDK owns the HTTP timeout policy for long-polling
shellctl requests. This adapter stays narrowly focused on translating SDK and
transport failures into ``ShellProviderError`` so the shell layer can return
tool observations instead of aborting the agent loop.
@@ -16,9 +16,9 @@ import time
from collections.abc import Awaitable
from collections.abc import Callable
from dataclasses import dataclass
from typing import Protocol, TypeVar, cast
from typing import Protocol, TypeVar
import httpx2 as httpx
import httpx
from dify_agent.adapters.shell.protocols import (
ShellCommandProtocol,
@@ -228,16 +228,8 @@ class ShellctlFileTransfer(ShellFileTransferProtocol):
@dataclass(slots=True)
class ShellctlResource(ShellResourceProtocol):
client: ShellctlClientProtocol
_commands: ShellCommandProtocol
_files: ShellFileTransferProtocol
@property
def commands(self) -> ShellCommandProtocol:
return self._commands
@property
def files(self) -> ShellFileTransferProtocol:
return self._files
commands: ShellCommandProtocol
files: ShellFileTransferProtocol
async def close(self) -> None:
try:
@@ -265,8 +257,8 @@ class ShellctlProvider(ShellProviderProtocol):
)
return ShellctlResource(
client=client,
_commands=ShellctlCommands(client=client),
_files=ShellctlFileTransfer(client=client),
commands=ShellctlCommands(client=client),
files=ShellctlFileTransfer(client=client),
)
@@ -277,15 +269,9 @@ def create_default_shellctl_client_factory(
output_limit: int = _SHELLCTL_OUTPUT_LIMIT_BYTES,
) -> ShellctlClientFactory:
def factory() -> ShellctlClientProtocol:
from shellctl.client import ShellctlClient
from shell_session_manager.shellctl.client import ShellctlClient
return cast(
ShellctlClientProtocol,
cast(
object,
ShellctlClient(entrypoint, token=token, output_limit=output_limit),
),
)
return ShellctlClient(entrypoint, token=token, output_limit=output_limit)
return factory
@@ -143,7 +143,6 @@ def download_file_from_environment(
url=environment.url,
auth_jwe=environment.auth_jwe,
file=file_mapping,
for_external=False,
)
if not hasattr(download_request, "filename") or not isinstance(download_request.filename, str):
raise AgentStubTransferError("signed file download response is missing filename")
@@ -208,10 +207,8 @@ def _request_uploaded_tool_file_download_url(*, url: str, auth_jwe: str, referen
file=AgentStubFileMapping(transfer_method="tool_file", reference=reference),
),
)
if not hasattr(download_request, "download_url") or not isinstance(download_request.download_url, str):
raise AgentStubTransferError("signed file download response is missing download_url")
download_url = download_request.download_url
if not download_url:
if not isinstance(download_url, str) or not download_url:
raise AgentStubTransferError("signed file download response is missing download_url")
return download_url
@@ -133,26 +133,6 @@ def config_skills_push(
Pass a directory such as ./skills/researcher that contains SKILL.md. Other files in that directory are
archived with the skill. Pushing a skill with an existing name replaces that config skill.
Skill directory requirements:
- Each PATH must be one skill directory; the directory basename is the config skill name.
- The directory must contain a top-level SKILL.md.
- SKILL.md must be non-empty UTF-8 Markdown.
- SKILL.md must start with YAML frontmatter matching this schema:
\b
---
name: <non-empty string>
description: <string>
---
- Symlinked files are rejected.
- Dependency/cache folders such as .git, __pycache__, .venv and node_modules should be manually cleared before push.
"""
_run_config_skills_push(paths=paths)
@@ -97,7 +97,6 @@ def request_agent_stub_file_download_sync(
url: str,
auth_jwe: str,
file: AgentStubFileMapping,
for_external: bool = True,
timeout: float | httpx.Timeout = 30.0,
sync_http_client: httpx.Client | None = None,
):
@@ -110,14 +109,12 @@ def request_agent_stub_file_download_sync(
url=endpoint.url,
auth_jwe=auth_jwe,
file=file,
for_external=for_external,
timeout=timeout,
)
return request_agent_stub_file_download_http_sync(
base_url=endpoint.url,
auth_jwe=auth_jwe,
file=file,
for_external=for_external,
timeout=timeout,
sync_http_client=sync_http_client,
)
@@ -124,7 +124,6 @@ def request_agent_stub_file_download_grpc_sync(
url: str,
auth_jwe: str,
file: AgentStubFileMapping,
for_external: bool = True,
timeout: float | httpx.Timeout = 30.0,
):
"""Request one signed download URL through the gRPC Agent Stub endpoint.
@@ -145,7 +144,7 @@ def request_agent_stub_file_download_grpc_sync(
auth_jwe=auth_jwe,
method_name="CreateFileDownloadRequest",
request_factory=lambda runtime: _require_conversions().proto_file_download_request(
runtime.agent_stub_pb2, file=file, for_external=for_external
runtime.agent_stub_pb2, file=file
),
response_parser=lambda response: _require_conversions().file_download_response_from_proto(response),
timeout=timeout,
@@ -117,14 +117,13 @@ def request_agent_stub_file_download_http_sync(
base_url: str,
auth_jwe: str,
file: AgentStubFileMapping,
for_external: bool = True,
timeout: float | httpx.Timeout = 30.0,
sync_http_client: httpx.Client | None = None,
) -> AgentStubFileDownloadResponse:
"""Request one signed download URL from the HTTP Agent Stub endpoint."""
try:
request_model = AgentStubFileDownloadRequest(file=file, for_external=for_external)
request_model = AgentStubFileDownloadRequest(file=file)
except ValidationError as exc:
raise AgentStubValidationError("invalid Agent Stub file download request") from exc
response = _post_agent_stub_json(
@@ -132,7 +131,7 @@ def request_agent_stub_file_download_http_sync(
auth_jwe=auth_jwe,
endpoint_name="file download request",
endpoint_url_factory=agent_stub_file_download_request_url,
request_body=request_model.model_dump_json(exclude_none=True, exclude_defaults=True),
request_body=request_model.model_dump_json(exclude_none=True),
timeout=timeout,
sync_http_client=sync_http_client,
)
@@ -1,6 +1,6 @@
# pyright: reportAttributeAccessIssue=false
# -*- coding: utf-8 -*-
# Generated by the protocol buffer compiler. DO NOT EDIT!
# NO CHECKED-IN PROTOBUF GENCODE
# source: dify/agent/stub/v1/agent_stub.proto
# Protobuf Python Version: 6.33.5
"""Generated protocol buffer code."""
@@ -12,7 +12,12 @@ from google.protobuf import symbol_database as _symbol_database
from google.protobuf.internal import builder as _builder
_runtime_version.ValidateProtobufRuntimeVersion(
_runtime_version.Domain.PUBLIC, 6, 33, 5, "", "dify/agent/stub/v1/agent_stub.proto"
_runtime_version.Domain.PUBLIC,
6,
33,
5,
"",
"dify/agent/stub/v1/agent_stub.proto",
)
# @@protoc_insertion_point(imports)
@@ -20,12 +25,12 @@ _sym_db = _symbol_database.Default()
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(
b'\n#dify/agent/stub/v1/agent_stub.proto\x12\x12\x64ify.agent.stub.v1"O\n\x0e\x43onnectRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\x05\x12\x0c\n\x04\x61rgv\x18\x02 \x03(\t\x12\x15\n\rmetadata_json\x18\x03 \x01(\t"8\n\x0f\x43onnectResponse\x12\x15\n\rconnection_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t"7\n\x11\x46ileUploadRequest\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x10\n\x08mimetype\x18\x02 \x01(\t"(\n\x12\x46ileUploadResponse\x12\x12\n\nupload_url\x18\x01 \x01(\t"f\n\x0b\x46ileMapping\x12\x17\n\x0ftransfer_method\x18\x01 \x01(\t\x12\x16\n\treference\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x10\n\x03url\x18\x03 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_referenceB\x06\n\x04_url"p\n\x13\x46ileDownloadRequest\x12-\n\x04\x66ile\x18\x01 \x01(\x0b\x32\x1f.dify.agent.stub.v1.FileMapping\x12\x19\n\x0c\x66or_external\x18\x02 \x01(\x08H\x00\x88\x01\x01\x42\x0f\n\r_for_external"r\n\x14\x46ileDownloadResponse\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x16\n\tmime_type\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x0c\n\x04size\x18\x03 \x01(\x03\x12\x14\n\x0c\x64ownload_url\x18\x04 \x01(\tB\x0c\n\n_mime_type2\xc0\x02\n\x10\x41gentStubService\x12R\n\x07\x43onnect\x12".dify.agent.stub.v1.ConnectRequest\x1a#.dify.agent.stub.v1.ConnectResponse\x12h\n\x17\x43reateFileUploadRequest\x12%.dify.agent.stub.v1.FileUploadRequest\x1a&.dify.agent.stub.v1.FileUploadResponse\x12n\n\x19\x43reateFileDownloadRequest\x12\'.dify.agent.stub.v1.FileDownloadRequest\x1a(.dify.agent.stub.v1.FileDownloadResponseb\x06proto3'
b'\n#dify/agent/stub/v1/agent_stub.proto\x12\x12\x64ify.agent.stub.v1"O\n\x0e\x43onnectRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\x05\x12\x0c\n\x04\x61rgv\x18\x02 \x03(\t\x12\x15\n\rmetadata_json\x18\x03 \x01(\t"8\n\x0f\x43onnectResponse\x12\x15\n\rconnection_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t"7\n\x11\x46ileUploadRequest\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x10\n\x08mimetype\x18\x02 \x01(\t"(\n\x12\x46ileUploadResponse\x12\x12\n\nupload_url\x18\x01 \x01(\t"f\n\x0b\x46ileMapping\x12\x17\n\x0ftransfer_method\x18\x01 \x01(\t\x12\x16\n\treference\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x10\n\x03url\x18\x03 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_referenceB\x06\n\x04_url"D\n\x13\x46ileDownloadRequest\x12-\n\x04\x66ile\x18\x01 \x01(\x0b\x32\x1f.dify.agent.stub.v1.FileMapping"r\n\x14\x46ileDownloadResponse\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x16\n\tmime_type\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x0c\n\x04size\x18\x03 \x01(\x03\x12\x14\n\x0c\x64ownload_url\x18\x04 \x01(\tB\x0c\n\n_mime_type2\xc0\x02\n\x10\x41gentStubService\x12R\n\x07\x43onnect\x12".dify.agent.stub.v1.ConnectRequest\x1a#.dify.agent.stub.v1.ConnectResponse\x12h\n\x17\x43reateFileUploadRequest\x12%.dify.agent.stub.v1.FileUploadRequest\x1a&.dify.agent.stub.v1.FileUploadResponse\x12n\n\x19\x43reateFileDownloadRequest\x12\'.dify.agent.stub.v1.FileDownloadRequest\x1a(.dify.agent.stub.v1.FileDownloadResponseb\x06proto3'
)
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "dify.agent.stub.v1.agent_stub_pb2", _globals)
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "dify_agent.agent_stub.grpc._generated.agent_stub_pb2", _globals)
if not _descriptor._USE_C_DESCRIPTORS:
DESCRIPTOR._loaded_options = None
_globals["_CONNECTREQUEST"]._serialized_start = 59
@@ -39,9 +44,9 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals["_FILEMAPPING"]._serialized_start = 297
_globals["_FILEMAPPING"]._serialized_end = 399
_globals["_FILEDOWNLOADREQUEST"]._serialized_start = 401
_globals["_FILEDOWNLOADREQUEST"]._serialized_end = 513
_globals["_FILEDOWNLOADRESPONSE"]._serialized_start = 515
_globals["_FILEDOWNLOADRESPONSE"]._serialized_end = 629
_globals["_AGENTSTUBSERVICE"]._serialized_start = 632
_globals["_AGENTSTUBSERVICE"]._serialized_end = 952
_globals["_FILEDOWNLOADREQUEST"]._serialized_end = 469
_globals["_FILEDOWNLOADRESPONSE"]._serialized_start = 471
_globals["_FILEDOWNLOADRESPONSE"]._serialized_end = 585
_globals["_AGENTSTUBSERVICE"]._serialized_start = 588
_globals["_AGENTSTUBSERVICE"]._serialized_end = 908
# @@protoc_insertion_point(module_scope)
@@ -1,69 +1,71 @@
from google.protobuf.internal import containers as _containers
from google.protobuf import descriptor as _descriptor
from google.protobuf import message as _message
from collections.abc import Iterable as _Iterable, Mapping as _Mapping
from typing import ClassVar as _ClassVar, Optional as _Optional, Union as _Union
from __future__ import annotations
DESCRIPTOR: _descriptor.FileDescriptor
from collections.abc import Iterable
class ConnectRequest(_message.Message):
__slots__ = ("protocol_version", "argv", "metadata_json")
PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int]
ARGV_FIELD_NUMBER: _ClassVar[int]
METADATA_JSON_FIELD_NUMBER: _ClassVar[int]
from google.protobuf.message import Message
class ConnectRequest(Message):
protocol_version: int
argv: _containers.RepeatedScalarFieldContainer[str]
argv: list[str]
metadata_json: str
def __init__(self, protocol_version: _Optional[int] = ..., argv: _Optional[_Iterable[str]] = ..., metadata_json: _Optional[str] = ...) -> None: ...
class ConnectResponse(_message.Message):
__slots__ = ("connection_id", "status")
CONNECTION_ID_FIELD_NUMBER: _ClassVar[int]
STATUS_FIELD_NUMBER: _ClassVar[int]
def __init__(
self,
*,
protocol_version: int = ...,
argv: Iterable[str] = ...,
metadata_json: str = ...,
) -> None: ...
class ConnectResponse(Message):
connection_id: str
status: str
def __init__(self, connection_id: _Optional[str] = ..., status: _Optional[str] = ...) -> None: ...
class FileUploadRequest(_message.Message):
__slots__ = ("filename", "mimetype")
FILENAME_FIELD_NUMBER: _ClassVar[int]
MIMETYPE_FIELD_NUMBER: _ClassVar[int]
def __init__(self, *, connection_id: str = ..., status: str = ...) -> None: ...
class FileUploadRequest(Message):
filename: str
mimetype: str
def __init__(self, filename: _Optional[str] = ..., mimetype: _Optional[str] = ...) -> None: ...
class FileUploadResponse(_message.Message):
__slots__ = ("upload_url",)
UPLOAD_URL_FIELD_NUMBER: _ClassVar[int]
def __init__(self, *, filename: str = ..., mimetype: str = ...) -> None: ...
class FileUploadResponse(Message):
upload_url: str
def __init__(self, upload_url: _Optional[str] = ...) -> None: ...
class FileMapping(_message.Message):
__slots__ = ("transfer_method", "reference", "url")
TRANSFER_METHOD_FIELD_NUMBER: _ClassVar[int]
REFERENCE_FIELD_NUMBER: _ClassVar[int]
URL_FIELD_NUMBER: _ClassVar[int]
def __init__(self, *, upload_url: str = ...) -> None: ...
class FileMapping(Message):
transfer_method: str
reference: str
url: str
def __init__(self, transfer_method: _Optional[str] = ..., reference: _Optional[str] = ..., url: _Optional[str] = ...) -> None: ...
class FileDownloadRequest(_message.Message):
__slots__ = ("file", "for_external")
FILE_FIELD_NUMBER: _ClassVar[int]
FOR_EXTERNAL_FIELD_NUMBER: _ClassVar[int]
def __init__(self, *, transfer_method: str = ..., reference: str = ..., url: str = ...) -> None: ...
def HasField(self, field_name: str) -> bool: ...
class FileDownloadRequest(Message):
file: FileMapping
for_external: bool
def __init__(self, file: _Optional[_Union[FileMapping, _Mapping]] = ..., for_external: _Optional[bool] = ...) -> None: ...
class FileDownloadResponse(_message.Message):
__slots__ = ("filename", "mime_type", "size", "download_url")
FILENAME_FIELD_NUMBER: _ClassVar[int]
MIME_TYPE_FIELD_NUMBER: _ClassVar[int]
SIZE_FIELD_NUMBER: _ClassVar[int]
DOWNLOAD_URL_FIELD_NUMBER: _ClassVar[int]
def __init__(self, *, file: FileMapping | None = ...) -> None: ...
class FileDownloadResponse(Message):
filename: str
mime_type: str
size: int
download_url: str
def __init__(self, filename: _Optional[str] = ..., mime_type: _Optional[str] = ..., size: _Optional[int] = ..., download_url: _Optional[str] = ...) -> None: ...
def __init__(
self,
*,
filename: str = ...,
mime_type: str = ...,
size: int = ...,
download_url: str = ...,
) -> None: ...
def HasField(self, field_name: str) -> bool: ...
@@ -98,19 +98,13 @@ def file_download_request_from_proto(message: agent_stub_pb2.FileDownloadRequest
"reference": message.file.reference if message.file.HasField("reference") else None,
"url": message.file.url if message.file.HasField("url") else None,
}
return AgentStubFileDownloadRequest.model_validate(
{
"file": file_mapping_kwargs,
"for_external": message.for_external if message.HasField("for_external") else True,
}
)
return AgentStubFileDownloadRequest.model_validate({"file": file_mapping_kwargs})
def proto_file_download_request(
pb2_module,
*,
file: AgentStubFileMapping,
for_external: bool = True,
) -> agent_stub_pb2.FileDownloadRequest:
"""Build one protobuf file-download request from the public DTO."""
mapping = pb2_module.FileMapping(transfer_method=file.transfer_method)
@@ -118,9 +112,7 @@ def proto_file_download_request(
mapping.reference = file.reference
if file.url is not None:
mapping.url = file.url
request = pb2_module.FileDownloadRequest(file=mapping)
request.for_external = for_external
return request
return pb2_module.FileDownloadRequest(file=mapping)
def file_download_response_from_proto(message: agent_stub_pb2.FileDownloadResponse) -> AgentStubFileDownloadResponse:
@@ -249,7 +249,6 @@ class AgentStubFileDownloadRequest(BaseModel):
"""Request body for one signed download URL allocation."""
file: AgentStubFileMapping
for_external: bool = True
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid")
@@ -161,8 +161,6 @@ class DifyApiAgentStubFileRequestHandler:
"invoke_from": execution_context.invoke_from,
"file": request.file.model_dump(mode="json", exclude_none=True),
}
if request.for_external is False:
payload["for_external"] = False
data = await self._post_inner_api("/inner/api/download/file/request", payload)
try:
return AgentStubFileDownloadResponse.model_validate(data)
@@ -141,18 +141,7 @@ shell_run script rules:
Tips:
- Python 3.12, uv, pip, Node.js, pnpm, and pnx are preinstalled in the local sandbox.
- For one-off Python dependencies, prefer a uv script with a PEP 723 dependency header or:
`uv run --with <package> python <script-or--c>`.
- For reusable Python CLI tools, use `uv tool install <tool>`; installed commands land in `$HOME/.local/bin`.
Run them by full path or add `$HOME/.local/bin` to PATH in the command that needs them.
- `python3 -m pip install --user <package>` also installs into `$HOME/.local`; add `$HOME/.local/bin` to PATH
when you need console scripts.
- For reusable Node.js CLIs, use user-level global installs:
`PNPM_HOME=$HOME/.local/share/pnpm PATH=$HOME/.local/share/pnpm/bin:$PATH pnpm add -g <package>`.
Installed commands land in `$PNPM_HOME/bin`; run them by full path or with the same PATH prefix.
- For one-off Node.js CLIs, prefer `pnx <command> [args]`.
- Do not install new packages into system or image tool paths such as `/usr/local`, `/usr`, or `/opt/dify-agent-tools`.
- When using Python, prefer a uv script with a PEP 723 dependency header.
- If you need MCP, install the MCP server in the shell environment and start that server when you use it.
Example shell_run script:
+3 -26
View File
@@ -31,14 +31,13 @@ both the JSON-safe final output or deferred tool call and the session snapshot;
there are no separate output or snapshot events to correlate.
"""
from collections.abc import AsyncIterable, Callable, Mapping
from collections.abc import AsyncIterable, Callable
from collections import Counter
from dataclasses import dataclass
from typing import Any, Literal, Protocol, cast, runtime_checkable
import httpx
from pydantic import JsonValue, TypeAdapter
from pydantic_ai.exceptions import ModelHTTPError
from pydantic_ai.messages import AgentStreamEvent, PartDeltaEvent, PartStartEvent, TextPart, TextPartDelta
from pydantic_ai.output import OutputSpec
from pydantic_ai.tools import DeferredToolRequests, DeferredToolResults
@@ -105,28 +104,6 @@ class AgentRunValidationError(ValueError):
"""Raised when a run request is valid JSON but cannot execute."""
def _run_failed_error_payload(exc: Exception) -> tuple[str, str | None]:
"""Return the public failed-run error text and structured reason."""
message = str(exc) or type(exc).__name__
reason: str | None = None
if isinstance(exc, ModelHTTPError):
body = exc.body
if isinstance(body, Mapping):
body_message = body.get("message")
if isinstance(body_message, str) and body_message:
message = body_message
error_type = body.get("error_type")
if isinstance(error_type, str) and error_type:
reason = error_type
if reason is None and exc.status_code == 429:
reason = "InvokeRateLimitError"
return message, reason
def _has_model_layer(request: CreateRunRequest) -> bool:
"""Return whether the public composition includes the reserved model layer."""
return any(layer.name == DIFY_AGENT_MODEL_LAYER_ID for layer in request.composition.layers)
@@ -188,8 +165,8 @@ class AgentRunRunner:
try:
outcome = await self._run_agent()
except Exception as exc:
message, reason = _run_failed_error_payload(exc)
_ = await emit_run_failed(self.sink, run_id=self.run_id, error=message, reason=reason)
message = str(exc) or type(exc).__name__
_ = await emit_run_failed(self.sink, run_id=self.run_id, error=message)
await self.sink.update_status(self.run_id, "failed", message)
raise
-124
View File
@@ -1,124 +0,0 @@
"""Public shellctl package exports.
This package stays lazy on purpose. Hot-path runtime helpers live outside the
`shellctl` package, and importing this package root should
not pull the full client/server/public DTO surface unless a caller explicitly
asks for those exports.
"""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from shellctl.client import (
ShellctlClient,
ShellctlClientError,
)
from shellctl.shared import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_BASE_URL,
DEFAULT_BASE_URL_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
DeleteJobResponse,
HealthResponse,
InputJobRequest,
JobInfo,
JobResult,
JobStatusName,
JobStatusView,
ListJobsResponse,
RunJobRequest,
TerminalSize,
TerminateJobRequest,
WaitJobRequest,
generate_job_id,
read_output_window,
tail_output_window,
)
__all__ = [
"DEFAULT_AUTH_TOKEN_ENV",
"DEFAULT_BASE_URL",
"DEFAULT_BASE_URL_ENV",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS",
"DEFAULT_GC_INTERVAL_SECONDS",
"DEFAULT_IDLE_FLUSH_SECONDS",
"DEFAULT_LIST_LIMIT",
"DEFAULT_OUTPUT_LIMIT_BYTES",
"DEFAULT_TERMINAL_COLS",
"DEFAULT_TERMINAL_ROWS",
"DEFAULT_TERMINATE_GRACE_SECONDS",
"DEFAULT_TIMEOUT_SECONDS",
"DeleteJobResponse",
"HealthResponse",
"InputJobRequest",
"JobInfo",
"JobResult",
"JobStatusName",
"JobStatusView",
"ListJobsResponse",
"RunJobRequest",
"ShellctlClient",
"ShellctlClientError",
"TerminalSize",
"TerminateJobRequest",
"WaitJobRequest",
"generate_job_id",
"read_output_window",
"tail_output_window",
]
_EXPORTS = {
"ShellctlClient": "shellctl.client",
"ShellctlClientError": "shellctl.client",
"DEFAULT_AUTH_TOKEN_ENV": "shellctl.shared",
"DEFAULT_BASE_URL": "shellctl.shared",
"DEFAULT_BASE_URL_ENV": "shellctl.shared",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS": "shellctl.shared",
"DEFAULT_GC_INTERVAL_SECONDS": "shellctl.shared",
"DEFAULT_IDLE_FLUSH_SECONDS": "shellctl.shared",
"DEFAULT_LIST_LIMIT": "shellctl.shared",
"DEFAULT_OUTPUT_LIMIT_BYTES": "shellctl.shared",
"DEFAULT_TERMINAL_COLS": "shellctl.shared",
"DEFAULT_TERMINAL_ROWS": "shellctl.shared",
"DEFAULT_TERMINATE_GRACE_SECONDS": "shellctl.shared",
"DEFAULT_TIMEOUT_SECONDS": "shellctl.shared",
"DeleteJobResponse": "shellctl.shared",
"HealthResponse": "shellctl.shared",
"InputJobRequest": "shellctl.shared",
"JobInfo": "shellctl.shared",
"JobResult": "shellctl.shared",
"JobStatusName": "shellctl.shared",
"JobStatusView": "shellctl.shared",
"ListJobsResponse": "shellctl.shared",
"RunJobRequest": "shellctl.shared",
"TerminalSize": "shellctl.shared",
"TerminateJobRequest": "shellctl.shared",
"WaitJobRequest": "shellctl.shared",
"generate_job_id": "shellctl.shared",
"read_output_window": "shellctl.shared",
"tail_output_window": "shellctl.shared",
}
def __getattr__(name: str) -> Any:
if name not in _EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module = import_module(_EXPORTS[name])
value = getattr(module, name) # noqa: no-new-getattr lazy export proxy
globals()[name] = value
return value
def __dir__() -> list[str]:
return sorted(set(globals()) | set(__all__))
-587
View File
@@ -1,587 +0,0 @@
"""Typer CLI for network-backed shellctl commands.
Job-management commands in this module intentionally stay on the SDK side of
the boundary: they parse CLI options, call `ShellctlClient`, and render compact
JSON. That keeps `shellctl --help` and `shellctl run --help` free of FastAPI,
SQLAlchemy, tmux, and local runtime bootstrap imports.
Only `serve` lazily imports server-side modules when that subcommand is
actually invoked.
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import NoReturn
import anyio
import httpx2 as httpx
import typer
from pydantic import BaseModel, ValidationError
from shellctl.client import ShellctlClient, ShellctlClientError
from shellctl.shared.constants import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_BASE_URL,
DEFAULT_BASE_URL_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
MAX_LIST_LIMIT,
MAX_OUTPUT_LIMIT_BYTES,
)
from shellctl.shared.schemas import (
DeleteJobResponse,
HealthResponse,
JobInfo,
JobResult,
JobStatusName,
JobStatusView,
RunJobRequest,
TerminalSize,
)
cli = typer.Typer(
no_args_is_help=True,
pretty_exceptions_enable=False,
rich_markup_mode=None,
)
@cli.command("health")
def health_command(
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help="Accepted for CLI consistency but ignored because /healthz is public.",
),
) -> None:
"""Call the public health endpoint and report JSON."""
del auth_token
async def action(client: ShellctlClient) -> HealthResponse:
return await client.health()
_run_client_action(
base_url=base_url,
auth_token=None,
action=action,
emit=_emit_model,
)
@cli.command("run")
def run_command(
script: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
cwd: Path | None = typer.Option(None, "--cwd"),
env: list[str] | None = typer.Option(None, "--env"),
timeout: float = typer.Option(DEFAULT_TIMEOUT_SECONDS, "--timeout"),
output_limit: int = typer.Option(DEFAULT_OUTPUT_LIMIT_BYTES, "--output-limit"),
idle_flush_seconds: float = typer.Option(
DEFAULT_IDLE_FLUSH_SECONDS,
"--idle-flush-seconds",
),
cols: int | None = typer.Option(None, "--cols"),
rows: int | None = typer.Option(None, "--rows"),
) -> None:
"""Create a job through the running shellctl server."""
request = _build_model(
RunJobRequest,
script=script,
cwd=str(cwd) if cwd is not None else None,
env=_parse_env(env),
terminal=_terminal_size(cols=cols, rows=rows),
timeout=timeout,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
)
async def action(client: ShellctlClient) -> JobResult:
return await client.run(
request.script,
cwd=request.cwd,
env=request.env,
timeout=request.timeout,
terminal=request.terminal,
)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
action=action,
emit=_emit_model,
)
@cli.command("wait")
def wait_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
offset: int = typer.Option(..., "--offset"),
timeout: float = typer.Option(DEFAULT_TIMEOUT_SECONDS, "--timeout"),
output_limit: int = typer.Option(DEFAULT_OUTPUT_LIMIT_BYTES, "--output-limit"),
idle_flush_seconds: float = typer.Option(
DEFAULT_IDLE_FLUSH_SECONDS,
"--idle-flush-seconds",
),
) -> None:
"""Wait for incremental output, completion, truncation, or timeout."""
async def action(client: ShellctlClient) -> JobResult:
return await client.wait(job_id, offset=offset, timeout=timeout)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
action=action,
emit=_emit_model,
)
@cli.command("status")
def status_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
) -> None:
"""Materialize the current status view for one job."""
async def action(client: ShellctlClient) -> JobStatusView:
return await client.status(job_id)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_model,
)
@cli.command("list")
def list_command(
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
status: JobStatusName | None = typer.Option(None, "--status"),
limit: int = typer.Option(DEFAULT_LIST_LIMIT, "--limit", min=1, max=MAX_LIST_LIMIT),
) -> None:
"""List recent jobs, optionally filtered by lifecycle status."""
async def action(client: ShellctlClient) -> list[JobInfo]:
return await client.list_jobs(status=status, limit=limit)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_job_list,
)
@cli.command("input")
def input_command(
job_id: str = typer.Argument(...),
text: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
offset: int = typer.Option(..., "--offset"),
timeout: float = typer.Option(DEFAULT_TIMEOUT_SECONDS, "--timeout"),
output_limit: int = typer.Option(DEFAULT_OUTPUT_LIMIT_BYTES, "--output-limit"),
idle_flush_seconds: float = typer.Option(
DEFAULT_IDLE_FLUSH_SECONDS,
"--idle-flush-seconds",
),
) -> None:
"""Send text input to a running job and wait for the next result window."""
async def action(client: ShellctlClient) -> JobResult:
return await client.input(job_id, text, offset=offset, timeout=timeout)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
action=action,
emit=_emit_model,
)
@cli.command("tail")
def tail_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
output_limit: int = typer.Option(
DEFAULT_OUTPUT_LIMIT_BYTES,
"--output-limit",
min=1,
max=MAX_OUTPUT_LIMIT_BYTES,
),
) -> None:
"""Read a UTF-8-safe output tail for one job."""
async def action(client: ShellctlClient) -> JobResult:
return await client.tail(job_id)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
action=action,
emit=_emit_model,
)
@cli.command("terminate")
def terminate_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
grace_seconds: float = typer.Option(
DEFAULT_TERMINATE_GRACE_SECONDS,
"--grace-seconds",
),
) -> None:
"""Terminate a job and return its materialized status."""
async def action(client: ShellctlClient) -> JobStatusView:
return await client.terminate(job_id, grace_seconds=grace_seconds)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_model,
)
@cli.command("delete")
def delete_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
force: bool = typer.Option(False, "--force"),
grace_seconds: float = typer.Option(
DEFAULT_TERMINATE_GRACE_SECONDS,
"--grace-seconds",
),
) -> None:
"""Delete a job row and artifacts, optionally terminating first."""
async def action(client: ShellctlClient) -> DeleteJobResponse:
return await client.delete(
job_id,
force=force,
grace_seconds=grace_seconds,
)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_model,
)
@cli.command("serve")
def serve_command(
listen: str = "127.0.0.1:8765",
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty to disable HTTP bearer auth."
),
),
state_dir: Path | None = None,
runtime_dir: Path | None = None,
gc_interval_seconds: float = typer.Option(
DEFAULT_GC_INTERVAL_SECONDS,
"--gc-interval-seconds",
),
gc_finished_job_retention_seconds: float = typer.Option(
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
"--gc-finished-job-retention-seconds",
),
) -> None:
"""Run the local shellctl FastAPI server via uvicorn."""
from shellctl.server.serve import (
serve_command as server_serve_command,
)
server_serve_command(
listen=listen,
auth_token=auth_token,
state_dir=state_dir,
runtime_dir=runtime_dir,
gc_interval_seconds=gc_interval_seconds,
gc_finished_job_retention_seconds=gc_finished_job_retention_seconds,
)
def main() -> None:
"""CLI entrypoint used by the console script and `python -m` invocations."""
cli()
def _parse_env(values: list[str] | None) -> dict[str, str] | None:
if not values:
return None
parsed: dict[str, str] = {}
for value in values:
if "=" not in value:
raise typer.BadParameter(
"env entries must use NAME=VALUE format",
param_hint="--env",
)
name, env_value = value.split("=", 1)
if not name:
raise typer.BadParameter(
"env names must be non-empty",
param_hint="--env",
)
parsed[name] = env_value
return parsed
def _terminal_size(*, cols: int | None, rows: int | None) -> TerminalSize | None:
if cols is None and rows is None:
return None
return _build_model(
TerminalSize,
cols=cols if cols is not None else DEFAULT_TERMINAL_COLS,
rows=rows if rows is not None else DEFAULT_TERMINAL_ROWS,
)
def _build_model[ModelT: BaseModel](model_type: type[ModelT], /, **data: object) -> ModelT:
try:
return model_type(**data)
except ValidationError as exc:
raise typer.BadParameter(_validation_error_message(exc)) from exc
async def _with_client[ResponseT](
base_url: str,
auth_token: str | None,
output_limit: int,
idle_flush_seconds: float,
action: Callable[[ShellctlClient], Awaitable[ResponseT]],
) -> ResponseT:
async with ShellctlClient(
base_url,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
token=auth_token,
) as client:
return await action(client)
def _run_client_action[ResponseT](
*,
base_url: str,
auth_token: str | None,
action: Callable[[ShellctlClient], Awaitable[ResponseT]],
emit: Callable[[ResponseT], None],
output_limit: int = DEFAULT_OUTPUT_LIMIT_BYTES,
idle_flush_seconds: float = DEFAULT_IDLE_FLUSH_SECONDS,
) -> None:
try:
payload = anyio.run(
_with_client,
base_url,
auth_token,
output_limit,
idle_flush_seconds,
action,
)
except ShellctlClientError as exc:
_emit_error_and_exit(exc.code, exc.message)
except httpx.TimeoutException:
_emit_error_and_exit("request_timeout", "request timed out")
except httpx.TransportError as exc:
_emit_error_and_exit("connection_error", str(exc))
emit(payload)
def _emit_model(model: BaseModel) -> None:
typer.echo(model.model_dump_json(exclude_none=True), color=False)
def _emit_job_list(jobs: list[JobInfo]) -> None:
typer.echo(
json.dumps(
[item.model_dump(mode="json", exclude_none=True) for item in jobs],
separators=(",", ":"),
),
color=False,
)
def _emit_error_and_exit(code: str, message: str) -> NoReturn:
typer.echo(
json.dumps(
{"error": {"code": code, "message": message}},
separators=(",", ":"),
),
err=True,
color=False,
)
raise typer.Exit(code=1)
def _validation_error_message(exc: ValidationError) -> str:
detail = exc.errors(include_url=False)[0]
location = ".".join(str(part) for part in detail.get("loc", ()))
message = detail["msg"]
return f"{location}: {message}" if location else str(message)
__all__ = [
"cli",
"delete_command",
"health_command",
"input_command",
"list_command",
"main",
"run_command",
"serve_command",
"status_command",
"tail_command",
"terminate_command",
"wait_command",
]

Some files were not shown because too many files have changed in this diff Show More