From af4f19840c8a45ea700fceac33cd0e34da7cc323 Mon Sep 17 00:00:00 2001 From: zyssyz123 <916125788@qq.com> Date: Mon, 21 Sep 2026 07:11:12 +0000 Subject: [PATCH] feat(agent): record E2B sandbox runtime usage (#42554) --- api/configs/extra/agent_backend_config.py | 12 + api/controllers/inner_api/__init__.py | 2 + .../inner_api/agent/runtime_usage.py | 100 +++ api/extensions/ext_celery.py | 21 +- ...00-e7b2a9c4d601_add_agent_sandbox_usage.py | 123 +++ api/models/__init__.py | 3 + api/models/agent_sandbox_usage.py | 117 +++ api/schedule/collect_agent_sandbox_usage.py | 70 ++ api/services/agent/runtime_usage_service.py | 680 +++++++++++++++++ .../inner_api/test_runtime_usage.py | 219 ++++++ .../unit_tests/extensions/test_celery_ssl.py | 2 + .../test_collect_agent_sandbox_usage.py | 179 +++++ .../agent/test_runtime_usage_service.py | 720 ++++++++++++++++++ .../services/agent/test_workspace_service.py | 64 ++ dify-agent/docs/dify-agent/guide/index.md | 8 +- .../docs/dify-agent/guide/runtime-metering.md | 156 ++++ dify-agent/src/dify_agent/server/app.py | 6 +- .../src/dify_agent/server/e2b_usage_client.py | 112 +++ .../dify_agent/server/e2b_usage_collector.py | 160 ++++ .../src/dify_agent/server/routes/e2b_usage.py | 99 +++ dify-agent/src/dify_agent/server/settings.py | 7 + .../dify_agent/runtime_backend/test_e2b.py | 118 ++- .../tests/local/dify_agent/server/test_app.py | 32 + .../server/test_e2b_usage_client.py | 65 ++ .../server/test_e2b_usage_collector.py | 193 +++++ .../dify_agent/server/test_e2b_usage_route.py | 176 +++++ .../local/dify_agent/server/test_settings.py | 44 ++ .../envs/core-services/dify-agent.env.example | 11 + docker/envs/core-services/shared.env.example | 7 + 29 files changed, 3501 insertions(+), 5 deletions(-) create mode 100644 api/controllers/inner_api/agent/runtime_usage.py create mode 100644 api/migrations/versions/2026_09_20_1200-e7b2a9c4d601_add_agent_sandbox_usage.py create mode 100644 api/models/agent_sandbox_usage.py create mode 100644 api/schedule/collect_agent_sandbox_usage.py create mode 100644 api/services/agent/runtime_usage_service.py create mode 100644 api/tests/unit_tests/controllers/inner_api/test_runtime_usage.py create mode 100644 api/tests/unit_tests/schedule/test_collect_agent_sandbox_usage.py create mode 100644 api/tests/unit_tests/services/agent/test_runtime_usage_service.py create mode 100644 dify-agent/docs/dify-agent/guide/runtime-metering.md create mode 100644 dify-agent/src/dify_agent/server/e2b_usage_client.py create mode 100644 dify-agent/src/dify_agent/server/e2b_usage_collector.py create mode 100644 dify-agent/src/dify_agent/server/routes/e2b_usage.py create mode 100644 dify-agent/tests/local/dify_agent/server/test_e2b_usage_client.py create mode 100644 dify-agent/tests/local/dify_agent/server/test_e2b_usage_collector.py create mode 100644 dify-agent/tests/local/dify_agent/server/test_e2b_usage_route.py diff --git a/api/configs/extra/agent_backend_config.py b/api/configs/extra/agent_backend_config.py index 8b782b5e3e2..13a9f00eafc 100644 --- a/api/configs/extra/agent_backend_config.py +++ b/api/configs/extra/agent_backend_config.py @@ -7,6 +7,18 @@ class AgentBackendConfig(BaseSettings): Configuration settings for the Agent backend runtime integration. """ + AGENT_SANDBOX_METERING_ENABLED: bool = Field(default=False, description="Persist independent E2B runtime usage.") + AGENT_SANDBOX_METERING_PROJECT_ID: str = Field( + default="", description="Allowed E2B project/team for this database." + ) + AGENT_SANDBOX_METERING_START_AT: str = Field( + default="", description="Immutable UTC activation instant, aligned to a whole second; required when enabled." + ) + AGENT_SANDBOX_METERING_INTERVAL_SECONDS: int | str = Field( + default=60, + description="Celery Beat interval in seconds; validated only when registering optional usage collection.", + ) + AGENT_BACKEND_BASE_URL: str | None = Field( description="Base URL for the Dify Agent backend service.", default=None, diff --git a/api/controllers/inner_api/__init__.py b/api/controllers/inner_api/__init__.py index 1656df8aaff..addeb44e00c 100644 --- a/api/controllers/inner_api/__init__.py +++ b/api/controllers/inner_api/__init__.py @@ -19,6 +19,7 @@ from . import mail as _mail from . import runtime_credentials as _runtime_credentials from .agent import files as _agent_files from .agent import llm as _agent_llm +from .agent import runtime_usage as _agent_runtime_usage from .agent import tools as _agent_tools from .app import dsl as _app_dsl from .app import file_grants as _app_file_grants @@ -34,6 +35,7 @@ __all__ = [ "_agent_config", "_agent_files", "_agent_llm", + "_agent_runtime_usage", "_agent_tools", "_app_dsl", "_app_file_grants", diff --git a/api/controllers/inner_api/agent/runtime_usage.py b/api/controllers/inner_api/agent/runtime_usage.py new file mode 100644 index 00000000000..8389bac5d14 --- /dev/null +++ b/api/controllers/inner_api/agent/runtime_usage.py @@ -0,0 +1,100 @@ +"""Trusted ingestion for independently persisted E2B execution usage.""" + +from flask import request +from flask_restx import Resource +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.inner_api import inner_api_ns +from controllers.inner_api.wraps import agent_inner_api_only +from fields.base import ResponseModel +from libs.exception import BaseHTTPException +from libs.helper import dump_response +from services.agent.runtime_usage_service import SandboxUsageError, SandboxUsageEvent, SandboxUsageService + + +class SandboxUsageHttpError(BaseHTTPException): + error_code = "sandbox_usage_invalid_request" + description = "Invalid sandbox usage request." + code = 400 + + def __init__(self, error: SandboxUsageError | None = None) -> None: + if error is not None: + self.error_code = error.code + self.description = error.code + self.code = error.status_code + super().__init__(self.description) + + +class SandboxUsagePayload(BaseModel): + project_id: str = Field(min_length=1, max_length=128) + events: list[SandboxUsageEvent] = Field(min_length=1, max_length=100) + model_config = ConfigDict(extra="forbid") + + +class SandboxUsageQuery(BaseModel): + project_id: str = Field(min_length=1, max_length=128) + model_config = ConfigDict(extra="forbid") + + +class SandboxUsageResponse(ResponseModel): + accepted: int + duplicates: int + conflicts: int + ignored: int + + +class SandboxUsageDiagnosticsResponse(ResponseModel): + unresolved_events: int = 0 + conflict_events: int = 0 + open_executions: int = 0 + unattributed_executions: int = 0 + + +class SandboxUsageStateResponse(ResponseModel): + enabled: bool + project_id: str + started_at: str | None + checkpoint_at: str | None + full_scan_at: str | None + diagnostics: SandboxUsageDiagnosticsResponse + + +register_schema_models(inner_api_ns, SandboxUsagePayload, SandboxUsageQuery) +register_response_schema_models(inner_api_ns, SandboxUsageResponse, SandboxUsageStateResponse) + + +@inner_api_ns.route("/agent/sandbox-usage/events") +class SandboxUsageEventsApi(Resource): + @agent_inner_api_only + @inner_api_ns.doc("inner_agent_sandbox_usage_events") + @inner_api_ns.expect(inner_api_ns.models[SandboxUsagePayload.__name__]) + @inner_api_ns.response(200, "Durably committed", inner_api_ns.models[SandboxUsageResponse.__name__]) + def post(self): + request.max_content_length = 1024 * 1024 + try: + payload = SandboxUsagePayload.model_validate(inner_api_ns.payload or {}) + result = SandboxUsageService.ingest(project_id=payload.project_id, events=payload.events) + except ValidationError as exc: + raise SandboxUsageHttpError() from exc + except SandboxUsageError as exc: + raise SandboxUsageHttpError(exc) from exc + return dump_response(SandboxUsageResponse, result) + + +@inner_api_ns.route("/agent/sandbox-usage/state") +class SandboxUsageStateApi(Resource): + @agent_inner_api_only + @inner_api_ns.doc("inner_agent_sandbox_usage_state", params=query_params_from_model(SandboxUsageQuery)) + @inner_api_ns.response( + 200, "Persisted activation and completed scans", inner_api_ns.models[SandboxUsageStateResponse.__name__] + ) + def get(self): + try: + query = SandboxUsageQuery.model_validate(request.args.to_dict(flat=True)) + result = SandboxUsageService.get_state(project_id=query.project_id) + except ValidationError as exc: + raise SandboxUsageHttpError() from exc + except SandboxUsageError as exc: + raise SandboxUsageHttpError(exc) from exc + return dump_response(SandboxUsageStateResponse, result) diff --git a/api/extensions/ext_celery.py b/api/extensions/ext_celery.py index 724fb4774cc..f4eff8a1012 100644 --- a/api/extensions/ext_celery.py +++ b/api/extensions/ext_celery.py @@ -1,6 +1,7 @@ +import logging import ssl from datetime import timedelta -from typing import Any +from typing import Any, NotRequired import pytz # type: ignore[import-untyped] from celery import Celery, Task @@ -14,6 +15,8 @@ from enums import DeploymentEdition from extensions.redis_names import normalize_redis_key_prefix from extensions.workflow_warm_shutdown import setup_workflow_warm_shutdown_handler +logger = logging.getLogger(__name__) + class _CelerySentinelKwargsDict(TypedDict): socket_timeout: float | None @@ -36,6 +39,7 @@ class CelerySSLOptionsDict(TypedDict): class CeleryBeatScheduleEntry(TypedDict): task: str schedule: crontab | timedelta + options: NotRequired[dict[str, Any]] def _enqueue_initial_community_telemetry_heartbeat(sender: Any, **_: Any) -> None: @@ -166,6 +170,7 @@ def init_app(app: DifyApp) -> Celery: setup_workflow_warm_shutdown_handler() imports = [ + "schedule.collect_agent_sandbox_usage", # optional background provider accounting "tasks.async_workflow_tasks", # trigger workers "tasks.collect_agent_resources_task", # retired Agent resource collection "tasks.trigger_processing_tasks", # async trigger processing @@ -182,6 +187,20 @@ def init_app(app: DifyApp) -> Celery: # if you add a new task, please add the switch to CeleryScheduleTasksConfig beat_schedule: dict[str, CeleryBeatScheduleEntry] = {} + if dify_config.AGENT_SANDBOX_METERING_ENABLED: + try: + interval = int(dify_config.AGENT_SANDBOX_METERING_INTERVAL_SECONDS) + if interval < 1: + raise ValueError("collection interval must be positive") + collection_schedule = timedelta(seconds=interval) + except (ValueError, OverflowError): + logger.warning("Skipping sandbox usage schedule: interval must be a valid positive number of seconds") + else: + beat_schedule["collect_agent_sandbox_usage"] = { + "task": "schedule.collect_agent_sandbox_usage.collect_agent_sandbox_usage", + "schedule": collection_schedule, + "options": {"expires": interval}, + } if dify_config.ENABLE_CONVERSATION_CLEANUP_TASK: imports.append("tasks.delete_conversation_task") beat_schedule["conversation_cleanup_sweeper"] = { diff --git a/api/migrations/versions/2026_09_20_1200-e7b2a9c4d601_add_agent_sandbox_usage.py b/api/migrations/versions/2026_09_20_1200-e7b2a9c4d601_add_agent_sandbox_usage.py new file mode 100644 index 00000000000..e18c0b03942 --- /dev/null +++ b/api/migrations/versions/2026_09_20_1200-e7b2a9c4d601_add_agent_sandbox_usage.py @@ -0,0 +1,123 @@ +"""Add independent sandbox execution accounting. + +Revision ID: e7b2a9c4d601 +Revises: d8e4a6b1c902 +Create Date: 2026-09-20 12:00:00 +""" + +import sqlalchemy as sa +from alembic import op + +from models.types import StringUUID + +revision = "e7b2a9c4d601" +down_revision = "d8e4a6b1c902" +branch_labels = None +depends_on = None + + +def _identity_columns(): + # UTC timestamps, explicit application values; not runtime measurements. + return [ + sa.Column("id", StringUUID(), primary_key=True, nullable=False), + sa.Column("provider", sa.String(32), nullable=False), + sa.Column("provider_project_id", sa.String(128), nullable=False), + sa.Column("allocation_id", StringUUID()), + sa.Column("tenant_id", StringUUID()), + sa.Column("app_id", StringUUID()), + sa.Column("agent_id", StringUUID()), + sa.Column("binding_id", StringUUID()), + sa.Column("workspace_id", StringUUID()), + sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp()), + sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.current_timestamp()), + ] + + +def upgrade(): + op.create_table( + "agent_sandbox_usage_events", + *_identity_columns(), + sa.Column("source", sa.String(32), nullable=False), + sa.Column("source_event_id", sa.String(255), nullable=False), + sa.Column("event_type", sa.String(96), nullable=False), + sa.Column("sandbox_id", sa.String(128)), + sa.Column("provider_execution_id", sa.String(128)), + sa.Column("operation_id", StringUUID()), + sa.Column("lease_id", StringUUID()), + sa.Column("attempt", sa.Integer()), + sa.Column("purpose", sa.String(32)), + sa.Column("correlation", sa.JSON(), nullable=False), + sa.Column("occurred_at", sa.DateTime()), + sa.Column("received_at", sa.DateTime(), nullable=False), + sa.Column("payload", sa.JSON(), nullable=False), + sa.Column("canonical_hash", sa.String(64), nullable=False), + sa.Column("payload_versions", sa.JSON(), nullable=False), + sa.Column("resolution", sa.JSON(), nullable=False), + sa.Column("projection_status", sa.String(32), nullable=False), + sa.Column("projection_error_code", sa.String(96)), + sa.Column("processed_at", sa.DateTime()), + sa.UniqueConstraint( + "provider", "provider_project_id", "source", "source_event_id", name="sandbox_usage_event_unique" + ), + sa.CheckConstraint("source IN ('application', 'provider')", name="sandbox_usage_event_source_check"), + ) + for name, columns in ( + ("sandbox_usage_event_sandbox_idx", ["provider_project_id", "sandbox_id", "occurred_at"]), + ("sandbox_usage_event_allocation_idx", ["allocation_id", "received_at"]), + ("sandbox_usage_event_owner_idx", ["provider_project_id", "binding_id", "event_type"]), + ("sandbox_usage_event_pending_idx", ["projection_status", "received_at"]), + ("sandbox_usage_checkpoint_idx", ["provider_project_id", "event_type", "purpose", "occurred_at"]), + ("sandbox_usage_diagnostics_idx", ["provider_project_id", "projection_status"]), + ("sandbox_usage_execution_idx", ["provider_project_id", "provider_execution_id", "projection_error_code"]), + ): + op.create_index(name, "agent_sandbox_usage_events", columns) + + op.create_table( + "agent_sandbox_executions", + *_identity_columns(), + sa.Column("sandbox_id", sa.String(128), nullable=False), + sa.Column("provider_execution_id", sa.String(128), nullable=False), + sa.Column("attribution_status", sa.String(32), nullable=False), + sa.Column("template_id", sa.String(128)), + sa.Column("template_build_id", sa.String(128)), + sa.Column("vcpu_count", sa.Integer()), + sa.Column("memory_mib", sa.BigInteger()), + sa.Column("started_at", sa.DateTime()), + sa.Column("started_at_source", sa.String(32)), + sa.Column("started_at_precision_ms", sa.Integer()), + sa.Column("terminal_event_at", sa.DateTime()), + sa.Column("metered_duration_ms", sa.BigInteger()), + sa.Column("measurement_source", sa.String(32), nullable=False), + sa.Column("state", sa.String(16), nullable=False), + sa.Column("quality", sa.String(16), nullable=False), + sa.Column("close_reason", sa.String(64)), + sa.Column("terminal_event_id", StringUUID()), + sa.Column("last_reconciled_at", sa.DateTime()), + sa.Column("revision", sa.Integer(), nullable=False), + sa.UniqueConstraint( + "provider", "provider_project_id", "provider_execution_id", name="sandbox_execution_unique" + ), + sa.CheckConstraint("metered_duration_ms IS NULL OR metered_duration_ms >= 0", name="sandbox_duration_check"), + sa.CheckConstraint("vcpu_count IS NULL OR vcpu_count > 0", name="sandbox_vcpu_check"), + sa.CheckConstraint("memory_mib IS NULL OR memory_mib > 0", name="sandbox_memory_check"), + sa.CheckConstraint( + "quality <> 'metered' OR (state = 'closed' AND measurement_source = 'provider_execution' " + "AND metered_duration_ms IS NOT NULL AND started_at IS NOT NULL " + "AND vcpu_count IS NOT NULL AND memory_mib IS NOT NULL)", + name="sandbox_metered_check", + ), + ) + for name, columns in ( + ("sandbox_execution_project_time_idx", ["provider_project_id", "started_at"]), + ("sandbox_execution_tenant_time_idx", ["tenant_id", "started_at"]), + ("sandbox_execution_sandbox_idx", ["provider_project_id", "sandbox_id", "started_at"]), + ("sandbox_execution_allocation_idx", ["allocation_id"]), + ("sandbox_execution_reconcile_idx", ["quality", "state", "last_reconciled_at"]), + ("sandbox_execution_diagnostics_idx", ["provider_project_id", "state", "attribution_status"]), + ): + op.create_index(name, "agent_sandbox_executions", columns) + + +def downgrade(): + op.drop_table("agent_sandbox_executions") + op.drop_table("agent_sandbox_usage_events") diff --git a/api/models/__init__.py b/api/models/__init__.py index ee32c8a8ee8..eedd7f4d598 100644 --- a/api/models/__init__.py +++ b/api/models/__init__.py @@ -30,6 +30,7 @@ from .agent import ( WorkflowAgentBindingType, WorkflowAgentNodeBinding, ) +from .agent_sandbox_usage import AgentSandboxExecution, AgentSandboxUsageEvent from .api_based_extension import APIBasedExtension, APIBasedExtensionPoint from .comment import ( WorkflowComment, @@ -171,6 +172,8 @@ __all__ = [ "AgentHomeSnapshot", "AgentIconType", "AgentKind", + "AgentSandboxExecution", + "AgentSandboxUsageEvent", "AgentScope", "AgentSkillBinding", "AgentSource", diff --git a/api/models/agent_sandbox_usage.py b/api/models/agent_sandbox_usage.py new file mode 100644 index 00000000000..2539895a6b6 --- /dev/null +++ b/api/models/agent_sandbox_usage.py @@ -0,0 +1,117 @@ +"""Independent E2B accounting records, retained after business resource deletion. + +All timestamps are explicit naive UTC values, following the API's database +convention. These records deliberately have no business-object foreign keys. +""" + +from datetime import datetime +from typing import Any + +import sqlalchemy as sa +from sqlalchemy.orm import Mapped, mapped_column + +from libs.datetime_utils import naive_utc_now + +from .base import Base, DefaultFieldsMixin +from .types import StringUUID + + +class AgentSandboxUsageEvent(DefaultFieldsMixin, Base): + __tablename__ = "agent_sandbox_usage_events" + __table_args__ = ( + sa.UniqueConstraint( + "provider", "provider_project_id", "source", "source_event_id", name="sandbox_usage_event_unique" + ), + sa.Index("sandbox_usage_event_sandbox_idx", "provider_project_id", "sandbox_id", "occurred_at"), + sa.Index("sandbox_usage_event_allocation_idx", "allocation_id", "received_at"), + sa.Index("sandbox_usage_event_owner_idx", "provider_project_id", "binding_id", "event_type"), + sa.Index("sandbox_usage_event_pending_idx", "projection_status", "received_at"), + sa.Index("sandbox_usage_checkpoint_idx", "provider_project_id", "event_type", "purpose", "occurred_at"), + sa.Index("sandbox_usage_diagnostics_idx", "provider_project_id", "projection_status"), + sa.Index( + "sandbox_usage_execution_idx", "provider_project_id", "provider_execution_id", "projection_error_code" + ), + sa.CheckConstraint("source IN ('application', 'provider')", name="sandbox_usage_event_source_check"), + ) + + provider: Mapped[str] = mapped_column(sa.String(32), nullable=False, default="e2b") + provider_project_id: Mapped[str] = mapped_column(sa.String(128), nullable=False) + source: Mapped[str] = mapped_column(sa.String(32), nullable=False) + source_event_id: Mapped[str] = mapped_column(sa.String(255), nullable=False) + event_type: Mapped[str] = mapped_column(sa.String(96), nullable=False) + sandbox_id: Mapped[str | None] = mapped_column(sa.String(128)) + provider_execution_id: Mapped[str | None] = mapped_column(sa.String(128)) + allocation_id: Mapped[str | None] = mapped_column(StringUUID) + operation_id: Mapped[str | None] = mapped_column(StringUUID) + lease_id: Mapped[str | None] = mapped_column(StringUUID) + attempt: Mapped[int | None] = mapped_column(sa.Integer) + tenant_id: Mapped[str | None] = mapped_column(StringUUID) + app_id: Mapped[str | None] = mapped_column(StringUUID) + agent_id: Mapped[str | None] = mapped_column(StringUUID) + binding_id: Mapped[str | None] = mapped_column(StringUUID) + workspace_id: Mapped[str | None] = mapped_column(StringUUID) + purpose: Mapped[str | None] = mapped_column(sa.String(32)) + correlation: Mapped[dict[str, Any]] = mapped_column(sa.JSON, nullable=False, default=dict) + occurred_at: Mapped[datetime | None] = mapped_column(sa.DateTime) + received_at: Mapped[datetime] = mapped_column(sa.DateTime, nullable=False, default=naive_utc_now) + payload: Mapped[dict[str, Any]] = mapped_column(sa.JSON, nullable=False) + canonical_hash: Mapped[str] = mapped_column(sa.String(64), nullable=False) + payload_versions: Mapped[list[dict[str, Any]]] = mapped_column(sa.JSON, nullable=False, default=list) + resolution: Mapped[dict[str, Any]] = mapped_column(sa.JSON, nullable=False, default=dict) + projection_status: Mapped[str] = mapped_column(sa.String(32), nullable=False, default="pending") + projection_error_code: Mapped[str | None] = mapped_column(sa.String(96)) + processed_at: Mapped[datetime | None] = mapped_column(sa.DateTime) + + +class AgentSandboxExecution(DefaultFieldsMixin, Base): + """One provider execution, not one request, lease, or connect call.""" + + __tablename__ = "agent_sandbox_executions" + __table_args__ = ( + sa.UniqueConstraint( + "provider", "provider_project_id", "provider_execution_id", name="sandbox_execution_unique" + ), + sa.Index("sandbox_execution_project_time_idx", "provider_project_id", "started_at"), + sa.Index("sandbox_execution_tenant_time_idx", "tenant_id", "started_at"), + sa.Index("sandbox_execution_sandbox_idx", "provider_project_id", "sandbox_id", "started_at"), + sa.Index("sandbox_execution_allocation_idx", "allocation_id"), + sa.Index("sandbox_execution_reconcile_idx", "quality", "state", "last_reconciled_at"), + sa.Index("sandbox_execution_diagnostics_idx", "provider_project_id", "state", "attribution_status"), + sa.CheckConstraint("metered_duration_ms IS NULL OR metered_duration_ms >= 0", name="sandbox_duration_check"), + sa.CheckConstraint("vcpu_count IS NULL OR vcpu_count > 0", name="sandbox_vcpu_check"), + sa.CheckConstraint("memory_mib IS NULL OR memory_mib > 0", name="sandbox_memory_check"), + sa.CheckConstraint( + "quality <> 'metered' OR (state = 'closed' AND measurement_source = 'provider_execution' " + "AND metered_duration_ms IS NOT NULL AND started_at IS NOT NULL " + "AND vcpu_count IS NOT NULL AND memory_mib IS NOT NULL)", + name="sandbox_metered_check", + ), + ) + + provider: Mapped[str] = mapped_column(sa.String(32), nullable=False, default="e2b") + provider_project_id: Mapped[str] = mapped_column(sa.String(128), nullable=False) + sandbox_id: Mapped[str] = mapped_column(sa.String(128), nullable=False) + provider_execution_id: Mapped[str] = mapped_column(sa.String(128), nullable=False) + allocation_id: Mapped[str | None] = mapped_column(StringUUID) + tenant_id: Mapped[str | None] = mapped_column(StringUUID) + app_id: Mapped[str | None] = mapped_column(StringUUID) + agent_id: Mapped[str | None] = mapped_column(StringUUID) + binding_id: Mapped[str | None] = mapped_column(StringUUID) + workspace_id: Mapped[str | None] = mapped_column(StringUUID) + attribution_status: Mapped[str] = mapped_column(sa.String(32), nullable=False, default="unresolved") + template_id: Mapped[str | None] = mapped_column(sa.String(128)) + template_build_id: Mapped[str | None] = mapped_column(sa.String(128)) + vcpu_count: Mapped[int | None] = mapped_column(sa.Integer) + memory_mib: Mapped[int | None] = mapped_column(sa.BigInteger) + started_at: Mapped[datetime | None] = mapped_column(sa.DateTime) + started_at_source: Mapped[str | None] = mapped_column(sa.String(32)) + started_at_precision_ms: Mapped[int | None] = mapped_column(sa.Integer) + terminal_event_at: Mapped[datetime | None] = mapped_column(sa.DateTime) + metered_duration_ms: Mapped[int | None] = mapped_column(sa.BigInteger) + measurement_source: Mapped[str] = mapped_column(sa.String(32), nullable=False, default="unknown") + state: Mapped[str] = mapped_column(sa.String(16), nullable=False, default="unknown") + quality: Mapped[str] = mapped_column(sa.String(16), nullable=False, default="pending") + close_reason: Mapped[str | None] = mapped_column(sa.String(64)) + terminal_event_id: Mapped[str | None] = mapped_column(StringUUID) + last_reconciled_at: Mapped[datetime | None] = mapped_column(sa.DateTime) + revision: Mapped[int] = mapped_column(sa.Integer, nullable=False, default=1) diff --git a/api/schedule/collect_agent_sandbox_usage.py b/api/schedule/collect_agent_sandbox_usage.py new file mode 100644 index 00000000000..27a7223ead3 --- /dev/null +++ b/api/schedule/collect_agent_sandbox_usage.py @@ -0,0 +1,70 @@ +"""Use existing Celery infrastructure to request one bounded E2B usage scan.""" + +import logging +from urllib.parse import urlsplit + +from celery import shared_task +from pydantic import BaseModel, ConfigDict, StrictBool + +from configs import dify_config +from core.helper import ssrf_proxy + +logger = logging.getLogger(__name__) + + +class _CollectionResponse(BaseModel): + completed: StrictBool + model_config = ConfigDict(extra="ignore") + + +@shared_task( + name="schedule.collect_agent_sandbox_usage.collect_agent_sandbox_usage", + queue="ops_trace", + soft_time_limit=270, + time_limit=300, + ignore_result=True, +) +def collect_agent_sandbox_usage() -> bool: + """Accounting failures fail this job only; the next scheduled scan retries. + + No provider credential is copied to the API. The Agent's authenticated + control plane owns its E2B client and executes one request-scoped scan. + """ + if not dify_config.AGENT_SANDBOX_METERING_ENABLED: + return False + try: + base_url = (dify_config.AGENT_BACKEND_BASE_URL or "").rstrip("/") + token = (dify_config.AGENT_BACKEND_API_TOKEN or "").strip() + project_id = dify_config.AGENT_SANDBOX_METERING_PROJECT_ID.strip() + parsed = urlsplit(base_url) + if ( + parsed.scheme not in {"http", "https"} + or not parsed.hostname + or parsed.username is not None + or parsed.password is not None + or parsed.query + or parsed.fragment + or not token + or not project_id + ): + raise ValueError("Sandbox usage collection requires backend URL, token, and project configuration") + response = ssrf_proxy.post( + f"{base_url}/internal/e2b/usage/collect", + headers={"Authorization": f"Bearer {token}"}, + json={"project_id": project_id}, + max_retries=0, + timeout=260, + ) + try: + response.raise_for_status() + completed = _CollectionResponse.model_validate(response.json()).completed + finally: + response.close() + if not completed: + raise RuntimeError("Sandbox usage collection did not complete its bounded scan") + except Exception as exc: + # Do not include credentials, response bodies, or provider metadata. + logger.warning("Background sandbox usage collection failed", extra={"error_type": type(exc).__name__}) + raise + logger.info("Background sandbox usage collection completed") + return True diff --git a/api/services/agent/runtime_usage_service.py b/api/services/agent/runtime_usage_service.py new file mode 100644 index 00000000000..43967bc1b98 --- /dev/null +++ b/api/services/agent/runtime_usage_service.py @@ -0,0 +1,680 @@ +"""Forward-only, provider-authoritative sandbox accounting. + +The activation event is also a project-scoped transaction lock. This keeps +deduplication, checkpoints, and execution projection atomic +across API workers without coupling billing records to business transactions. +Business creation/lookups never call accounting. Ownership is resolved only +while collecting provider events. No network calls belong in this service. +""" + +import hashlib +import json +from collections.abc import Callable, Mapping +from datetime import UTC, datetime +from typing import Any, Literal, TypeVar +from uuid import UUID + +import sqlalchemy as sa +from pydantic import BaseModel, ConfigDict, Field, JsonValue +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +from configs import dify_config +from core.db.session_factory import session_factory +from libs.datetime_utils import naive_utc_now +from models.agent import Agent, AgentWorkspace, AgentWorkspaceBinding +from models.agent_sandbox_usage import AgentSandboxExecution, AgentSandboxUsageEvent +from models.model import App + +_T = TypeVar("_T") +_COLLECTOR_EVENT_TYPES = {"collector_checkpoint", "collector_retention_gap"} +_TERMINAL_TYPES = {"sandbox.lifecycle.paused", "sandbox.lifecycle.killed"} +_NONTERMINAL_TYPES = {"sandbox.lifecycle.created", "sandbox.lifecycle.resumed", "sandbox.lifecycle.updated"} + + +class SandboxUsageError(ValueError): + def __init__(self, code: str, *, status_code: int = 400) -> None: + self.code = code + self.status_code = status_code + super().__init__(code) + + +class SandboxUsageEvent(BaseModel): + id: str = Field(min_length=1, max_length=255) + source: Literal["application", "provider"] + type: str = Field(min_length=1, max_length=96) + timestamp: str | None = None + sandbox_id: str | None = Field(default=None, max_length=128) + execution_id: str | None = Field(default=None, max_length=128) + payload: dict[str, JsonValue] + model_config = ConfigDict(extra="forbid") + + +def _utc(value: object) -> datetime | None: + if value is None: + return None + if not isinstance(value, str): + raise SandboxUsageError("invalid_timestamp") + try: + result = datetime.fromisoformat(value) + except ValueError as exc: + raise SandboxUsageError("invalid_timestamp") from exc + if result.tzinfo is None: + raise SandboxUsageError("timestamp_requires_timezone") + return result.astimezone(UTC).replace(tzinfo=None) + + +def _iso(value: datetime | None) -> str | None: + return value.replace(tzinfo=UTC).isoformat().replace("+00:00", "Z") if value else None + + +def _hash(payload: Mapping[str, Any]) -> str: + return hashlib.sha256(json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()).hexdigest() + + +def _mapping(value: object) -> dict[str, Any]: + return dict(value) if isinstance(value, dict) else {} + + +def _pick(data: Mapping[str, Any], *names: str) -> Any: + for name in names: + if name in data and data[name] is not None: + return data[name] + return None + + +def _positive_int(value: object, *, allow_zero: bool = False) -> int | None: + if value is None: + return None + if type(value) is not int or value < (0 if allow_zero else 1) or value > 2**63 - 1: + raise SandboxUsageError("invalid_provider_quantity") + return value + + +def _provider_values(event: SandboxUsageEvent, project_id: str) -> dict[str, Any]: + """Normalize the current v2 API and webhook spellings, not old event versions.""" + raw = event.payload + data = _mapping(_pick(raw, "eventData", "event_data", "data")) + execution = _mapping(data.get("execution")) + team = _pick(raw, "sandboxTeamId", "sandbox_team_id") + if team != project_id: + raise SandboxUsageError("provider_project_mismatch", status_code=403) + raw_id = raw.get("id") + raw_type = raw.get("type") + sandbox_id = _pick(raw, "sandboxId", "sandbox_id") + execution_id = _pick(raw, "sandboxExecutionId", "sandbox_execution_id") + if raw_id != event.id or raw_type != event.type: + raise SandboxUsageError("provider_envelope_mismatch") + for declared, actual in ((event.sandbox_id, sandbox_id), (event.execution_id, execution_id)): + if declared is not None and declared != actual: + raise SandboxUsageError("provider_envelope_mismatch") + for identifier in (sandbox_id, execution_id): + if identifier is not None and (not isinstance(identifier, str) or not identifier or len(identifier) > 128): + raise SandboxUsageError("invalid_provider_identifier") + return { + "id": event.id, + "version": raw.get("version"), + "type": event.type, + "timestamp": _iso(_utc(raw.get("timestamp"))), + "sandbox_id": sandbox_id, + "execution_id": execution_id, + "project_id": project_id, + "template_id": _pick(raw, "sandboxTemplateId", "sandbox_template_id"), + "template_build_id": _pick(raw, "sandboxBuildId", "sandbox_build_id"), + "started_at": _iso(_utc(_pick(execution, "started_at", "startedAt"))), + "duration_ms": _positive_int(_pick(execution, "execution_time", "executionTime"), allow_zero=True), + "vcpu_count": _positive_int(_pick(execution, "vcpu_count", "vcpuCount")), + "memory_mib": _positive_int(_pick(execution, "memory_mb", "memoryMb")), + "metadata": _mapping(_pick(data, "sandbox_metadata", "sandboxMetadata")), + } + + +def _enrich(existing: dict[str, Any], incoming: dict[str, Any]) -> dict[str, Any] | None: + """Accept absent-field enrichment; never silently choose contradictory facts.""" + merged = dict(existing) + for key, value in incoming.items(): + old = merged.get(key) + if isinstance(old, dict) and isinstance(value, dict): + value = _enrich(old, value) + if value is None: + return None + elif old is not None and value is not None and old != value: + return None + if value is not None: + merged[key] = value + return merged + + +class SandboxUsageService: + @staticmethod + def _config(project_id: str | None = None) -> tuple[str, datetime]: + if not dify_config.AGENT_SANDBOX_METERING_ENABLED: + raise SandboxUsageError("sandbox_metering_disabled", status_code=503) + configured = dify_config.AGENT_SANDBOX_METERING_PROJECT_ID.strip() + if not configured or len(configured) > 128: + raise SandboxUsageError("sandbox_metering_project_not_configured", status_code=503) + if project_id is not None and project_id != configured: + raise SandboxUsageError("sandbox_metering_project_not_allowed", status_code=403) + try: + start = _utc(dify_config.AGENT_SANDBOX_METERING_START_AT) + except SandboxUsageError as exc: + raise SandboxUsageError("sandbox_metering_start_not_configured", status_code=503) from exc + if start is None or start.microsecond: + raise SandboxUsageError("sandbox_metering_start_requires_whole_second", status_code=503) + return configured, start + + @classmethod + def _write(cls, project_id: str, action: Callable[[Session, datetime], _T]) -> _T: + _, start = cls._config(project_id) + # The first activation can race. Retrying the whole short transaction + # also handles concurrent unique-key inserts without losing an event. + for attempt in range(3): + try: + with session_factory.create_session() as session, session.begin(): + activation = cls._event(session, project_id, "application", "metering_activated", lock=True) + if activation is None: + activation = AgentSandboxUsageEvent( + provider="e2b", + provider_project_id=project_id, + source="application", + source_event_id="metering_activated", + event_type="metering_activated", + occurred_at=start, + payload={"started_at": _iso(start)}, + canonical_hash=_hash({"started_at": _iso(start)}), + projection_status="applied", + ) + session.add(activation) + session.flush() + elif _utc(activation.payload.get("started_at")) != start: + raise SandboxUsageError("sandbox_metering_start_is_immutable", status_code=409) + result = action(session, start) + return result + except IntegrityError: + if attempt == 2: + raise + raise AssertionError("unreachable") + + @staticmethod + def _event( + session: Session, project_id: str, source: str, event_id: str, *, lock: bool = False + ) -> AgentSandboxUsageEvent | None: + query = sa.select(AgentSandboxUsageEvent).where( + AgentSandboxUsageEvent.provider == "e2b", + AgentSandboxUsageEvent.provider_project_id == project_id, + AgentSandboxUsageEvent.source == source, + AgentSandboxUsageEvent.source_event_id == event_id, + ) + return session.scalar(query.with_for_update() if lock else query) + + @classmethod + def get_state(cls, *, project_id: str) -> dict[str, Any]: + if not dify_config.AGENT_SANDBOX_METERING_ENABLED: + return { + "enabled": False, + "project_id": project_id, + "started_at": None, + "checkpoint_at": None, + "full_scan_at": None, + "diagnostics": { + "unresolved_events": 0, + "conflict_events": 0, + "open_executions": 0, + "unattributed_executions": 0, + }, + } + + def read(session: Session, start: datetime) -> dict[str, Any]: + checkpoints = session.execute( + sa.select(AgentSandboxUsageEvent.purpose, sa.func.max(AgentSandboxUsageEvent.occurred_at)) + .where( + AgentSandboxUsageEvent.provider_project_id == project_id, + AgentSandboxUsageEvent.source == "application", + AgentSandboxUsageEvent.event_type == "collector_checkpoint", + AgentSandboxUsageEvent.projection_status == "applied", + ) + .group_by(AgentSandboxUsageEvent.purpose) + ).all() + ends = [end for _purpose, end in checkpoints if end is not None] + full = [end for purpose, end in checkpoints if purpose == "full" and end is not None] + event_scope = sa.and_( + AgentSandboxUsageEvent.provider == "e2b", + AgentSandboxUsageEvent.provider_project_id == project_id, + sa.or_( + AgentSandboxUsageEvent.source == "provider", + AgentSandboxUsageEvent.event_type.in_(_COLLECTOR_EVENT_TYPES), + ), + ) + execution_scope = AgentSandboxExecution.provider_project_id == project_id + diagnostics = { + "unresolved_events": session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxUsageEvent) + .where(event_scope, AgentSandboxUsageEvent.projection_status.in_(["pending", "unresolved"])) + ) + or 0, + "conflict_events": session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxUsageEvent) + .where(event_scope, AgentSandboxUsageEvent.projection_status == "conflict") + ) + or 0, + "open_executions": session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxExecution) + .where(execution_scope, AgentSandboxExecution.state == "open") + ) + or 0, + "unattributed_executions": session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxExecution) + .where(execution_scope, AgentSandboxExecution.attribution_status != "resolved") + ) + or 0, + } + return { + "enabled": True, + "project_id": project_id, + "started_at": _iso(start), + "checkpoint_at": _iso(max(ends, default=None)), + "full_scan_at": _iso(max(full, default=None)), + "diagnostics": diagnostics, + } + + _, configured_start = cls._config(project_id) + with session_factory.create_session() as session: + activation = cls._event(session, project_id, "application", "metering_activated") + if activation is not None: + if _utc(activation.payload.get("started_at")) != configured_start: + raise SandboxUsageError("sandbox_metering_start_is_immutable", status_code=409) + return read(session, configured_start) + # Only the first lookup initializes T0. Ordinary polling is SELECT-only, + # does not lock the project row, and cannot delay an ingestion transaction. + return cls._write(project_id, read) + + @classmethod + def ingest(cls, *, project_id: str, events: list[SandboxUsageEvent]) -> dict[str, int]: + if len(events) > 100: + raise SandboxUsageError("too_many_usage_events") + + def ingest_batch(session: Session, start: datetime) -> dict[str, int]: + counts = {"accepted": 0, "duplicates": 0, "conflicts": 0, "ignored": 0} + for event in events: + counts[cls._ingest_one(session, project_id, start, event)] += 1 + return counts + + return cls._write(project_id, ingest_batch) + + @classmethod + def _ingest_one(cls, session: Session, project_id: str, start: datetime, event: SandboxUsageEvent) -> str: + if event.source == "provider": + canonical = _provider_values(event, project_id) + begun = _utc(canonical["started_at"]) + occurred = _utc(canonical["timestamp"]) + if (begun is not None and begun < start) or (begun is None and occurred is not None and occurred < start): + previous = cls._event(session, project_id, "provider", event.id) + physical = ( + session.scalar( + sa.select(AgentSandboxExecution).where( + AgentSandboxExecution.provider_project_id == project_id, + AgentSandboxExecution.provider_execution_id == canonical["execution_id"], + ) + ) + if canonical["execution_id"] + else None + ) + # New old executions are excluded; corrections to known facts + # must still go through conflict detection and preserve evidence. + if physical is None and previous is None: + return "ignored" + if physical is None or physical.started_at is None: + if physical is not None: + session.delete(physical) + for pending in session.scalars( + sa.select(AgentSandboxUsageEvent).where( + AgentSandboxUsageEvent.provider_project_id == project_id, + AgentSandboxUsageEvent.source == "provider", + AgentSandboxUsageEvent.provider_execution_id == canonical["execution_id"], + ) + ): + pending.projection_status = "applied" + pending.projection_error_code = "execution_predates_activation" + if previous is not None: + digest = _hash(canonical) + if not any(version.get("hash") == digest for version in previous.payload_versions): + previous.payload_versions = [ + *previous.payload_versions, + { + "hash": digest, + "payload": event.payload, + "received_at": _iso(naive_utc_now()), + }, + ] + previous.resolution = {**previous.resolution, "scope": "excluded_before_activation"} + return "ignored" + else: + if event.type not in _COLLECTOR_EVENT_TYPES: + raise SandboxUsageError("unsupported_application_event") + occurred = _utc(event.timestamp) + if occurred is not None and occurred < start: + return "ignored" + canonical = event.model_dump(mode="json") + if event.type == "collector_checkpoint": + cls._validate_checkpoint(event.payload, start) + elif event.type == "collector_retention_gap": + lower = _utc(event.payload.get("uncovered_start")) + upper = _utc(event.payload.get("uncovered_end")) + if ( + lower is None + or upper is None + or occurred is None + or not start <= lower < upper <= occurred + or _positive_int(event.payload.get("assumed_retention_seconds")) is None + ): + raise SandboxUsageError("invalid_collector_retention_gap") + digest = _hash(canonical) + row = cls._event(session, project_id, event.source, event.id) + if row is not None: + if row.canonical_hash == digest: + if event.source == "provider" and row.projection_status != "conflict": + cls._project_provider(session, row, canonical, start) + return "duplicates" + versions = list(row.payload_versions or []) + if not any(version.get("hash") == digest for version in versions): + versions.append({"hash": digest, "payload": event.payload, "received_at": _iso(naive_utc_now())}) + row.payload_versions = versions + old = _mapping(row.resolution.get("canonical")) + # Attribution labels are optional. A bad/changed label must not + # invalidate otherwise identical provider-measured runtime. + merged = ( + _enrich( + {key: value for key, value in old.items() if key != "metadata"}, + {key: value for key, value in canonical.items() if key != "metadata"}, + ) + if old and event.source == "provider" + else None + ) + if merged is not None: + merged["metadata"] = canonical.get("metadata") or old.get("metadata", {}) + if merged is None or row.projection_status == "conflict": + cls._conflict(session, row, "event_content_conflict") + return "conflicts" + if merged == old: + return "duplicates" + canonical = merged + row.resolution = {"canonical": canonical} + row.canonical_hash = _hash(canonical) + else: + row = AgentSandboxUsageEvent( + provider="e2b", + provider_project_id=project_id, + source=event.source, + source_event_id=event.id, + event_type=event.type, + sandbox_id=event.sandbox_id, + provider_execution_id=event.execution_id, + occurred_at=occurred, + payload=event.payload, + canonical_hash=digest, + resolution={"canonical": canonical}, + ) + session.add(row) + session.flush() + if event.source == "provider": + cls._project_provider(session, row, canonical, start) + else: + row.projection_status = "applied" + if event.type == "collector_checkpoint": + row.occurred_at = _utc(event.payload["window_end"]) + row.purpose = str(event.payload["mode"]) + elif event.type == "collector_retention_gap": + row.projection_status = "unresolved" + row.projection_error_code = "provider_retention_gap" + row.processed_at = naive_utc_now() + session.flush() + return "conflicts" if row.projection_status == "conflict" else "accepted" + + @staticmethod + def _validate_checkpoint(payload: Mapping[str, Any], start: datetime) -> None: + lower = _utc(payload.get("window_start")) + upper = _utc(payload.get("window_end")) + scan = _utc(payload.get("scan_started_at")) + if ( + payload.get("completed") is not True + or payload.get("mode") not in {"incremental", "full"} + or lower is None + or upper is None + or scan is None + or not start <= lower <= upper <= scan + or scan > naive_utc_now() + or type(payload.get("pages")) is not int + or payload["pages"] < 1 + or type(payload.get("events")) is not int + or payload["events"] < 0 + ): + raise SandboxUsageError("invalid_collector_checkpoint") + + @staticmethod + def _verified_owner(session: Session, sandbox_id: str, metadata: dict[str, Any]) -> AgentWorkspaceBinding | None: + """Resolve optional labels through the existing tenant-owned resource chain. + + Metadata alone is never ownership evidence. Retired rows are still valid + historical associations until physical collection removes them. + """ + binding_id = metadata.get("dify.binding_id") + tenant_id = metadata.get("dify.tenant_id") + if not isinstance(binding_id, str) or not isinstance(tenant_id, str): + return None + try: + binding_id, tenant_id = str(UUID(binding_id)), str(UUID(tenant_id)) + except ValueError: + return None + binding = session.scalar( + sa.select(AgentWorkspaceBinding) + .join( + AgentWorkspace, + sa.and_( + AgentWorkspace.id == AgentWorkspaceBinding.workspace_id, + AgentWorkspace.tenant_id == AgentWorkspaceBinding.tenant_id, + AgentWorkspace.app_id == AgentWorkspaceBinding.app_id, + ), + ) + .join( + App, + sa.and_( + App.id == AgentWorkspaceBinding.app_id, + App.tenant_id == AgentWorkspaceBinding.tenant_id, + ), + ) + .join( + Agent, + sa.and_( + Agent.id == AgentWorkspaceBinding.agent_id, + Agent.tenant_id == AgentWorkspaceBinding.tenant_id, + ), + ) + .where( + AgentWorkspaceBinding.id == binding_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.backend_binding_ref == sandbox_id, + AgentWorkspace.backend_workspace_ref == sandbox_id, + ) + ) + if binding is None: + return None + for key, expected in ( + ("dify.agent_id", binding.agent_id), + ("dify.workspace_id", binding.workspace_id), + ): + supplied = metadata.get(key) + if supplied is not None and supplied != expected: + return None + return binding + + @classmethod + def _attribute_execution(cls, session: Session, execution: AgentSandboxExecution, metadata: dict[str, Any]) -> None: + owner = cls._verified_owner(session, execution.sandbox_id, metadata) + if owner is None or execution.attribution_status == "conflict": + # Existing snapshots survive missing/deleted business resources or + # malformed labels. No fallback to unverified allocation metadata. + return + expected = (owner.tenant_id, owner.app_id, owner.agent_id, owner.id, owner.workspace_id) + retained = ( + execution.tenant_id, + execution.app_id, + execution.agent_id, + execution.binding_id, + execution.workspace_id, + ) + if execution.binding_id is not None: + if retained != expected: + execution.attribution_status = "conflict" + return + other_owner = session.scalar( + sa.select(AgentSandboxExecution.id) + .where( + AgentSandboxExecution.provider == execution.provider, + AgentSandboxExecution.provider_project_id == execution.provider_project_id, + AgentSandboxExecution.sandbox_id == execution.sandbox_id, + AgentSandboxExecution.attribution_status == "resolved", + AgentSandboxExecution.binding_id != owner.id, + ) + .limit(1) + ) + if other_owner is not None: + execution.attribution_status = "conflict" + return + execution.tenant_id, execution.app_id, execution.agent_id = owner.tenant_id, owner.app_id, owner.agent_id + execution.binding_id, execution.workspace_id = owner.id, owner.workspace_id + execution.attribution_status = "resolved" + + @classmethod + def _project_provider( + cls, session: Session, row: AgentSandboxUsageEvent, data: dict[str, Any], start: datetime + ) -> None: + row.sandbox_id = data["sandbox_id"] + row.provider_execution_id = data["execution_id"] + started_at = _utc(data["started_at"]) + if data["version"] != "v2" or not row.sandbox_id or not row.provider_execution_id: + row.projection_status = "unresolved" + row.projection_error_code = "incomplete_provider_execution" + return + if started_at is not None and started_at < start: + cls._conflict(session, row, "execution_start_scope_conflict") + return + excluded = session.scalar( + sa.select(AgentSandboxUsageEvent.id) + .where( + AgentSandboxUsageEvent.provider_project_id == row.provider_project_id, + AgentSandboxUsageEvent.source == "provider", + AgentSandboxUsageEvent.provider_execution_id == row.provider_execution_id, + AgentSandboxUsageEvent.projection_error_code == "execution_predates_activation", + ) + .limit(1) + ) + if excluded is not None: + if started_at is not None: + cls._conflict(session, row, "execution_start_scope_conflict") + else: + row.projection_status = "applied" + row.projection_error_code = "execution_predates_activation" + return + execution = session.scalar( + sa.select(AgentSandboxExecution).where( + AgentSandboxExecution.provider == "e2b", + AgentSandboxExecution.provider_project_id == row.provider_project_id, + AgentSandboxExecution.provider_execution_id == row.provider_execution_id, + ) + ) + if execution is None: + execution = AgentSandboxExecution( + provider="e2b", + provider_project_id=row.provider_project_id, + sandbox_id=row.sandbox_id, + provider_execution_id=row.provider_execution_id, + started_at=started_at, + started_at_source="provider_execution" if started_at else None, + started_at_precision_ms=1000 if started_at else None, + state="open", + quality="pending", + attribution_status="unresolved", + ) + session.add(execution) + session.flush() + if execution.sandbox_id != row.sandbox_id or ( + execution.started_at is not None and started_at is not None and execution.started_at != started_at + ): + cls._conflict(session, row, "execution_identity_conflict", execution) + return + if started_at is not None and execution.started_at is None: + execution.started_at = started_at + execution.started_at_source = "provider_execution" + execution.started_at_precision_ms = 1000 + cls._attribute_execution(session, execution, _mapping(data["metadata"])) + for field, previous in ( + ("template_id", execution.template_id), + ("template_build_id", execution.template_build_id), + ("vcpu_count", execution.vcpu_count), + ("memory_mib", execution.memory_mib), + ): + value = data[field] + if previous is not None and value is not None and previous != value: + cls._conflict(session, row, "execution_resources_conflict", execution) + return + if value is not None: + setattr(execution, field, value) + if row.event_type in _TERMINAL_TYPES: + duration = data["duration_ms"] + if ( + execution.metered_duration_ms is not None + and duration is not None + and execution.metered_duration_ms != duration + ): + cls._conflict(session, row, "execution_duration_conflict", execution) + return + execution.state = "closed" + if duration is not None: + execution.metered_duration_ms = duration + if execution.quality != "conflict" and all( + value is not None + for value in ( + execution.metered_duration_ms, + execution.started_at, + execution.vcpu_count, + execution.memory_mib, + ) + ): + execution.quality = "metered" + execution.measurement_source = "provider_execution" + if execution.terminal_event_id is None: + execution.terminal_event_id = row.id + execution.terminal_event_at = row.occurred_at + execution.close_reason = row.event_type.rsplit(".", 1)[-1] + execution.last_reconciled_at = naive_utc_now() + execution.revision += 1 + if execution.quality == "conflict": + row.projection_status = "conflict" + row.projection_error_code = "execution_conflict" + elif execution.quality == "metered" or row.event_type in _NONTERMINAL_TYPES: + row.projection_status = "applied" + row.projection_error_code = None + else: + row.projection_status = "unresolved" + row.projection_error_code = "incomplete_provider_execution" + + @staticmethod + def _conflict( + session: Session, row: AgentSandboxUsageEvent, code: str, execution: AgentSandboxExecution | None = None + ) -> None: + row.projection_status = "conflict" + row.projection_error_code = code + if execution is None and row.provider_execution_id: + execution = session.scalar( + sa.select(AgentSandboxExecution).where( + AgentSandboxExecution.provider_project_id == row.provider_project_id, + AgentSandboxExecution.provider_execution_id == row.provider_execution_id, + ) + ) + if execution: + execution.quality = "conflict" + execution.revision += 1 diff --git a/api/tests/unit_tests/controllers/inner_api/test_runtime_usage.py b/api/tests/unit_tests/controllers/inner_api/test_runtime_usage.py new file mode 100644 index 00000000000..fc084e63bd4 --- /dev/null +++ b/api/tests/unit_tests/controllers/inner_api/test_runtime_usage.py @@ -0,0 +1,219 @@ +"""Authenticated HTTP contract backed by the real accounting service and DB.""" + +from collections.abc import Iterator +from typing import TypedDict +from uuid import uuid4 + +import pytest +import sqlalchemy as sa +from flask import Flask +from flask.testing import FlaskClient +from pydantic import JsonValue +from sqlalchemy.orm import Session, sessionmaker + +from controllers.inner_api import bp as inner_api_bp +from controllers.inner_api import inner_api_ns +from models.agent_sandbox_usage import AgentSandboxUsageEvent +from tests.unit_tests.config_override import config_overrides_context + +PROJECT = "431de237-596f-4d59-8a85-20a9846bf243" + + +class _EventPayload(TypedDict): + id: str + source: str + type: str + timestamp: str + payload: dict[str, JsonValue] + + +class _BatchPayload(TypedDict): + project_id: str + events: list[_EventPayload] + + +@pytest.fixture +def client() -> Iterator[FlaskClient]: + app = Flask(__name__) + app.config["TESTING"] = True + app.register_blueprint(inner_api_bp) + with config_overrides_context( + PLUGIN_DAEMON_KEY="test-daemon", + INNER_API_KEY_FOR_PLUGIN="test-inner", + AGENT_SANDBOX_METERING_ENABLED=True, + AGENT_SANDBOX_METERING_PROJECT_ID=PROJECT, + AGENT_SANDBOX_METERING_START_AT="2026-09-20T00:00:00Z", + ): + yield app.test_client() + + +def payload() -> _BatchPayload: + event_id = str(uuid4()) + return { + "project_id": PROJECT, + "events": [ + { + "id": event_id, + "source": "provider", + "type": "sandbox.lifecycle.paused", + "timestamp": "2026-09-20T00:01:00Z", + "payload": { + "id": event_id, + "version": "v2", + "type": "sandbox.lifecycle.paused", + "timestamp": "2026-09-20T00:01:00Z", + "sandboxId": "sandbox-test", + "sandboxExecutionId": "execution-test", + "sandboxTeamId": PROJECT, + "eventData": { + "execution": { + "started_at": "2026-09-20T00:00:02Z", + "execution_time": 12345, + "memory_mb": 1024, + "vcpu_count": 2, + } + }, + }, + } + ], + } + + +def test_state_requires_inner_auth(client: FlaskClient) -> None: + response = client.get(f"/inner/api/agent/sandbox-usage/state?project_id={PROJECT}") + assert response.status_code == 404 + + +def test_state_exposes_persisted_scope_without_secret_or_environment(client: FlaskClient) -> None: + response = client.get( + f"/inner/api/agent/sandbox-usage/state?project_id={PROJECT}", headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 200 + assert response.json == { + "enabled": True, + "project_id": PROJECT, + "started_at": "2026-09-20T00:00:00Z", + "checkpoint_at": None, + "full_scan_at": None, + "diagnostics": { + "unresolved_events": 0, + "conflict_events": 0, + "open_executions": 0, + "unattributed_executions": 0, + }, + } + + +@pytest.mark.parametrize("query", ["", "?project_id=", f"?project_id={PROJECT}&unexpected=value"]) +def test_state_rejects_invalid_query_before_activation( + client: FlaskClient, sqlite_session_factory: sessionmaker[Session], query: str +) -> None: + response = client.get(f"/inner/api/agent/sandbox-usage/state{query}", headers={"X-Inner-Api-Key": "test-inner"}) + assert response.status_code == 400 + body = response.get_json() + assert isinstance(body, dict) + assert body["code"] == "sandbox_usage_invalid_request" + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxUsageEvent)) == 0 + + +def test_state_rejects_project_outside_allowlist_without_activation( + client: FlaskClient, sqlite_session_factory: sessionmaker[Session] +) -> None: + response = client.get( + "/inner/api/agent/sandbox-usage/state?project_id=other-project", headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 403 + body = response.get_json() + assert isinstance(body, dict) + assert body["code"] == "sandbox_metering_project_not_allowed" + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxUsageEvent)) == 0 + + +def test_state_reports_start_change_conflict_and_preserves_activation(client: FlaskClient) -> None: + url = f"/inner/api/agent/sandbox-usage/state?project_id={PROJECT}" + headers = {"X-Inner-Api-Key": "test-inner"} + original = client.get(url, headers=headers) + assert original.status_code == 200 + with config_overrides_context(AGENT_SANDBOX_METERING_START_AT="2026-09-21T00:00:00Z"): + changed = client.get(url, headers=headers) + assert changed.status_code == 409 + body = changed.get_json() + assert isinstance(body, dict) + assert body["code"] == "sandbox_metering_start_is_immutable" + restored = client.get(url, headers=headers) + assert restored.status_code == 200 + assert restored.json == original.json + + +def test_state_reports_invalid_start_as_configuration_error(client: FlaskClient) -> None: + with config_overrides_context(AGENT_SANDBOX_METERING_START_AT="2026-09-20T00:00:00"): + response = client.get( + f"/inner/api/agent/sandbox-usage/state?project_id={PROJECT}", headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 503 + body = response.get_json() + assert isinstance(body, dict) + assert body["code"] == "sandbox_metering_start_not_configured" + + +def test_events_ack_only_after_persisted_and_duplicate_is_idempotent( + client: FlaskClient, sqlite_session_factory: sessionmaker[Session] +) -> None: + body = payload() + for expected in ( + {"accepted": 1, "duplicates": 0, "conflicts": 0, "ignored": 0}, + {"accepted": 0, "duplicates": 1, "conflicts": 0, "ignored": 0}, + ): + response = client.post( + "/inner/api/agent/sandbox-usage/events", json=body, headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 200 + assert response.json == expected + with sqlite_session_factory() as session: + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxUsageEvent) + .where(AgentSandboxUsageEvent.source_event_id == body["events"][0]["id"]) + ) + == 1 + ) + + +def test_ingestion_rejects_other_project_and_oversized_batch(client: FlaskClient) -> None: + body = payload() + body["project_id"] = "other-project" + response = client.post( + "/inner/api/agent/sandbox-usage/events", json=body, headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 403 + body = payload() + body["events"] *= 101 + response = client.post( + "/inner/api/agent/sandbox-usage/events", json=body, headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 400 + + +def test_ingestion_payload_limit(client: FlaskClient) -> None: + body = payload() + body["events"][0]["payload"] = {"oversized": "x" * (1024 * 1024)} + response = client.post( + "/inner/api/agent/sandbox-usage/events", json=body, headers={"X-Inner-Api-Key": "test-inner"} + ) + assert response.status_code == 413 + + +def test_schema_and_http_exclude_business_operation_envelope_fields(client: FlaskClient) -> None: + schema = inner_api_ns.models["SandboxUsageEvent"].__schema__ + assert set(schema["properties"]) == {"id", "source", "type", "timestamp", "sandbox_id", "execution_id", "payload"} + assert schema["additionalProperties"] is False + body = payload() + response = client.post( + "/inner/api/agent/sandbox-usage/events", + json={**body, "events": [{**body["events"][0], "allocation_id": str(uuid4())}]}, + headers={"X-Inner-Api-Key": "test-inner"}, + ) + assert response.status_code == 400 diff --git a/api/tests/unit_tests/extensions/test_celery_ssl.py b/api/tests/unit_tests/extensions/test_celery_ssl.py index c5df0a3b7ff..d6495785d71 100644 --- a/api/tests/unit_tests/extensions/test_celery_ssl.py +++ b/api/tests/unit_tests/extensions/test_celery_ssl.py @@ -162,6 +162,7 @@ class TestCelerySSLConfiguration: # Mock all the scheduler configs mock_config.CELERY_BEAT_SCHEDULER_TIME = 1 + mock_config.AGENT_SANDBOX_METERING_ENABLED = False mock_config.ENABLE_CONVERSATION_CLEANUP_TASK = False mock_config.CONVERSATION_CLEANUP_TASK_INTERVAL = 5 mock_config.ENABLE_CLEAN_EMBEDDING_CACHE_TASK = False @@ -212,6 +213,7 @@ class TestCelerySSLConfiguration: mock_config.CELERY_TASK_ANNOTATIONS = {} mock_config.CELERY_BEAT_SCHEDULER_TIME = 1 + mock_config.AGENT_SANDBOX_METERING_ENABLED = False mock_config.ENABLE_CONVERSATION_CLEANUP_TASK = True mock_config.CONVERSATION_CLEANUP_TASK_INTERVAL = 5 mock_config.ENABLE_CLEAN_EMBEDDING_CACHE_TASK = False diff --git a/api/tests/unit_tests/schedule/test_collect_agent_sandbox_usage.py b/api/tests/unit_tests/schedule/test_collect_agent_sandbox_usage.py new file mode 100644 index 00000000000..a9c0691b7e5 --- /dev/null +++ b/api/tests/unit_tests/schedule/test_collect_agent_sandbox_usage.py @@ -0,0 +1,179 @@ +"""The scheduled collector reuses existing Celery workers and bounded HTTP.""" + +from collections.abc import Iterator +from datetime import timedelta +from unittest.mock import MagicMock + +import httpx +import pytest +from celery.contrib.testing.app import setup_default_app + +from configs.extra.agent_backend_config import AgentBackendConfig +from dify_app import DifyApp +from extensions import ext_celery +from schedule.collect_agent_sandbox_usage import collect_agent_sandbox_usage +from tests.unit_tests.config_override import config_overrides_context + +PROJECT = "431de237-596f-4d59-8a85-20a9846bf243" + + +@pytest.fixture(autouse=True) +def restore_celery_app_state() -> Iterator[None]: + # init_app changes Celery's process-wide current/default app; close() alone + # leaves later shared tasks pointing at the now-unregistered test app. + with setup_default_app(collect_agent_sandbox_usage.app): + yield + + +@pytest.fixture(autouse=True) +def configured_collection() -> Iterator[None]: + with config_overrides_context( + AGENT_SANDBOX_METERING_ENABLED=True, + AGENT_SANDBOX_METERING_PROJECT_ID=PROJECT, + AGENT_BACKEND_BASE_URL="http://agent_backend:5001/", + AGENT_BACKEND_API_TOKEN="backend-token", + ): + yield + + +def response(payload: object, status: int = 200) -> httpx.Response: + return httpx.Response( + status, json=payload, request=httpx.Request("POST", "http://agent_backend:5001/internal/e2b/usage/collect") + ) + + +def test_task_uses_existing_queue_deadlines_and_single_authenticated_call(monkeypatch: pytest.MonkeyPatch) -> None: + post = MagicMock(return_value=response({"completed": True})) + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + assert collect_agent_sandbox_usage() is True + post.assert_called_once_with( + "http://agent_backend:5001/internal/e2b/usage/collect", + headers={"Authorization": "Bearer backend-token"}, + json={"project_id": PROJECT}, + max_retries=0, + timeout=260, + ) + dispatch = MagicMock() + monkeypatch.setattr(collect_agent_sandbox_usage.app, "send_task", dispatch) + collect_agent_sandbox_usage.apply_async() + assert dispatch.call_args.kwargs["queue"] == "ops_trace" + assert collect_agent_sandbox_usage.soft_time_limit == 270 + assert collect_agent_sandbox_usage.time_limit == 300 + + +def test_disabled_job_makes_no_http_request(monkeypatch: pytest.MonkeyPatch) -> None: + post = MagicMock() + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + with config_overrides_context(AGENT_SANDBOX_METERING_ENABLED=False): + assert collect_agent_sandbox_usage() is False + post.assert_not_called() + + +@pytest.mark.parametrize( + "missing", ["AGENT_BACKEND_BASE_URL", "AGENT_BACKEND_API_TOKEN", "AGENT_SANDBOX_METERING_PROJECT_ID"] +) +def test_incomplete_configuration_only_fails_collection_job(missing: str, monkeypatch: pytest.MonkeyPatch) -> None: + post = MagicMock() + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + with config_overrides_context(**{missing: ""}): + with pytest.raises(ValueError, match="requires backend URL, token, and project"): + collect_agent_sandbox_usage() + post.assert_not_called() + + +@pytest.mark.parametrize("payload", [{"completed": False}, {"completed": "true"}, {"unexpected": True}]) +def test_incomplete_or_invalid_scan_is_a_job_failure_without_retry( + payload: dict[str, object], monkeypatch: pytest.MonkeyPatch +) -> None: + post = MagicMock(return_value=response(payload)) + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + with pytest.raises((ValueError, RuntimeError)): + collect_agent_sandbox_usage() + assert post.call_count == 1 + + +@pytest.mark.parametrize("status", [403, 503, 504]) +def test_backend_failure_does_not_trigger_http_retry(status: int, monkeypatch: pytest.MonkeyPatch) -> None: + post = MagicMock(return_value=response({"detail": "unavailable"}, status)) + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + with pytest.raises(httpx.HTTPStatusError): + collect_agent_sandbox_usage() + assert post.call_count == 1 + + +def test_transport_failure_is_logged_without_credentials_or_raw_message( + monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + post = MagicMock(side_effect=httpx.ReadTimeout("do-not-log-secret")) + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + with pytest.raises(httpx.ReadTimeout): + collect_agent_sandbox_usage() + assert post.call_count == 1 + assert "Background sandbox usage collection failed" in caplog.text + assert "do-not-log-secret" not in caplog.text + assert "backend-token" not in caplog.text + + +@pytest.mark.parametrize("enabled", [False, True]) +@pytest.mark.parametrize("interval", [90, "90"]) +def test_beat_registers_only_enabled_collection_with_expiring_schedule( + enabled: bool, interval: int | str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(ext_celery, "setup_workflow_warm_shutdown_handler", lambda: None) + with config_overrides_context( + AGENT_SANDBOX_METERING_ENABLED=enabled, + AGENT_SANDBOX_METERING_INTERVAL_SECONDS=interval, + DISABLE_TELEMETRY=True, + ): + celery = ext_celery.init_app(DifyApp(__name__)) + try: + assert "schedule.collect_agent_sandbox_usage" in celery.conf.imports + schedule = celery.conf.beat_schedule.get("collect_agent_sandbox_usage") + if enabled: + assert schedule == { + "task": "schedule.collect_agent_sandbox_usage.collect_agent_sandbox_usage", + "schedule": timedelta(seconds=90), + "options": {"expires": 90}, + } + else: + assert schedule is None + finally: + celery.close() + + +@pytest.mark.parametrize("raw_interval", ["0", "-1", "not-a-number", "1.5", "9" * 30]) +def test_invalid_interval_env_only_disables_optional_schedule( + raw_interval: str, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + monkeypatch.setenv("AGENT_SANDBOX_METERING_INTERVAL_SECONDS", raw_interval) + config = AgentBackendConfig() + assert raw_interval == config.AGENT_SANDBOX_METERING_INTERVAL_SECONDS + monkeypatch.setattr(ext_celery, "setup_workflow_warm_shutdown_handler", lambda: None) + with config_overrides_context( + AGENT_SANDBOX_METERING_INTERVAL_SECONDS=config.AGENT_SANDBOX_METERING_INTERVAL_SECONDS, + ENABLE_CONVERSATION_CLEANUP_TASK=True, + CONVERSATION_CLEANUP_TASK_INTERVAL=2, + DISABLE_TELEMETRY=True, + ): + app = DifyApp(__name__) + celery = ext_celery.init_app(app) + try: + assert app.extensions["celery"] is celery + assert "collect_agent_sandbox_usage" not in celery.conf.beat_schedule + assert "schedule.collect_agent_sandbox_usage" in celery.conf.imports + assert celery.conf.beat_schedule["conversation_cleanup_sweeper"] == { + "task": "tasks.delete_conversation_task.sweep_deleted_conversations", + "schedule": timedelta(minutes=2), + } + finally: + celery.close() + assert "Skipping sandbox usage schedule" in caplog.text + assert all(raw_interval not in record.getMessage() for record in caplog.records) + + +def test_manual_collection_does_not_depend_on_beat_interval(monkeypatch: pytest.MonkeyPatch) -> None: + post = MagicMock(return_value=response({"completed": True})) + monkeypatch.setattr("schedule.collect_agent_sandbox_usage.ssrf_proxy.post", post) + with config_overrides_context(AGENT_SANDBOX_METERING_INTERVAL_SECONDS="invalid"): + assert collect_agent_sandbox_usage() is True + post.assert_called_once() diff --git a/api/tests/unit_tests/services/agent/test_runtime_usage_service.py b/api/tests/unit_tests/services/agent/test_runtime_usage_service.py new file mode 100644 index 00000000000..cc29dadf456 --- /dev/null +++ b/api/tests/unit_tests/services/agent/test_runtime_usage_service.py @@ -0,0 +1,720 @@ +"""Accounting invariants tested against real SQLAlchemy persistence.""" + +from collections.abc import Iterator +from copy import deepcopy +from datetime import datetime +from uuid import uuid4 + +import pytest +import sqlalchemy as sa +from pydantic import JsonValue +from sqlalchemy import event as sqlalchemy_event +from sqlalchemy.dialects import postgresql +from sqlalchemy.engine import Connection, Engine +from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.sql import ClauseElement + +from models.agent import ( + Agent, + AgentConfigVersionKind, + AgentScope, + AgentSource, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) +from models.agent_sandbox_usage import AgentSandboxExecution, AgentSandboxUsageEvent +from models.model import App +from services.agent.runtime_usage_service import SandboxUsageError, SandboxUsageEvent, SandboxUsageService +from tests.unit_tests.config_override import config_overrides_context + +PROJECT = "431de237-596f-4d59-8a85-20a9846bf243" +START = "2026-09-20T00:00:00Z" + + +@pytest.fixture(autouse=True) +def metering_config(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr("services.agent.runtime_usage_service.naive_utc_now", lambda: datetime(2026, 9, 21)) + with config_overrides_context( + AGENT_SANDBOX_METERING_ENABLED=True, + AGENT_SANDBOX_METERING_PROJECT_ID=PROJECT, + AGENT_SANDBOX_METERING_START_AT=START, + ): + yield + + +def provider_event( + *, + event_id: str = "event-1", + execution_id: str = "execution-1", + sandbox_id: str = "sandbox-1", + duration: int | None = 12345, + started_at: str | None = "2026-09-20T00:00:02Z", + kind: str = "paused", + metadata: dict[str, JsonValue] | None = None, +) -> SandboxUsageEvent: + return SandboxUsageEvent.model_validate( + { + "id": event_id, + "source": "provider", + "type": f"sandbox.lifecycle.{kind}", + "sandbox_id": sandbox_id, + "execution_id": execution_id, + "payload": { + "id": event_id, + "version": "v2", + "type": f"sandbox.lifecycle.{kind}", + "timestamp": "2026-09-20T00:00:14.551745815Z", + "sandboxId": sandbox_id, + "sandboxExecutionId": execution_id, + "sandboxTeamId": PROJECT, + "sandboxTemplateId": "template-1", + "sandboxBuildId": "build-1", + "eventData": { + "execution": { + "started_at": started_at, + "execution_time": duration, + "memory_mb": 1024, + "vcpu_count": 2, + }, + "sandbox_metadata": metadata or {}, + }, + }, + } + ) + + +def ingest(*events: SandboxUsageEvent) -> dict[str, int]: + return SandboxUsageService.ingest(project_id=PROJECT, events=list(events)) + + +def execution(factory: sessionmaker[Session]) -> AgentSandboxExecution: + with factory() as session: + return session.scalars(sa.select(AgentSandboxExecution)).one() + + +def test_duplicate_terminal_and_api_webhook_aliases_count_once(sqlite_session_factory: sessionmaker[Session]) -> None: + event = provider_event() + assert ingest(event)["accepted"] == 1 + alias = deepcopy(event.payload) + for camel, snake in ( + ("sandboxId", "sandbox_id"), + ("sandboxExecutionId", "sandbox_execution_id"), + ("sandboxTeamId", "sandbox_team_id"), + ("sandboxTemplateId", "sandbox_template_id"), + ("sandboxBuildId", "sandbox_build_id"), + ("eventData", "event_data"), + ): + alias[snake] = alias.pop(camel) + assert ingest(event.model_copy(update={"payload": alias}))["duplicates"] == 1 + assert ingest(provider_event(event_id="killed-event", kind="killed"))["accepted"] == 1 + row = execution(sqlite_session_factory) + assert row.metered_duration_ms == 12345 + assert row.quality == "metered" + assert row.started_at == datetime(2026, 9, 20, 0, 0, 2) + assert row.terminal_event_at == datetime(2026, 9, 20, 0, 0, 14, 551745) + + +def test_same_sandbox_resume_is_a_new_execution(sqlite_session_factory: sessionmaker[Session]) -> None: + ingest(provider_event(), provider_event(event_id="second", execution_id="execution-2", duration=6789)) + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.sum(AgentSandboxExecution.metered_duration_ms))) == 19134 + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxExecution)) == 2 + + +def test_state_reports_pending_and_conflict_counts_without_owner_ids() -> None: + created = provider_event(event_id="created", kind="created", duration=None, started_at=None) + created.payload["eventData"] = dict[str, JsonValue]() + ingest(created) + assert SandboxUsageService.get_state(project_id=PROJECT)["diagnostics"] == { + "unresolved_events": 0, + "conflict_events": 0, + "open_executions": 1, + "unattributed_executions": 1, + } + ingest(provider_event()) + ingest(provider_event(duration=1)) + diagnostics = SandboxUsageService.get_state(project_id=PROJECT)["diagnostics"] + assert diagnostics == { + "unresolved_events": 0, + "conflict_events": 1, + "open_executions": 0, + "unattributed_executions": 1, + } + + +def test_unknown_duration_stays_null_and_enrichment_is_replayable( + sqlite_session_factory: sessionmaker[Session], +) -> None: + missing = provider_event(duration=None) + ingest(missing) + assert ingest(missing)["duplicates"] == 1 + assert execution(sqlite_session_factory).metered_duration_ms is None + assert execution(sqlite_session_factory).quality == "pending" + assert ingest(provider_event())["accepted"] == 1 + assert execution(sqlite_session_factory).metered_duration_ms == 12345 + assert ingest(missing)["duplicates"] == 1 + with sqlite_session_factory() as session: + event = session.scalars( + sa.select(AgentSandboxUsageEvent).where(AgentSandboxUsageEvent.source_event_id == "event-1") + ).one() + assert event.payload["eventData"]["execution"]["execution_time"] is None + assert event.payload_versions + assert event.resolution["canonical"]["duration_ms"] == 12345 + + +def test_conflicting_event_keeps_both_payloads_and_quarantines_usage( + sqlite_session_factory: sessionmaker[Session], +) -> None: + ingest(provider_event()) + assert ingest(provider_event(duration=999))["conflicts"] == 1 + assert execution(sqlite_session_factory).quality == "conflict" + assert execution(sqlite_session_factory).metered_duration_ms == 12345 + with sqlite_session_factory() as session: + event = session.scalars( + sa.select(AgentSandboxUsageEvent).where(AgentSandboxUsageEvent.source_event_id == "event-1") + ).one() + assert event.payload_versions[0]["payload"]["eventData"]["execution"]["execution_time"] == 999 + ingest(provider_event(event_id="new-terminal")) + assert execution(sqlite_session_factory).quality == "conflict" + + +def test_conflicting_executions_never_add_or_overwrite_duration(sqlite_session_factory: sessionmaker[Session]) -> None: + ingest(provider_event()) + assert ingest(provider_event(event_id="other-terminal", duration=1))["conflicts"] == 1 + assert execution(sqlite_session_factory).metered_duration_ms == 12345 + + +def test_late_start_does_not_reopen_closed_execution(sqlite_session_factory: sessionmaker[Session]) -> None: + ingest(provider_event()) + ingest(provider_event(event_id="start", kind="resumed", duration=None)) + row = execution(sqlite_session_factory) + assert (row.state, row.quality, row.metered_duration_ms) == ("closed", "metered", 12345) + + +def test_activation_survives_restarts_and_config_changes_are_rejected() -> None: + state = SandboxUsageService.get_state(project_id=PROJECT) + assert state["started_at"] == START + assert SandboxUsageService.get_state(project_id=PROJECT) == state + with config_overrides_context(AGENT_SANDBOX_METERING_START_AT="2026-09-21T00:00:00Z"): + with pytest.raises(SandboxUsageError, match="immutable"): + SandboxUsageService.get_state(project_id=PROJECT) + + +def test_no_old_data_even_when_terminal_arrives_after_activation(sqlite_session_factory: sessionmaker[Session]) -> None: + assert ingest(provider_event(started_at="2026-09-19T23:59:59Z"))["ignored"] == 1 + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxExecution)) == 0 + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxUsageEvent) + .where(AgentSandboxUsageEvent.source == "provider") + ) + == 0 + ) + + +def test_t0_filter_cannot_hide_contradictory_existing_metered_event( + sqlite_session_factory: sessionmaker[Session], +) -> None: + ingest(provider_event()) + assert ingest(provider_event(started_at="2026-09-19T23:59:59Z"))["conflicts"] == 1 + assert execution(sqlite_session_factory).quality == "conflict" + + +def test_unknown_start_proven_before_t0_removes_pending_execution( + sqlite_session_factory: sessionmaker[Session], +) -> None: + unknown = provider_event(started_at=None, duration=None) + ingest(unknown) + assert ingest(provider_event(started_at="2026-09-19T23:59:59Z"))["ignored"] == 1 + ingest(provider_event(event_id="late-start", kind="created", started_at=None, duration=None)) + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxExecution)) == 0 + + +def test_unusable_provider_data_is_visible_but_not_metered(sqlite_session_factory: sessionmaker[Session]) -> None: + event = provider_event(started_at=None) + assert ingest(event)["accepted"] == 1 + with sqlite_session_factory() as session: + assert session.scalars(sa.select(AgentSandboxExecution)).one().quality == "pending" + row = session.scalars( + sa.select(AgentSandboxUsageEvent).where(AgentSandboxUsageEvent.source == "provider") + ).one() + assert row.projection_status == "unresolved" + + +def test_created_without_execution_data_then_terminal_and_late_start( + sqlite_session_factory: sessionmaker[Session], +) -> None: + created = provider_event(event_id="created", kind="created", duration=None, started_at=None) + created.payload["eventData"] = dict[str, JsonValue]() + assert ingest(created)["accepted"] == 1 + row = execution(sqlite_session_factory) + assert row.started_at is None + assert row.state == "open" + assert row.metered_duration_ms is None + ingest(provider_event()) + assert execution(sqlite_session_factory).quality == "metered" + assert ingest(created)["duplicates"] == 1 + late = created.model_copy(deep=True, update={"id": "late"}) + late.payload["id"] = "late" + ingest(late) + row = execution(sqlite_session_factory) + assert (row.state, row.quality, row.metered_duration_ms) == ("closed", "metered", 12345) + + +def test_missing_provider_project_is_rejected() -> None: + event = provider_event() + event.payload.pop("sandboxTeamId") + with pytest.raises(SandboxUsageError, match="provider_project_mismatch"): + ingest(event) + + +def test_retention_gap_is_saved_without_advancing_checkpoint(sqlite_session_factory: sessionmaker[Session]) -> None: + with config_overrides_context(AGENT_SANDBOX_METERING_START_AT="2026-09-01T00:00:00Z"): + event = SandboxUsageEvent( + id="gap", + source="application", + type="collector_retention_gap", + timestamp=START, + payload={ + "uncovered_start": "2026-09-01T00:00:00Z", + "uncovered_end": "2026-09-13T00:00:00Z", + "assumed_retention_seconds": 604800, + }, + ) + assert ingest(event)["accepted"] == 1 + assert SandboxUsageService.get_state(project_id=PROJECT)["checkpoint_at"] is None + with sqlite_session_factory() as session: + row = session.scalars( + sa.select(AgentSandboxUsageEvent).where(AgentSandboxUsageEvent.source_event_id == "gap") + ).one() + assert row.projection_status == "unresolved" + assert row.projection_error_code == "provider_retention_gap" + + +def test_checkpoint_only_advances_after_completed_scan() -> None: + event = SandboxUsageEvent( + id="checkpoint-1", + source="application", + type="collector_checkpoint", + timestamp=START, + payload={ + "scan_started_at": START, + "window_start": START, + "window_end": START, + "completed": True, + "mode": "full", + "pages": 1, + "events": 0, + }, + ) + ingest(event) + state = SandboxUsageService.get_state(project_id=PROJECT) + assert state["checkpoint_at"] == START + assert state["full_scan_at"] == START + with pytest.raises(SandboxUsageError, match="invalid_collector_checkpoint"): + ingest(event.model_copy(update={"id": "partial", "payload": {**event.payload, "completed": False}})) + + +def test_forbidden_project_and_reserved_application_events_fail_without_partial_batch( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with pytest.raises(SandboxUsageError, match="not_allowed"): + SandboxUsageService.ingest(project_id="other-project", events=[provider_event()]) + reserved = SandboxUsageEvent(id="forged", source="application", type="allocation_registered", payload={}) + with pytest.raises(SandboxUsageError, match="unsupported_application_event"): + ingest(provider_event(), reserved) + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxExecution)) == 0 + + +def test_provider_team_must_match_configured_project() -> None: + event = provider_event() + event.payload["sandboxTeamId"] = "other-project" + with pytest.raises(SandboxUsageError, match="provider_project_mismatch"): + ingest(event) + + +def test_disabled_ingestion_does_not_ack( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with config_overrides_context(AGENT_SANDBOX_METERING_ENABLED=False): + assert SandboxUsageService.get_state(project_id=PROJECT)["enabled"] is False + with pytest.raises(SandboxUsageError, match="disabled"): + ingest(provider_event()) + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxUsageEvent)) == 0 + + +@pytest.mark.parametrize( + ("timestamp", "error"), + [ + (123, "invalid_timestamp"), + ("not-a-time", "invalid_timestamp"), + ("2026-09-20T00:00:14", "timestamp_requires_timezone"), + ], +) +def test_invalid_provider_time_rolls_back_entire_batch( + timestamp: JsonValue, error: str, sqlite_session_factory: sessionmaker[Session] +) -> None: + bad = provider_event(event_id="bad-time", execution_id="bad-execution") + bad.payload["timestamp"] = timestamp + with pytest.raises(SandboxUsageError, match=error): + ingest(provider_event(), bad) + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxExecution)) == 0 + + +@pytest.mark.parametrize("quantity", [-1, True, 1.5, 2**63]) +def test_invalid_duration_is_rejected_before_any_accounting(quantity: JsonValue) -> None: + event = provider_event() + data = event.payload["eventData"] + assert isinstance(data, dict) + details = data["execution"] + assert isinstance(details, dict) + details["execution_time"] = quantity + with pytest.raises(SandboxUsageError, match="invalid_provider_quantity"): + ingest(event) + + +@pytest.mark.parametrize("field", ["id", "sandboxId", "sandboxExecutionId"]) +def test_raw_identity_cannot_disagree_with_delivery_envelope(field: str) -> None: + event = provider_event() + event.payload[field] = "different-id" + with pytest.raises(SandboxUsageError, match="provider_envelope_mismatch"): + ingest(event) + + +def test_invalid_provider_identifier_without_envelope_hint_is_rejected() -> None: + event = provider_event().model_copy(update={"sandbox_id": None}) + event.payload["sandboxId"] = "" + with pytest.raises(SandboxUsageError, match="invalid_provider_identifier"): + ingest(event) + + +@pytest.mark.parametrize( + ("configured_start", "error"), + [ + ("invalid", "sandbox_metering_start_not_configured"), + ("2026-09-20T00:00:00.001Z", "sandbox_metering_start_requires_whole_second"), + ], +) +def test_invalid_activation_configuration_cannot_create_ledger_scope(configured_start: str, error: str) -> None: + with config_overrides_context(AGENT_SANDBOX_METERING_START_AT=configured_start): + with pytest.raises(SandboxUsageError, match=error) as result: + SandboxUsageService.get_state(project_id=PROJECT) + assert result.value.status_code == 503 + + +@pytest.mark.parametrize( + ("key", "value"), + [ + ("dify.usage_allocation_id", 42), + ("dify.usage_allocation_id", "invalid-uuid"), + ("dify.binding_id", 42), + ("dify.binding_id", "invalid-uuid"), + ], +) +def test_invalid_owner_metadata_does_not_discard_metered_usage( + key: str, value: JsonValue, sqlite_session_factory: sessionmaker[Session] +) -> None: + assert ingest(provider_event(metadata={key: value}))["accepted"] == 1 + row = execution(sqlite_session_factory) + assert row.quality == "metered" + assert row.metered_duration_ms == 12345 + assert row.tenant_id is None + assert row.attribution_status == "unresolved" + + +@pytest.mark.parametrize("change", ["sandbox", "started_at", "resources"]) +def test_execution_identity_and_resource_conflicts_preserve_original_usage( + change: str, sqlite_session_factory: sessionmaker[Session] +) -> None: + ingest(provider_event()) + different = provider_event(event_id="conflicting-terminal") + if change == "sandbox": + different = provider_event(event_id="conflicting-terminal", sandbox_id="other-sandbox") + elif change == "started_at": + different = provider_event(event_id="conflicting-terminal", started_at="2026-09-20T00:00:03Z") + else: + data = different.payload["eventData"] + assert isinstance(data, dict) + details = data["execution"] + assert isinstance(details, dict) + details["vcpu_count"] = 4 + assert ingest(different)["conflicts"] == 1 + row = execution(sqlite_session_factory) + assert row.quality == "conflict" + assert row.sandbox_id == "sandbox-1" + assert row.vcpu_count == 2 + assert row.metered_duration_ms == 12345 + + +def test_unsupported_provider_contract_is_visible_without_fabricated_execution( + sqlite_session_factory: sessionmaker[Session], +) -> None: + event = provider_event() + event.payload["version"] = "unsupported" + assert ingest(event)["accepted"] == 1 + with sqlite_session_factory() as session: + assert session.scalar(sa.select(sa.func.count()).select_from(AgentSandboxExecution)) == 0 + row = session.scalars( + sa.select(AgentSandboxUsageEvent).where(AgentSandboxUsageEvent.source_event_id == event.id) + ).one() + assert row.projection_status == "unresolved" + assert row.projection_error_code == "incomplete_provider_execution" + + +def business_owner( + session: Session, + *, + binding_id: str | None = None, + tenant_id: str | None = None, + sandbox_id: str = "sandbox-1", +) -> AgentWorkspaceBinding: + tenant_id = tenant_id or str(uuid4()) + app_id, agent_id, workspace_id = (str(uuid4()) for _ in range(3)) + session.add_all( + [ + App( + id=app_id, + tenant_id=tenant_id, + name="Metering test", + description="", + mode="agent", + enable_site=False, + enable_api=False, + max_active_requests=None, + ), + Agent( + id=agent_id, tenant_id=tenant_id, name=str(uuid4()), scope=AgentScope.ROSTER, source=AgentSource.ROSTER + ), + AgentWorkspace( + id=workspace_id, + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id=str(uuid4()), + owner_scope_key="root", + backend_workspace_ref=sandbox_id, + ), + ] + ) + binding = AgentWorkspaceBinding( + id=binding_id or str(uuid4()), + tenant_id=tenant_id, + app_id=app_id, + agent_id=agent_id, + workspace_id=workspace_id, + agent_config_version_id=str(uuid4()), + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref=sandbox_id, + ) + session.add(binding) + session.commit() + return binding + + +def owner_metadata(binding: AgentWorkspaceBinding) -> dict[str, JsonValue]: + return { + "dify.binding_id": binding.id, + "dify.tenant_id": binding.tenant_id, + "dify.agent_id": binding.agent_id, + "dify.workspace_id": binding.workspace_id, + } + + +def test_provider_usage_without_business_rows_remains_metered(sqlite_session_factory: sessionmaker[Session]) -> None: + ingest(provider_event(metadata={"dify.binding_id": str(uuid4()), "dify.tenant_id": str(uuid4())})) + row = execution(sqlite_session_factory) + assert row.quality == "metered" + assert row.metered_duration_ms == 12345 + assert row.attribution_status == "unresolved" + assert row.tenant_id is None + assert row.allocation_id is None + + +def test_background_attribution_uses_verified_business_chain( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + binding = business_owner(sqlite_session) + ingest(provider_event(metadata=owner_metadata(binding))) + row = execution(sqlite_session_factory) + assert (row.tenant_id, row.app_id, row.agent_id, row.binding_id, row.workspace_id) == ( + binding.tenant_id, + binding.app_id, + binding.agent_id, + binding.id, + binding.workspace_id, + ) + assert row.attribution_status == "resolved" + assert row.allocation_id is None + with sqlite_session_factory() as session: + assert ( + session.scalar( + sa.select(sa.func.count()) + .select_from(AgentSandboxUsageEvent) + .where(AgentSandboxUsageEvent.event_type == "allocation_registered") + ) + == 0 + ) + + +def test_duplicate_provider_replay_can_resolve_business_rows_committed_later( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + binding_id, tenant_id = str(uuid4()), str(uuid4()) + event = provider_event(metadata={"dify.binding_id": binding_id, "dify.tenant_id": tenant_id}) + ingest(event) + assert execution(sqlite_session_factory).attribution_status == "unresolved" + binding = business_owner(sqlite_session, binding_id=binding_id, tenant_id=tenant_id) + assert ingest(event)["duplicates"] == 1 + row = execution(sqlite_session_factory) + assert row.binding_id == binding.id + assert row.attribution_status == "resolved" + assert row.metered_duration_ms == 12345 + + +@pytest.mark.parametrize("missing", ["binding", "workspace", "app", "agent"]) +def test_deleted_business_rows_do_not_erase_retained_owner_or_usage( + missing: str, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + binding = business_owner(sqlite_session) + event = provider_event(metadata=owner_metadata(binding)) + ingest(event) + if missing == "binding": + sqlite_session.delete(binding) + elif missing == "workspace": + sqlite_session.execute(sa.delete(AgentWorkspace).where(AgentWorkspace.id == binding.workspace_id)) + elif missing == "app": + sqlite_session.execute(sa.delete(App).where(App.id == binding.app_id)) + else: + sqlite_session.execute(sa.delete(Agent).where(Agent.id == binding.agent_id)) + sqlite_session.commit() + assert ingest(event)["duplicates"] == 1 + row = execution(sqlite_session_factory) + assert (row.binding_id, row.tenant_id) == (binding.id, binding.tenant_id) + assert row.attribution_status == "resolved" + assert row.metered_duration_ms == 12345 + + +@pytest.mark.parametrize( + "broken_link", + [ + "tenant_metadata", + "agent_metadata", + "workspace_metadata", + "binding_sandbox", + "workspace_sandbox", + "app_tenant", + "agent_tenant", + ], +) +def test_unverified_business_chain_never_guesses_an_owner( + broken_link: str, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + binding = business_owner(sqlite_session) + metadata = owner_metadata(binding) + if broken_link == "tenant_metadata": + metadata["dify.tenant_id"] = str(uuid4()) + elif broken_link in {"agent_metadata", "workspace_metadata"}: + metadata[f"dify.{broken_link.removesuffix('_metadata')}_id"] = str(uuid4()) + elif broken_link == "binding_sandbox": + binding.backend_binding_ref = "different-sandbox" + elif broken_link == "workspace_sandbox": + sqlite_session.execute( + sa.update(AgentWorkspace) + .where(AgentWorkspace.id == binding.workspace_id) + .values(backend_workspace_ref="different-sandbox") + ) + elif broken_link == "app_tenant": + sqlite_session.execute(sa.update(App).where(App.id == binding.app_id).values(tenant_id=str(uuid4()))) + else: + sqlite_session.execute(sa.update(Agent).where(Agent.id == binding.agent_id).values(tenant_id=str(uuid4()))) + sqlite_session.commit() + ingest(provider_event(metadata=metadata)) + row = execution(sqlite_session_factory) + assert row.quality == "metered" + assert row.tenant_id is None + assert row.attribution_status == "unresolved" + + +@pytest.mark.parametrize("bad_value", [None, 42, "invalid-uuid"]) +def test_bad_attribution_labels_do_not_invalidate_measured_runtime( + bad_value: JsonValue, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + binding = business_owner(sqlite_session) + ingest(provider_event(metadata=owner_metadata(binding))) + changed = owner_metadata(binding) + changed["dify.tenant_id"] = bad_value + assert ingest(provider_event(metadata=changed))["accepted"] == 1 + row = execution(sqlite_session_factory) + assert row.quality == "metered" + assert row.attribution_status == "resolved" + assert row.tenant_id == binding.tenant_id + assert row.metered_duration_ms == 12345 + + +def test_owner_conflict_does_not_reassign_history_or_discard_project_usage( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + first, second = business_owner(sqlite_session), business_owner(sqlite_session) + ingest(provider_event(metadata=owner_metadata(first))) + ingest(provider_event(event_id="other-owner", metadata=owner_metadata(second))) + row = execution(sqlite_session_factory) + assert row.attribution_status == "conflict" + assert row.binding_id == first.id + assert row.tenant_id == first.tenant_id + assert row.quality == "metered" + assert row.metered_duration_ms == 12345 + + +def test_same_sandbox_cannot_change_owner_across_executions( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + first, second = business_owner(sqlite_session), business_owner(sqlite_session) + ingest(provider_event(metadata=owner_metadata(first))) + resumed = provider_event(event_id="new-execution", execution_id="execution-2", metadata=owner_metadata(second)) + ingest(resumed) + # A subsequent replay with the original owner cannot erase the conflicting + # association. Both physical executions remain counted at project level. + ingest(provider_event(event_id="later-event", execution_id="execution-2", metadata=owner_metadata(first))) + with sqlite_session_factory() as session: + rows = list( + session.scalars(sa.select(AgentSandboxExecution).order_by(AgentSandboxExecution.provider_execution_id)) + ) + assert [(row.attribution_status, row.binding_id) for row in rows] == [ + ("resolved", first.id), + ("conflict", None), + ] + assert [row.quality for row in rows] == ["metered", "metered"] + assert sum(row.metered_duration_ms or 0 for row in rows) == 24690 + + +@pytest.mark.parametrize("event_type", ["operation_requested", "operation_observed", "allocation_registered"]) +def test_application_operations_are_no_longer_accepted(event_type: str) -> None: + event = SandboxUsageEvent(id="legacy-operation", source="application", type=event_type, payload={}) + with pytest.raises(SandboxUsageError, match="unsupported_application_event"): + ingest(event) + + +def test_steady_state_lookup_only_selects_without_locking_activation(sqlite_session: Session) -> None: + engine = sqlite_session.get_bind() + assert isinstance(engine, Engine) + initial = SandboxUsageService.get_state(project_id=PROJECT) + statements: list[str] = [] + + def observe(_connection: Connection, statement: ClauseElement, *_args: object) -> None: + statements.append(str(statement.compile(dialect=postgresql.dialect()))) + + sqlalchemy_event.listen(engine, "before_execute", observe) + try: + assert SandboxUsageService.get_state(project_id=PROJECT) == initial + finally: + sqlalchemy_event.remove(engine, "before_execute", observe) + assert statements + assert all(statement.lstrip().startswith("SELECT") for statement in statements) + assert all("FOR UPDATE" not in statement for statement in statements) diff --git a/api/tests/unit_tests/services/agent/test_workspace_service.py b/api/tests/unit_tests/services/agent/test_workspace_service.py index 4758a20410a..afdc4cfcd0a 100644 --- a/api/tests/unit_tests/services/agent/test_workspace_service.py +++ b/api/tests/unit_tests/services/agent/test_workspace_service.py @@ -637,3 +637,67 @@ def test_workspace_collection_final_delete_failure_propagates( assert exc_info.value is error client.destroy_execution_binding_sync.assert_called_once() + + +def test_binding_creation_ignores_unavailable_metering_configuration( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + apply_config_overrides( + monkeypatch, + AGENT_SANDBOX_METERING_ENABLED=True, + AGENT_SANDBOX_METERING_PROJECT_ID="", + AGENT_SANDBOX_METERING_START_AT="invalid", + ) + client = _backend_client() + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + binding = AgentWorkspaceService.create_binding( + session=sqlite_session, + scope=_scope(), + agent_id="agent-1", + base_home_snapshot_id=None, + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + sqlite_session.commit() + assert binding.backend_binding_ref == "binding-ref" + client.create_execution_binding_sync.assert_called_once() + + +def test_binding_lookups_ignore_bad_metering_configuration_and_remain_tenant_scoped( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + apply_config_overrides( + monkeypatch, + AGENT_SANDBOX_METERING_ENABLED=True, + AGENT_SANDBOX_METERING_PROJECT_ID="", + AGENT_SANDBOX_METERING_START_AT="invalid", + ) + workspace, binding = _workspace(), _binding() + sqlite_session.add_all([workspace, binding]) + sqlite_session.commit() + assert ( + AgentWorkspaceService.get_active_binding( + session=sqlite_session, + tenant_id=binding.tenant_id, + binding_id=binding.id, + expected_owner_scope=_scope(), + ) + is binding + ) + assert ( + AgentWorkspaceService.resolve_active_binding_for_scope( + session=sqlite_session, + scope=_scope(), + agent_id=binding.agent_id, + ) + is binding + ) + assert ( + AgentWorkspaceService.get_active_binding( + session=sqlite_session, + tenant_id="another-tenant", + binding_id=binding.id, + expected_owner_scope=_scope(), + ) + is None + ) diff --git a/dify-agent/docs/dify-agent/guide/index.md b/dify-agent/docs/dify-agent/guide/index.md index ae5613d836d..9fb5fe0eb96 100644 --- a/dify-agent/docs/dify-agent/guide/index.md +++ b/dify-agent/docs/dify-agent/guide/index.md @@ -2,7 +2,10 @@ This guide describes how to run the MVP Dify Agent API server. The server is implemented in `dify-agent/src/dify_agent/server/app.py` and uses Redis for run -records and per-run event streams only. +records and per-run event streams. Optional E2B usage collection is triggered by +the API's existing Celery Beat schedule through a separate one-shot endpoint. +It adds no operation logging, polling loop, leader election or accounting context +to business runtime execution. ## Default local startup @@ -32,6 +35,9 @@ also reads `.env` and `dify-agent/.env` when present. OpenShell-specific settings are listed in the [OpenShell configuration reference](openshell.md#configuration). +Independent E2B execution accounting, activation and rollout checks are described +in [E2B runtime metering](runtime-metering.md). + | Environment variable | Default | Description | | --- | --- | --- | | `DIFY_AGENT_REDIS_URL` | `redis://localhost:6379/0` | Redis connection URL. | diff --git a/dify-agent/docs/dify-agent/guide/runtime-metering.md b/dify-agent/docs/dify-agent/guide/runtime-metering.md new file mode 100644 index 00000000000..abae852ae0b --- /dev/null +++ b/dify-agent/docs/dify-agent/guide/runtime-metering.md @@ -0,0 +1,156 @@ +# E2B runtime metering + +One ledger row represents one E2B `sandboxExecutionId`, from start/resume to +pause/kill. Provider execution time and resource sizes determine compute usage; +LLM token usage, message latency and local SDK-call duration do not. + +## Scheduled collection and boundaries + +The Dify API's existing Celery Beat schedule dispatches +`schedule.collect_agent_sandbox_usage.collect_agent_sandbox_usage` every 60 seconds +by default, through the existing `ops_trace` worker queue. The job calls +`POST /internal/e2b/usage/collect` on the Agent backend using its existing Bearer +service token and a `project_id` body. The E2B key stays in the Agent backend. + +The provider-specific endpoint performs one bounded scan, using request-scoped +HTTP clients. It reads E2B lifecycle events and posts them to the Dify API's +idempotent usage ingestion endpoint. It returns `{"completed": true}` only after +all scanned pages and the checkpoint have been acknowledged. Partial scans return +`completed=false`; provider/API errors fail this accounting request only. + +There is no Agent-process polling loop, Redis leader election, log outbox, +observation worker, or requested/observed operation reporting. Generic leases, +Agent runs, files, snapshots and binding lookups do not participate in metering. +Their timing, cancellation handling and SDK retries retain their normal behavior. +Metering misconfiguration is checked when collection is invoked, rather than +becoming a prerequisite for business runtime startup or binding operations. + +The endpoint has a 240-second deadline. The scheduled HTTP caller has a +260-second timeout, and the Celery task configures 270/300-second soft/hard limits. +Pool support differs: gevent workers do not implement soft time limits, and +blocking work cannot rely on their hard limit. The collection and HTTP deadlines +are the primary bounds; verify the deployed worker pool before relying on Celery +limits. See the [Celery time-limit documentation](https://docs.celeryq.dev/en/stable/userguide/workers.html#time-limits). +Scheduled messages expire after one collection interval to avoid stale backlog. +There is no automatic HTTP retry; future scheduled scans can redeliver provider +IDs safely. Overlapping tasks may duplicate reads, but cannot multiply usage. + +## Data and attribution + +The API owns `agent_sandbox_usage_events` and `agent_sandbox_executions`. +Ingestion persists provider facts and projects execution rows atomically; the +execution key is `(provider, provider_project_id, provider_execution_id)`. +A short project transaction lock protects ingestion. Binding creation/lookups do +not register allocations, write accounting rows or acquire that lock. +Steady-state usage-state reads do not take a project `FOR UPDATE` lock; activation +can be initialized once when first needed. + +E2B's existing sandbox metadata contains binding, tenant, Agent and workspace +identifiers. During background ingestion, optional attribution checks those IDs +against the existing Binding → Workspace → App/Agent tenant-owned chain and the +physical sandbox reference. It never trusts labels alone. Missing, malformed or +already-deleted business records leave attribution unresolved while preserving +complete project-level provider usage. A previously verified ownership snapshot +survives later business deletion. Conflicting ownership is not silently reassigned. + +No accounting callback or transaction is added to the business path. There is no +pre-create allocation journal. Older diagnostic/allocation columns and records +are not destructively migrated in this revision; new accounting accepts only +provider events and collector checkpoint/coverage-gap control records. Table +retention is unchanged, with no new automatic cleanup policy. + +## Configuration + +Apply the existing usage-table migration before first enabling metering. +Configure Dify API, Beat and the workers consuming `ops_trace` consistently: + +```dotenv +AGENT_SANDBOX_METERING_ENABLED=true +AGENT_SANDBOX_METERING_PROJECT_ID= +AGENT_SANDBOX_METERING_START_AT= +AGENT_SANDBOX_METERING_INTERVAL_SECONDS=60 +``` + +The task reuses `AGENT_BACKEND_BASE_URL` and `AGENT_BACKEND_API_TOKEN`; the +existing SSRF-safe HTTP policy must permit that internal service. No E2B key is +copied into API or worker configuration. + +Configure the Agent backend's one-shot endpoint: + +```dotenv +DIFY_AGENT_SANDBOX_METERING_ENABLED=true +DIFY_AGENT_E2B_PROJECT_ID= +DIFY_AGENT_SANDBOX_METERING_OVERLAP_SECONDS=900 +DIFY_AGENT_SANDBOX_METERING_FULL_SCAN_INTERVAL_SECONDS=3600 +DIFY_AGENT_SANDBOX_METERING_MAX_PAGES=1000 +``` + +It reuses the existing E2B key, `DIFY_AGENT_API_TOKEN`, inner API URL and inner API +key. The new endpoint requires a nonempty Bearer-token configuration even when +legacy control-plane routes permit unauthenticated local development. +`DIFY_AGENT_SANDBOX_METERING_POLL_INTERVAL_SECONDS` is no longer used: the interval +belongs to Celery Beat on the API side. Both metering feature flags default off. + +Activation T0 is persisted and immutable. Only provider executions started at or +after T0 are metered; starting before T0 and ending after it does not include an +old execution. Resuming an older sandbox after T0 creates a new execution and is +included. There is no historical backfill and no environment column; databases +separate environments. A project ID also scopes and allowlists ingestion. + +## Failure semantics + +- Only complete provider terminal execution data populates metered duration. A + local `finally`, SDK expiry deadline or business success is not a billing clock. +- Missing terminal fields remain pending; conflicts remain visible. Late starts + cannot reopen closed executions, and duplicate facts do not add duration twice. +- Every scan starts provider offset pagination from zero, overlaps prior coverage + and periodically rescans the retained window. Failed, cancelled or page-bounded + scans never claim a completed checkpoint. SQL idempotency protects redelivery. +- The assumed default provider event retention is seven days. Older uncovered + intervals are recorded as gaps rather than reconstructed or counted as zero. +- A failed scheduled task affects accounting only. No accounting network request + is added to create/connect/pause/kill, lease helpers, runner or file operations. + +The API's private ingestion endpoints remain: +`GET /inner/api/agent/sandbox-usage/state?project_id=...` and +`POST /inner/api/agent/sandbox-usage/events`, authenticated with `X-Inner-Api-Key`. +Event batches are limited to 100 records and 1 MiB, acknowledged after commit. +They are for the isolated collector, not business-operation logging. + +## Querying usage + +For complete `quality='metered'` executions: + +```text +sandbox_hours = sum(metered_duration_ms) / 3600000 +vcpu_hours = sum(metered_duration_ms * vcpu_count) / 3600000 +ram_gib_hours = sum(metered_duration_ms * memory_mib) / (3600000 * 1024) +``` + +Project totals include complete executions with unresolved attribution; report +that count separately. Tenant/app reports require verified resolved attribution. +For UTC windows `[A, B)`, clip execution intervals using start plus provider +duration. Start time can have second precision, so cross-boundary allocation is +not claimed to be millisecond-accurate; complete-window totals conserve duration. + +## Rollout checks + +1. Verify Beat's schedule, task registration and an existing worker consuming + `ops_trace`, along with image revisions, flags, project and immutable T0. +2. Verify the authenticated one-shot endpoint, real scheduled execution and + advancing checkpoints. The Agent must not start a metering poller or leader. +3. Keep a real business binding until collection resolves its owner, then retire + it and confirm previously copied accounting remains. Also verify complete + project usage for test sandboxes without business records. +4. Exercise model success/failure/cancellation, files/list/read/download/errors, + snapshots/restoration and resource cleanup. Isolated bad accounting settings + or failing accounting-service stubs must not make business lookups fail. +5. Compare physical execution IDs, duration and resources against actual E2B + events; replay scans/events and confirm totals do not increase. +6. Verify old outbox and collector-leader keys are absent after deployment. Do + not alter run-store keys or stop shared services to inject accounting faults. + +Disable the metering flags and restart the relevant processes to stop collection; +business runtime behavior remains the same. Keep the ledger and original T0. +The independent stdout-only sandbox file-error fix is tracked separately from +this accounting integration. diff --git a/dify-agent/src/dify_agent/server/app.py b/dify-agent/src/dify_agent/server/app.py index 5af355f376c..cd5f1d650b0 100644 --- a/dify-agent/src/dify_agent/server/app.py +++ b/dify-agent/src/dify_agent/server/app.py @@ -5,7 +5,9 @@ instances for plugin-daemon and Dify API inner calls, route wiring, and a process-local scheduler. Run execution happens in background ``asyncio`` tasks rather than request handlers, so client disconnects do not cancel the agent runtime. Redis persists run records and per-run event streams with configured -retention only; it is not used as a job queue. Agenton layers and providers +retention. Optional provider accounting is exposed as a separate one-shot HTTP +endpoint driven by Celery Beat; it starts no tasks during lifespan startup and +adds no reporting to business operations. Redis is not used as a run job queue. Agenton layers and providers stay state-only: they borrow the lifespan-owned clients through the runner and receive runtime-backend and Shell settings through provider construction rather than reading environment variables themselves. The standard server mounts the @@ -29,6 +31,7 @@ from dify_agent.runtime.compositor_factory import create_default_layer_providers from dify_agent.runtime.run_scheduler import RunScheduler from dify_agent.server.auth import create_bearer_token_dependency from dify_agent.server.observability import configure_server_observability +from dify_agent.server.routes.e2b_usage import create_e2b_usage_router from dify_agent.server.routes.runs import create_runs_router from dify_agent.server.routes.execution_bindings import create_execution_bindings_router from dify_agent.server.routes.home_snapshots import create_home_snapshots_router @@ -146,6 +149,7 @@ def create_app(settings: ServerSettings | None = None) -> FastAPI: control_plane_router.include_router(create_home_snapshots_router(lambda: home_snapshot_service)) control_plane_router.include_router(create_binding_files_router(lambda: binding_file_service)) app.include_router(control_plane_router) + app.include_router(create_e2b_usage_router(resolved_settings)) app.include_router( create_agent_stub_router( token_codec=agent_stub_token_codec, diff --git a/dify-agent/src/dify_agent/server/e2b_usage_client.py b/dify-agent/src/dify_agent/server/e2b_usage_client.py new file mode 100644 index 00000000000..5bd36ad6496 --- /dev/null +++ b/dify-agent/src/dify_agent/server/e2b_usage_client.py @@ -0,0 +1,112 @@ +"""Provider-specific client for E2B usage facts and completed-scan checkpoints. + +Only the scheduled collector uses this client. Business operations do not report +accounting events or borrow this connection pool. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import UTC, datetime +import json +import logging +from typing import Any + +import httpx +from pydantic import BaseModel, ConfigDict, Field + +logger = logging.getLogger(__name__) +_MAX_BATCH_BYTES = 1024 * 1024 + + +class E2BUsageState(BaseModel): + enabled: bool + project_id: str + started_at: datetime | None = None + checkpoint_at: datetime | None = None + full_scan_at: datetime | None = None + + model_config = ConfigDict(extra="ignore") + + +class _IngestResult(BaseModel): + accepted: int = Field(ge=0) + duplicates: int = Field(ge=0) + conflicts: int = Field(ge=0) + ignored: int = Field(ge=0) + + +@dataclass(slots=True) +class E2BUsageApiClient: + """Use a collection-scoped inner client; credentials never reach E2B.""" + + client: httpx.AsyncClient + base_url: str + api_key: str + project_id: str + + async def get_state(self) -> E2BUsageState: + response = await self.client.get( + f"{self.base_url.rstrip('/')}/inner/api/agent/sandbox-usage/state", + params={"project_id": self.project_id}, + headers={"X-Inner-Api-Key": self.api_key}, + timeout=15.0, + ) + response.raise_for_status() + state = E2BUsageState.model_validate(response.json()) + if state.project_id != self.project_id: + raise ValueError("sandbox usage state project mismatch") + for value in (state.started_at, state.checkpoint_at, state.full_scan_at): + if value is not None and value.utcoffset() is None: + raise ValueError("sandbox usage state requires timezone-aware timestamps") + if state.enabled and state.started_at is None: + raise ValueError("enabled sandbox usage state requires a fixed activation time") + return state + + async def post_events(self, events: list[dict[str, Any]]) -> None: + if not events or len(events) > 100: + raise ValueError("sandbox usage batches require 1 to 100 events") + # The API limits the whole HTTP body, not just its event count. Splitting + # here covers provider pages. Repeated scheduled scans use stable + # provider IDs, so partial success cannot multiply execution usage. + batch: list[dict[str, Any]] = [] + for event in events: + candidate = [*batch, event] + body = self._encode_batch(candidate) + if len(body) > _MAX_BATCH_BYTES: + if not batch: + raise ValueError("sandbox usage event exceeds API request size") + await self._post_batch(self._encode_batch(batch), len(batch)) + batch = [event] + if len(self._encode_batch(batch)) > _MAX_BATCH_BYTES: + raise ValueError("sandbox usage event exceeds API request size") + else: + batch = candidate + if batch: + await self._post_batch(self._encode_batch(batch), len(batch)) + + def _encode_batch(self, events: list[dict[str, Any]]) -> bytes: + return json.dumps( + {"project_id": self.project_id, "events": events}, + ensure_ascii=False, + allow_nan=False, + separators=(",", ":"), + ).encode() + + async def _post_batch(self, body: bytes, count: int) -> None: + response = await self.client.post( + f"{self.base_url.rstrip('/')}/inner/api/agent/sandbox-usage/events", + headers={"X-Inner-Api-Key": self.api_key, "Content-Type": "application/json"}, + content=body, + timeout=15.0, + ) + response.raise_for_status() + result = _IngestResult.model_validate(response.json()) + if result.accepted + result.duplicates + result.conflicts + result.ignored != count: + raise ValueError("sandbox usage acknowledgement does not cover the entire batch") + if result.conflicts: + logger.error("sandbox usage ingestion reported conflicts", extra={"conflicts": result.conflicts}) + + +def utc_now() -> datetime: + return datetime.now(UTC) diff --git a/dify-agent/src/dify_agent/server/e2b_usage_collector.py b/dify-agent/src/dify_agent/server/e2b_usage_collector.py new file mode 100644 index 00000000000..9ee2fcb9c21 --- /dev/null +++ b/dify-agent/src/dify_agent/server/e2b_usage_collector.py @@ -0,0 +1,160 @@ +"""A single bounded E2B event scan, invoked by the existing Celery Beat schedule. + +Offsets are used within one bounded scan only. Every scan starts at zero, with +an overlap window and periodic rescan of the provider retention window. Neither +redelivery nor overlapping scheduled scans can inflate the API's idempotent +execution ledger. A partial scan never advances coverage. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime, timedelta +import logging +from typing import Any, Callable +from uuid import uuid4 + +import httpx + +from dify_agent.server.e2b_usage_client import E2BUsageApiClient, utc_now + +logger = logging.getLogger(__name__) +_PROVIDER_RETENTION = timedelta(days=7) +_EVENTS_URL = "https://api.e2b.app/events/sandboxes" + + +def _field(event: dict[str, Any], camel: str, snake: str) -> Any: + return event.get(camel, event.get(snake)) + + +def _timestamp(value: Any) -> datetime | None: + if value is None: + return None + if not isinstance(value, str): + raise ValueError("provider event timestamp is invalid") + parsed = datetime.fromisoformat(value.replace("Z", "+00:00")) + if parsed.utcoffset() is None: + raise ValueError("provider event timestamp must have a timezone") + return parsed + + +def provider_event(raw: dict[str, Any], *, project_id: str) -> dict[str, Any]: + """Preserve the raw provider payload; API owns normalization and T0 filtering.""" + project = _field(raw, "sandboxTeamId", "sandbox_team_id") + if project != project_id: + raise ValueError("E2B event project mismatch or missing project") + event_id = raw.get("id") + event_type = raw.get("type") + if not isinstance(event_id, str) or not event_id or not isinstance(event_type, str): + raise ValueError("E2B event missing stable identity") + _timestamp(raw.get("timestamp")) + return { + "id": event_id, + "source": "provider", + "type": event_type, + "timestamp": raw.get("timestamp"), + "sandbox_id": _field(raw, "sandboxId", "sandbox_id"), + "execution_id": _field(raw, "sandboxExecutionId", "sandbox_execution_id"), + "payload": raw, + } + + +@dataclass(slots=True) +class E2BUsageCollector: + provider_client: httpx.AsyncClient + usage_client: E2BUsageApiClient + api_key: str + project_id: str + overlap_seconds: int = 900 + full_scan_interval_seconds: int = 3600 + max_pages: int = 1000 + clock: Callable[[], datetime] = utc_now + + async def collect_once(self) -> bool: + state = await self.usage_client.get_state() + if not state.enabled: + return False + if state.started_at is None or state.project_id != self.project_id: + raise ValueError("sandbox usage activation/project state is invalid") + now = self.clock() + if state.started_at > now: + return False + retention_floor = now - _PROVIDER_RETENTION + full = ( + state.full_scan_at is None or (now - state.full_scan_at).total_seconds() >= self.full_scan_interval_seconds + ) + anchor = state.started_at if full else (state.checkpoint_at or state.started_at) + lower = max(state.started_at, retention_floor, anchor - timedelta(seconds=self.overlap_seconds)) + if (state.checkpoint_at or state.started_at) < retention_floor: + # The API checkpoint also records window_start, so a retention gap is + # visible instead of being misrepresented as complete T0 coverage. + logger.error("sandbox usage collection gap exceeds provider retention") + await self.usage_client.post_events( + [ + { + "id": str(uuid4()), + "source": "application", + "type": "collector_retention_gap", + "timestamp": now.isoformat(), + "payload": { + "uncovered_start": (state.checkpoint_at or state.started_at).isoformat(), + "uncovered_end": retention_floor.isoformat(), + "assumed_retention_seconds": int(_PROVIDER_RETENTION.total_seconds()), + }, + } + ] + ) + count = 0 + complete = False + pages = 0 + for page in range(self.max_pages): + response = await self.provider_client.get( + _EVENTS_URL, + headers={"X-API-Key": self.api_key}, + params={"limit": 100, "offset": page * 100, "orderAsc": "false"}, + timeout=15.0, + ) + response.raise_for_status() + raw_events = response.json() + if not isinstance(raw_events, list) or len(raw_events) > 100: + raise ValueError("E2B events response must be a page of up to 100 events") + pages += 1 + events: list[dict[str, Any]] = [] + all_older = bool(raw_events) + for raw in raw_events: + if not isinstance(raw, dict): + raise ValueError("invalid E2B event") + event = provider_event(raw, project_id=self.project_id) + occurred = _timestamp(raw.get("timestamp")) + if occurred is None or occurred >= lower: + all_older = False + events.append(event) + if events: + await self.usage_client.post_events(events) + count += len(events) + if len(raw_events) < 100 or all_older: + complete = True + break + if not complete: + logger.warning("sandbox usage scan reached page bound; checkpoint unchanged", extra={"pages": pages}) + return False + await self.usage_client.post_events( + [ + { + "id": str(uuid4()), + "source": "application", + "type": "collector_checkpoint", + "timestamp": now.isoformat(), + "payload": { + "scan_started_at": now.isoformat(), + "window_start": lower.isoformat(), + "window_end": now.isoformat(), + "completed": True, + "mode": "full" if full else "incremental", + "pages": pages, + "events": count, + }, + } + ] + ) + return True diff --git a/dify-agent/src/dify_agent/server/routes/e2b_usage.py b/dify-agent/src/dify_agent/server/routes/e2b_usage.py new file mode 100644 index 00000000000..35a1f37d775 --- /dev/null +++ b/dify-agent/src/dify_agent/server/routes/e2b_usage.py @@ -0,0 +1,99 @@ +"""One-shot provider accounting endpoint, invoked by the API's Celery task. + +The E2B credential stays in its existing service. Collection owns separate HTTP +clients, has no Redis leader or recurring task, and never touches runtime leases. +Invalid optional accounting configuration fails only this endpoint. +""" + +import asyncio +import logging + +import httpx +from fastapi import APIRouter, Depends, HTTPException +from pydantic import BaseModel, ConfigDict, Field + +from dify_agent.server.auth import create_bearer_token_dependency +from dify_agent.server.e2b_usage_client import E2BUsageApiClient +from dify_agent.server.e2b_usage_collector import E2BUsageCollector +from dify_agent.server.settings import ServerSettings + +logger = logging.getLogger(__name__) +_COLLECTION_TIMEOUT_SECONDS = 240.0 + + +class E2BUsageCollectRequest(BaseModel): + project_id: str = Field(min_length=1, max_length=128) + model_config = ConfigDict(extra="forbid") + + +class E2BUsageCollectResponse(BaseModel): + completed: bool + + +class _CollectionOptions(BaseModel): + overlap_seconds: int = Field(ge=1) + full_scan_interval_seconds: int = Field(ge=1) + max_pages: int = Field(ge=1, le=10000) + + +def create_e2b_usage_router(settings: ServerSettings) -> APIRouter: + async def require_configured_auth() -> None: + # Other legacy control-plane routes may permit a missing token. New + # collection requests must never expose the project key without auth. + if not settings.api_token: + raise HTTPException(status_code=503, detail="e2b_usage_collection_auth_unconfigured") + + router = APIRouter( + prefix="/internal/e2b/usage", + dependencies=[Depends(require_configured_auth), create_bearer_token_dependency(settings.api_token)], + ) + + @router.post("/collect", response_model=E2BUsageCollectResponse) + async def collect(request: E2BUsageCollectRequest) -> E2BUsageCollectResponse: + project_id = settings.e2b_project_id.strip() + if ( + not settings.sandbox_metering_enabled + or settings.runtime_backend != "e2b" + or not settings.e2b_api_key + or not project_id + or not settings.inner_api_key + ): + raise HTTPException(status_code=503, detail="e2b_usage_collection_unavailable") + if request.project_id != project_id: + raise HTTPException(status_code=403, detail="e2b_usage_project_mismatch") + try: + options = _CollectionOptions.model_validate( + { + "overlap_seconds": settings.sandbox_metering_overlap_seconds, + "full_scan_interval_seconds": settings.sandbox_metering_full_scan_interval_seconds, + "max_pages": settings.sandbox_metering_max_pages, + } + ) + async with ( + asyncio.timeout(_COLLECTION_TIMEOUT_SECONDS), + httpx.AsyncClient(timeout=15.0, trust_env=False) as provider_client, + httpx.AsyncClient(timeout=15.0, trust_env=False) as inner_client, + ): + collector = E2BUsageCollector( + provider_client=provider_client, + usage_client=E2BUsageApiClient( + client=inner_client, + base_url=settings.inner_api_url, + api_key=settings.inner_api_key, + project_id=project_id, + ), + api_key=settings.e2b_api_key, + project_id=project_id, + overlap_seconds=options.overlap_seconds, + full_scan_interval_seconds=options.full_scan_interval_seconds, + max_pages=options.max_pages, + ) + return E2BUsageCollectResponse(completed=await collector.collect_once()) + except TimeoutError as exc: + logger.warning("E2B usage collection exceeded its deadline") + raise HTTPException(status_code=504, detail="e2b_usage_collection_timed_out") from exc + except Exception as exc: + logger.warning("E2B usage collection failed", extra={"error_type": type(exc).__name__}) + raise HTTPException(status_code=503, detail="e2b_usage_collection_failed") from exc + + return router diff --git a/dify-agent/src/dify_agent/server/settings.py b/dify-agent/src/dify_agent/server/settings.py index a7722a67f37..2f0ab43183e 100644 --- a/dify-agent/src/dify_agent/server/settings.py +++ b/dify-agent/src/dify_agent/server/settings.py @@ -88,6 +88,13 @@ class ServerSettings(BaseSettings): le=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, ) e2b_shellctl_port: int = Field(default=5004, ge=1, le=65535) + sandbox_metering_enabled: bool = False + e2b_project_id: str = "" + # Optional accounting validates its numeric configuration at invocation, + # so a bad metering value cannot prevent ordinary runtime startup. + sandbox_metering_overlap_seconds: int | str = 900 + sandbox_metering_full_scan_interval_seconds: int | str = 3600 + sandbox_metering_max_pages: int | str = 1000 openshell_gateway_endpoint: str | None = None openshell_workspace: str = "default" openshell_bearer_token: str | None = None diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py index 932b66bb34d..445e862a402 100644 --- a/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py @@ -2,7 +2,7 @@ from __future__ import annotations import asyncio import posixpath -from collections.abc import Callable +from collections.abc import Awaitable, Callable from dataclasses import dataclass, field import logging from typing import cast @@ -191,7 +191,7 @@ class _ControlPlane: def _mock_http( monkeypatch: pytest.MonkeyPatch, - handler: Callable[[httpx.Request], httpx.Response], + handler: Callable[[httpx.Request], httpx.Response | Awaitable[httpx.Response]], ) -> list[httpx.AsyncClient]: original_async_client = httpx.AsyncClient transport = httpx.MockTransport(handler) @@ -1082,3 +1082,117 @@ async def test_e2b_acquire_preserves_health_failure_when_close_and_pause_fail( assert client.calls == 3 assert data_plane.close_calls == 1 assert sandbox.pauses == [True] + + +@pytest.mark.anyio +async def test_create_task_cancellation_during_initialization_kills_sandbox( + monkeypatch: pytest.MonkeyPatch, +) -> None: + entered = asyncio.Event() + + async def blocked_make_dir(_self: _Files, _path: str) -> bool: + entered.set() + await asyncio.Event().wait() + return True + + monkeypatch.setattr(_Files, "make_dir", blocked_make_dir) + control = _ControlPlane() + backend = _binding_backend(control) + task = asyncio.create_task( + backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant", + agent_id="agent", + binding_id="binding", + workspace_id="workspace", + existing_workspace_ref=None, + ) + ) + ) + async with asyncio.timeout(2): + await entered.wait() + task.cancel("cancel initialization") + with pytest.raises(asyncio.CancelledError, match="cancel initialization"): + await task + assert len(control.created) == 1 + assert control.sandboxes["sandbox-1"].killed == 1 + + +@pytest.mark.anyio +async def test_acquire_task_cancellation_during_health_closes_transport_and_pauses( + monkeypatch: pytest.MonkeyPatch, +) -> None: + entered = asyncio.Event() + + async def handler(_request: httpx.Request) -> httpx.Response: + entered.set() + await asyncio.Event().wait() + return httpx.Response(200, json={"status": "ok"}) + + clients = _mock_http(monkeypatch, handler) + backend, sandbox = _connected_backend() + task = asyncio.create_task(backend.acquire(sandbox.sandbox_id)) + async with asyncio.timeout(2): + await entered.wait() + task.cancel("cancel acquisition") + with pytest.raises(asyncio.CancelledError, match="cancel acquisition"): + await task + assert clients[0].is_closed + assert sandbox.pauses == [True] + + +@pytest.mark.anyio +async def test_release_task_cancellation_keeps_pause_retries_and_original_cancel( + sleep_delays: list[float], +) -> None: + entered = asyncio.Event() + + class BlockingCloseDataPlane: + async def close(self) -> None: + entered.set() + await asyncio.Event().wait() + + backend, sandbox = _connected_backend() + sandbox.pause_errors = [_transport_error(e2b_httpx.ReadTimeout), _transport_error(e2b_httpx.ReadTimeout)] + lease = E2BRuntimeLease( + sandbox=sandbox, # pyright: ignore[reportArgumentType] + data_plane=cast(ShellctlRuntimeLease, cast(object, BlockingCloseDataPlane())), + ) + task = asyncio.create_task(backend.release(lease)) + async with asyncio.timeout(2): + await entered.wait() + task.cancel("cancel release") + with pytest.raises(asyncio.CancelledError, match="cancel release"): + await task + assert sandbox.pauses == [True, True] + assert sleep_delays == [0.25] + + +@pytest.mark.anyio +async def test_acquire_primary_health_error_survives_task_cancel_during_compensation( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Preserve the pre-metering best-effort compensation contract: a failure in + # cleanup must not replace the primary acquisition error. + entered = asyncio.Event() + + async def blocked_pause(sandbox: _Sandbox, keep_memory: bool = True) -> bool: + sandbox.pauses.append(keep_memory) + entered.set() + await asyncio.Event().wait() + return True + + monkeypatch.setattr(_Sandbox, "pause", blocked_pause) + clients = _mock_http( + monkeypatch, + lambda _request: httpx.Response(401, json={"error": {"code": "unauthorized", "message": "bad token"}}), + ) + backend, sandbox = _connected_backend() + task = asyncio.create_task(backend.acquire(sandbox.sandbox_id)) + async with asyncio.timeout(2): + await entered.wait() + task.cancel("cancel compensation") + with pytest.raises(BindingAcquireError, match="bad token"): + await task + assert clients[0].is_closed + assert sandbox.pauses == [True] diff --git a/dify-agent/tests/local/dify_agent/server/test_app.py b/dify-agent/tests/local/dify_agent/server/test_app.py index 7fa209c8bd1..23c9da4343b 100644 --- a/dify-agent/tests/local/dify_agent/server/test_app.py +++ b/dify-agent/tests/local/dify_agent/server/test_app.py @@ -461,3 +461,35 @@ def test_server_settings_use_generic_outbound_http_args_for_shared_clients() -> assert "outbound_http_max_connections" in model_fields assert "outbound_http_max_keepalive_connections" in model_fields assert "outbound_http_keepalive_expiry" in model_fields + + +@pytest.mark.parametrize("enabled", [False, True]) +def test_optional_metering_never_starts_a_collector_or_blocks_runtime_startup( + monkeypatch: pytest.MonkeyPatch, enabled: bool +) -> None: + import dify_agent.server.routes.e2b_usage as usage_route + + _patch_app_lifecycle(monkeypatch) + + def unexpected_collector(*args: object, **kwargs: object) -> None: + raise AssertionError("collector must only be constructed by an explicit scheduled request") + + monkeypatch.setattr(usage_route, "E2BUsageCollector", unexpected_collector) + settings = ServerSettings( + _env_file=None, + sandbox_metering_enabled=enabled, + runtime_backend="local", + api_token="control-token", + e2b_project_id="", + e2b_api_key=None, + inner_api_key=None, + ) + with TestClient(create_app(settings)) as client: + assert client.get("/openapi.json").status_code == 200 + response = client.post( + "/internal/e2b/usage/collect", + headers={"Authorization": "Bearer control-token"}, + json={"project_id": "project"}, + ) + assert response.status_code == 503 + assert FakeRunScheduler.created[-1].shutdown_called diff --git a/dify-agent/tests/local/dify_agent/server/test_e2b_usage_client.py b/dify-agent/tests/local/dify_agent/server/test_e2b_usage_client.py new file mode 100644 index 00000000000..14bc57f264b --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_e2b_usage_client.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import asyncio +import json +from typing import Any + +import httpx +import pytest + +from dify_agent.server.e2b_usage_client import E2BUsageApiClient + + +def _event(event_id: str = "event-1") -> dict[str, Any]: + return {"id": event_id, "source": "provider", "type": "sandbox.lifecycle.paused", "payload": {}} + + +def _ack(request: httpx.Request) -> httpx.Response: + count = len(json.loads(request.content)["events"]) + return httpx.Response(200, json={"accepted": count, "duplicates": 0, "conflicts": 0, "ignored": 0}) + + +def test_client_rejects_wrong_project_state_and_incomplete_ack() -> None: + async def scenario() -> None: + async def handler(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response( + 200, json={"enabled": True, "project_id": "wrong", "started_at": "2026-09-20T00:00:00Z"} + ) + return httpx.Response(200, json={"accepted": 0, "duplicates": 0, "conflicts": 0, "ignored": 0}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http: + client = E2BUsageApiClient(http, "http://inner", "inner", "project") + with pytest.raises(ValueError, match="project mismatch"): + await client.get_state() + with pytest.raises(ValueError, match="entire batch"): + await client.post_events([_event()]) + + asyncio.run(scenario()) + + +def test_large_batches_respect_api_byte_limit_and_keep_stable_ids_on_retry() -> None: + async def scenario() -> None: + bodies: list[list[dict[str, Any]]] = [] + fail_second = True + + async def handler(request: httpx.Request) -> httpx.Response: + assert len(request.content) <= 1024 * 1024 + assert request.headers["Content-Type"] == "application/json" + batch = json.loads(request.content)["events"] + bodies.append(batch) + if fail_second and len(bodies) == 2: + return httpx.Response(503) + return httpx.Response(200, json={"accepted": len(batch), "duplicates": 0, "conflicts": 0, "ignored": 0}) + + events = [{**_event(str(i)), "payload": {"provider_metadata": "x" * 60000}} for i in range(40)] + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http: + client = E2BUsageApiClient(http, "http://inner", "inner", "project") + with pytest.raises(httpx.HTTPStatusError): + await client.post_events(events) + fail_second = False + await client.post_events(events) + assert bodies[0] == bodies[2] + assert [event["id"] for batch in bodies[2:] for event in batch] == [str(i) for i in range(40)] + + asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/server/test_e2b_usage_collector.py b/dify-agent/tests/local/dify_agent/server/test_e2b_usage_collector.py new file mode 100644 index 00000000000..07512848ca7 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_e2b_usage_collector.py @@ -0,0 +1,193 @@ +from __future__ import annotations + +import asyncio +from datetime import UTC, datetime, timedelta +import json +from typing import Any + +import httpx +import pytest + +from dify_agent.server.e2b_usage_collector import E2BUsageCollector, provider_event +from dify_agent.server.e2b_usage_client import E2BUsageApiClient + +NOW = datetime(2026, 9, 20, 1, tzinfo=UTC) +T0 = NOW - timedelta(hours=1) + + +def _provider(event_id: str = "event-1", timestamp: datetime = NOW) -> dict[str, Any]: + return { + "version": "v2", + "id": event_id, + "type": "sandbox.lifecycle.paused", + "sandboxId": "sandbox-1", + "sandboxTeamId": "project", + "sandboxExecutionId": "execution-1", + "timestamp": timestamp.isoformat(), + "eventData": { + "execution": {"execution_time": 1234, "memory_mb": 1024, "vcpu_count": 2, "started_at": T0.isoformat()} + }, + } + + +def _state(**kwargs: Any) -> dict[str, Any]: + return {"enabled": True, "project_id": "project", "started_at": T0.isoformat(), **kwargs} + + +def test_actual_v2_camel_case_shape_is_preserved_and_project_checked() -> None: + raw = _provider() + raw["timestamp"] = "2026-09-20T01:00:00.551745815Z" + event = provider_event(raw, project_id="project") + assert event["id"] == raw["id"] + assert event["execution_id"] == "execution-1" + assert event["payload"] == raw + with pytest.raises(ValueError, match="project mismatch"): + provider_event(raw, project_id="other") + + +def test_complete_scan_acks_all_pages_before_checkpoint_and_resets_offset() -> None: + async def scenario() -> None: + offsets: list[int] = [] + batches: list[list[dict[str, Any]]] = [] + + async def provider(request: httpx.Request) -> httpx.Response: + assert request.headers["X-API-Key"] == "provider-only" + assert "X-Inner-Api-Key" not in request.headers + assert request.url.params["orderAsc"] == "false" + offset = int(request.url.params["offset"]) + offsets.append(offset) + if offset == 0: + return httpx.Response(200, json=[_provider(str(i)) for i in range(100)]) + return httpx.Response(200, json=[_provider("old", T0 - timedelta(seconds=1))]) + + async def inner(request: httpx.Request) -> httpx.Response: + assert "X-API-Key" not in request.headers + if request.method == "GET": + return httpx.Response(200, json=_state()) + batch = json.loads(request.content)["events"] + batches.append(batch) + return httpx.Response(200, json={"accepted": len(batch), "duplicates": 0, "conflicts": 0, "ignored": 0}) + + async with ( + httpx.AsyncClient(transport=httpx.MockTransport(provider)) as p, + httpx.AsyncClient(transport=httpx.MockTransport(inner)) as i, + ): + collector = E2BUsageCollector( + p, + E2BUsageApiClient(i, "http://inner", "inner", "project"), + "provider-only", + "project", + clock=lambda: NOW, + ) + assert await collector.collect_once() + assert len(batches) == 2 + assert len(batches[0]) == 100 + checkpoint = batches[1][0] + assert checkpoint["type"] == "collector_checkpoint" + assert checkpoint["payload"]["window_start"] == T0.isoformat() + assert checkpoint["payload"]["pages"] == 2 + assert await collector.collect_once() + assert offsets == [0, 100, 0, 100] + + asyncio.run(scenario()) + + +@pytest.mark.parametrize("failure", ["max_pages", "api_failure", "wrong_project"]) +def test_incomplete_scan_never_advances_checkpoint(failure: str) -> None: + async def scenario() -> None: + events: list[dict[str, Any]] = [] + + async def provider(request: httpx.Request) -> httpx.Response: + raw = _provider() + if failure == "wrong_project": + raw["sandboxTeamId"] = "other" + return httpx.Response(200, json=[raw] * 100) + + async def inner(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json=_state()) + if failure == "api_failure": + return httpx.Response(503) + batch = json.loads(request.content)["events"] + events.extend(batch) + return httpx.Response(200, json={"accepted": len(batch), "duplicates": 0, "conflicts": 0, "ignored": 0}) + + async with ( + httpx.AsyncClient(transport=httpx.MockTransport(provider)) as p, + httpx.AsyncClient(transport=httpx.MockTransport(inner)) as i, + ): + collector = E2BUsageCollector( + p, + E2BUsageApiClient(i, "http://inner", "inner", "project"), + "provider", + "project", + max_pages=1, + clock=lambda: NOW, + ) + if failure == "max_pages": + assert not await collector.collect_once() + elif failure == "api_failure": + with pytest.raises(httpx.HTTPStatusError): + await collector.collect_once() + else: + with pytest.raises(ValueError, match="project mismatch"): + await collector.collect_once() + assert all(event["type"] != "collector_checkpoint" for event in events) + + asyncio.run(scenario()) + + +def test_incremental_overlap_and_retention_gap_are_explicit() -> None: + async def scenario() -> None: + state = _state(checkpoint_at=(NOW - timedelta(minutes=5)).isoformat(), full_scan_at=NOW.isoformat()) + sent: list[dict[str, Any]] = [] + + async def provider(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=[]) + + async def inner(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json=state) + batch = json.loads(request.content)["events"] + sent.extend(batch) + return httpx.Response(200, json={"accepted": len(batch), "duplicates": 0, "conflicts": 0, "ignored": 0}) + + async with ( + httpx.AsyncClient(transport=httpx.MockTransport(provider)) as p, + httpx.AsyncClient(transport=httpx.MockTransport(inner)) as i, + ): + collector = E2BUsageCollector( + p, E2BUsageApiClient(i, "http://inner", "inner", "project"), "provider", "project", clock=lambda: NOW + ) + await collector.collect_once() + assert sent[-1]["payload"]["window_start"] == (NOW - timedelta(minutes=20)).isoformat() + state.update(started_at=(NOW - timedelta(days=10)).isoformat(), checkpoint_at=None, full_scan_at=None) + await collector.collect_once() + assert sent[-2]["type"] == "collector_retention_gap" + assert sent[-1]["payload"]["window_start"] == (NOW - timedelta(days=7)).isoformat() + + asyncio.run(scenario()) + + +def test_disabled_or_future_activation_does_not_contact_e2b() -> None: + async def scenario() -> None: + async def provider(request: httpx.Request) -> httpx.Response: + raise AssertionError("provider must not be queried before activation") + + state = _state(enabled=False) + + async def inner(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=state) + + async with ( + httpx.AsyncClient(transport=httpx.MockTransport(provider)) as p, + httpx.AsyncClient(transport=httpx.MockTransport(inner)) as i, + ): + collector = E2BUsageCollector( + p, E2BUsageApiClient(i, "http://inner", "inner", "project"), "provider", "project", clock=lambda: NOW + ) + assert not await collector.collect_once() + state.update(enabled=True, started_at=(NOW + timedelta(hours=1)).isoformat()) + assert not await collector.collect_once() + + asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/server/test_e2b_usage_route.py b/dify-agent/tests/local/dify_agent/server/test_e2b_usage_route.py new file mode 100644 index 00000000000..65e02fc9c96 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_e2b_usage_route.py @@ -0,0 +1,176 @@ +from __future__ import annotations + +import asyncio +from datetime import UTC, datetime, timedelta +import json +from typing import Any + +from fastapi import FastAPI +from fastapi.testclient import TestClient +import httpx +import pytest + +import dify_agent.server.routes.e2b_usage as route_module +from dify_agent.server.routes.e2b_usage import create_e2b_usage_router +from dify_agent.server.settings import ServerSettings + + +def _settings(**overrides: Any) -> ServerSettings: + config = { + "_env_file": None, + "sandbox_metering_enabled": True, + "runtime_backend": "e2b", + "api_token": "control-only", + "e2b_api_key": "provider-only", + "e2b_project_id": "project", + "inner_api_key": "inner-only", + "inner_api_url": "http://inner", + } + return ServerSettings(**(config | overrides)) + + +def _app(settings: ServerSettings) -> FastAPI: + app = FastAPI() + app.include_router(create_e2b_usage_router(settings)) + return app + + +@pytest.mark.parametrize( + ("settings", "headers", "project", "status"), + [ + (_settings(), {}, "project", 401), + (_settings(), {"Authorization": "Bearer wrong"}, "project", 401), + (_settings(api_token=None), {}, "project", 503), + (_settings(sandbox_metering_enabled=False), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(e2b_project_id=" "), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(e2b_api_key=None), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(inner_api_key=None), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(runtime_backend="local"), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(sandbox_metering_max_pages=0), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(sandbox_metering_max_pages="invalid"), {"Authorization": "Bearer control-only"}, "project", 503), + (_settings(sandbox_metering_overlap_seconds=-1), {"Authorization": "Bearer control-only"}, "project", 503), + ( + _settings(sandbox_metering_full_scan_interval_seconds=""), + {"Authorization": "Bearer control-only"}, + "project", + 503, + ), + (_settings(), {"Authorization": "Bearer control-only"}, "wrong-project", 403), + (_settings(), {"Authorization": "Bearer control-only"}, "", 422), + ], +) +def test_collection_auth_and_configuration_errors_are_local_to_endpoint( + settings: ServerSettings, headers: dict[str, str], project: str, status: int, monkeypatch: pytest.MonkeyPatch +) -> None: + def unexpected_collector(*args: object, **kwargs: object) -> None: + raise AssertionError("invalid collection must not query the provider") + + monkeypatch.setattr(route_module, "E2BUsageCollector", unexpected_collector) + with TestClient(_app(settings)) as client: + response = client.post("/internal/e2b/usage/collect", headers=headers, json={"project_id": project}) + assert response.status_code == status + assert "provider-only" not in response.text and "inner-only" not in response.text + + +@pytest.mark.parametrize("failure", [None, "provider", "ack", "bounded"]) +def test_explicit_request_scans_once_and_only_completed_scan_commits_checkpoint( + monkeypatch: pytest.MonkeyPatch, failure: str | None +) -> None: + calls: list[str] = [] + batches: list[list[dict[str, Any]]] = [] + clients: list[httpx.AsyncClient] = [] + now = datetime.now(UTC) + event = { + "version": "v2", + "id": "event-1", + "type": "sandbox.lifecycle.paused", + "sandboxTeamId": "project", + "sandboxId": "sandbox", + "sandboxExecutionId": "execution", + "timestamp": now.isoformat(), + "eventData": {}, + } + + async def transport(request: httpx.Request) -> httpx.Response: + calls.append(request.url.host) + if request.url.host == "api.e2b.app": + assert request.headers["X-API-Key"] == "provider-only" + assert "X-Inner-Api-Key" not in request.headers + if failure == "provider": + return httpx.Response(503) + return httpx.Response(200, json=[event] * (100 if failure == "bounded" else 1)) + assert request.url.host == "inner" + assert request.headers["X-Inner-Api-Key"] == "inner-only" + assert "X-API-Key" not in request.headers + if request.method == "GET": + return httpx.Response( + 200, + json={"enabled": True, "project_id": "project", "started_at": (now - timedelta(hours=1)).isoformat()}, + ) + batch = json.loads(request.content)["events"] + batches.append(batch) + return httpx.Response( + 200, + json={"accepted": 0 if failure == "ack" else len(batch), "duplicates": 0, "conflicts": 0, "ignored": 0}, + ) + + original_client = httpx.AsyncClient + + def client_factory(**kwargs: Any) -> httpx.AsyncClient: + assert kwargs["trust_env"] is False + result = original_client(transport=httpx.MockTransport(transport), **kwargs) + clients.append(result) + return result + + monkeypatch.setattr(route_module.httpx, "AsyncClient", client_factory) + with TestClient(_app(_settings(sandbox_metering_max_pages="1", sandbox_metering_overlap_seconds="900"))) as client: + assert calls == [] # no startup polling, leadership or accounting IO + response = client.post( + "/internal/e2b/usage/collect", + headers={"Authorization": "Bearer control-only"}, + json={"project_id": "project"}, + ) + assert len(clients) == 2 and all(client.is_closed for client in clients) + assert calls.count("api.e2b.app") == 1 + checkpoints = [event for batch in batches for event in batch if event["type"] == "collector_checkpoint"] + if failure in {"provider", "ack"}: + assert response.status_code == 503 + assert checkpoints == [] + else: + assert response.status_code == 200 + assert response.json() == {"completed": failure != "bounded"} + assert len(checkpoints) == (0 if failure == "bounded" else 1) + + +def test_collection_deadline_closes_clients_without_checkpoint(monkeypatch: pytest.MonkeyPatch) -> None: + clients: list[httpx.AsyncClient] = [] + posts: list[httpx.Request] = [] + + async def transport(request: httpx.Request) -> httpx.Response: + if request.url.host == "inner": + if request.method == "POST": + posts.append(request) + return httpx.Response( + 200, json={"enabled": True, "project_id": "project", "started_at": datetime.now(UTC).isoformat()} + ) + await asyncio.Event().wait() + raise AssertionError("unreachable") + + original_client = httpx.AsyncClient + + def client_factory(**kwargs: Any) -> httpx.AsyncClient: + result = original_client(transport=httpx.MockTransport(transport), **kwargs) + clients.append(result) + return result + + monkeypatch.setattr(route_module.httpx, "AsyncClient", client_factory) + monkeypatch.setattr(route_module, "_COLLECTION_TIMEOUT_SECONDS", 0.02) + with TestClient(_app(_settings())) as client: + response = client.post( + "/internal/e2b/usage/collect", + headers={"Authorization": "Bearer control-only"}, + json={"project_id": "project"}, + ) + assert response.status_code == 504 + assert posts == [] + assert len(clients) == 2 and all(client.is_closed for client in clients) diff --git a/dify-agent/tests/local/dify_agent/server/test_settings.py b/dify-agent/tests/local/dify_agent/server/test_settings.py index 3f975355579..561ac117caa 100644 --- a/dify-agent/tests/local/dify_agent/server/test_settings.py +++ b/dify-agent/tests/local/dify_agent/server/test_settings.py @@ -524,3 +524,47 @@ def test_server_settings_rejects_non_array_shell_redact_patterns(monkeypatch: py with pytest.raises(ValueError, match="must be a JSON array"): _ = settings.get_shell_redact_patterns() + + +@pytest.mark.parametrize( + "overrides", + [{"runtime_backend": "local"}, {"e2b_api_key": None}, {"e2b_project_id": " "}, {"inner_api_key": None}], +) +def test_metering_configuration_is_checked_only_by_collection_endpoint(overrides: dict[str, object]) -> None: + config: dict[str, object] = { + "sandbox_metering_enabled": True, + "runtime_backend": "e2b", + "e2b_api_key": "provider-key", + "e2b_project_id": "project", + "inner_api_key": "inner-key", + "_env_file": None, + } + config.update(overrides) + # Optional accounting must not prevent unrelated runtime startup. Missing + # project/credentials/backend compatibility are checked by the one-shot route. + settings = ServerSettings(**config) + assert settings.sandbox_metering_enabled + + +def test_metering_defaults_off_and_reads_env_when_explicitly_enabled(monkeypatch: pytest.MonkeyPatch) -> None: + assert not ServerSettings(_env_file=None).sandbox_metering_enabled + monkeypatch.setenv("DIFY_AGENT_SANDBOX_METERING_ENABLED", "true") + monkeypatch.setenv("DIFY_AGENT_RUNTIME_BACKEND", "e2b") + monkeypatch.setenv("DIFY_AGENT_E2B_API_KEY", "provider-key") + monkeypatch.setenv("DIFY_AGENT_E2B_PROJECT_ID", "project") + monkeypatch.setenv("DIFY_AGENT_INNER_API_KEY", "inner-key") + settings = ServerSettings(_env_file=None) + assert settings.sandbox_metering_enabled + assert settings.sandbox_metering_max_pages == 1000 + assert settings.sandbox_metering_overlap_seconds == 900 + + +@pytest.mark.parametrize("value", ["0", "-1", "not-a-number"]) +@pytest.mark.parametrize("enabled", ["true", "false"]) +def test_invalid_optional_metering_number_does_not_block_settings_startup( + monkeypatch: pytest.MonkeyPatch, value: str, enabled: str +) -> None: + monkeypatch.setenv("DIFY_AGENT_SANDBOX_METERING_ENABLED", enabled) + monkeypatch.setenv("DIFY_AGENT_SANDBOX_METERING_MAX_PAGES", value) + settings = ServerSettings(_env_file=None) + assert settings.sandbox_metering_max_pages == value diff --git a/docker/envs/core-services/dify-agent.env.example b/docker/envs/core-services/dify-agent.env.example index ab26fbe46b7..b5fea249412 100644 --- a/docker/envs/core-services/dify-agent.env.example +++ b/docker/envs/core-services/dify-agent.env.example @@ -49,6 +49,17 @@ DIFY_AGENT_E2B_TEMPLATE=difys-default-team/dify-agent-local-sandbox DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS=3600 DIFY_AGENT_E2B_SHELLCTL_PORT=5004 +# Independent sandbox usage ledger. API owns fixed activation time T0; no backfill. +# Enable only with matching API metering project settings and events API access. +# Celery Beat in the API triggers a protected one-shot collection endpoint. +# No operation logging, Agent polling loop, or collector Redis leader. +# Collection requires the existing DIFY_AGENT_API_TOKEN to be configured. +DIFY_AGENT_SANDBOX_METERING_ENABLED=false +DIFY_AGENT_E2B_PROJECT_ID= +DIFY_AGENT_SANDBOX_METERING_OVERLAP_SECONDS=900 +DIFY_AGENT_SANDBOX_METERING_FULL_SCAN_INTERVAL_SECONDS=3600 +DIFY_AGENT_SANDBOX_METERING_MAX_PAGES=1000 + # OpenShell backend (DIFY_AGENT_RUNTIME_BACKEND=openshell): sandboxes run on a # self-hosted NVIDIA OpenShell gateway. Setup and validation: # dify-agent/docs/dify-agent/guide/openshell.md diff --git a/docker/envs/core-services/shared.env.example b/docker/envs/core-services/shared.env.example index 3b38ef23c3b..2615732b34c 100644 --- a/docker/envs/core-services/shared.env.example +++ b/docker/envs/core-services/shared.env.example @@ -127,6 +127,13 @@ AGENT_BACKEND_STREAM_MAX_RECONNECTS=3 # Client timeout for Agent backend calls that may carry a Home Snapshot transfer. # Must stay above DIFY_AGENT_ENTERPRISE_SANDBOX_SNAPSHOT_TIMEOUT on the Agent backend. AGENT_BACKEND_HOME_SNAPSHOT_TIMEOUT_SECONDS=45 +# Independent E2B accounting, collected only by API Celery Beat/workers. +# Uses the existing ops_trace queue; business binding operations never call metering. +# Enable after migrations. START_AT is an immutable, whole-second UTC T0. +AGENT_SANDBOX_METERING_ENABLED=false +AGENT_SANDBOX_METERING_PROJECT_ID= +AGENT_SANDBOX_METERING_START_AT= +AGENT_SANDBOX_METERING_INTERVAL_SECONDS=60 # Outer execution limit for Agent Apps. Agent Backend may stop the internal # model/tool loop sooner when DIFY_AGENT_RUN_TIMEOUT_SECONDS is lower. APP_MAX_EXECUTION_TIME=3600