Compare commits

...
Author SHA1 Message Date
QuantumGhostandGitHub b5e35fc2fc fix: remove hardcoded sandbox path in configuration file (#38618) 2026-07-09 13:30:56 +00:00
wangxiaoleiandGitHub 7bfbb2bbe8 fix: fix auth prefix duplicate (#38616) 2026-07-09 12:55:58 +00:00
QuantumGhostandGitHub d177998255 chore: Bump version to 1.16.0-rc1 (#38600) 2026-07-09 11:42:09 +00:00
wangxiaoleiGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
74ee665af6 fix: fix miss session param (#38612)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-09 11:26:16 +00:00
Stephen ZhouandGitHub d655153a3d fix(web): update docs links (#38591) 2026-07-09 11:25:31 +00:00
zyssyz123GitHub盐粒 YanliyyhJoelautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>QuantumGhostyyh盐粒 Yanli
925f97be20 feat: daily sync (#38593)
Co-authored-by: 盐粒 Yanli <[email protected]>
Co-authored-by: yyh <[email protected]>
Co-authored-by: Joel <[email protected]>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: QuantumGhost <[email protected]>
Co-authored-by: yyh <[email protected]>
Co-authored-by: 盐粒 Yanli <[email protected]>
2026-07-09 10:35:15 +00:00
9252d81826 docs: remove Dify Premium on AWS Marketplace section from all READMEs (#38607)
Co-authored-by: Claude Fable 5 <[email protected]>
2026-07-09 10:09:59 +00:00
Harsh KashyapGitHubHarsh Kashyapautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Harsh Kashyap
9d5819a9c1 fix(api): ignore invalid utf8 cache payloads (#37835)
Co-authored-by: Harsh Kashyap <[email protected]>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Harsh Kashyap <[email protected]>
2026-07-09 09:52:49 +00:00
Ingram ZGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
99e3b1a401 fix: harden workflow archive DB retries (#38170)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-09 09:46:28 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
45957225cd chore: batch example #38419 (#38474)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-09 09:14:52 +00:00
2f99652203 fix: chunk workflow failure tracking data (#38598)
Co-authored-by: CodingOnStar <[email protected]>
2026-07-09 08:33:47 +00:00
dc1131b6df refactor(tests): replace logger mocks with caplog in trace provider tests (#38569)
Co-authored-by: Cursor <[email protected]>
2026-07-09 08:33:28 +00:00
Stephen ZhouandGitHub a5a7c762a3 refactor(web): split app context state atoms (#38588) 2026-07-09 07:42:02 +00:00
Coding On StarGitHubCodingOnStarautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
063e390c5d fix(web): preserve attribution from auth redirect (#38583)
Co-authored-by: CodingOnStar <[email protected]>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-09 07:40:55 +00:00
JoelandGitHub 3b3c25273a fix: guard chat tree against out-of-order parents (#38590) 2026-07-09 07:05:05 +00:00
555 changed files with 15240 additions and 1415 deletions
-2
View File
@@ -125,8 +125,6 @@ All of Dify's offerings come with corresponding APIs, so you could effortlessly
- **Dify for enterprise / organizations<br/>**
We provide additional enterprise-centric features. [Send us an email](mailto:[email protected]?subject=%5BGitHub%5DBusiness%20License%20Inquiry) to discuss your enterprise needs. <br/>
> For startups and small businesses using AWS, check out [Dify Premium on AWS Marketplace](https://aws.amazon.com/marketplace/pp/prodview-t22mebxzwjhu6) and deploy it to your own AWS VPC with one click. It's an affordable AMI offering with the option to create apps with custom logo and branding.
## Staying ahead
Star Dify on GitHub and be instantly notified of new releases.
+3
View File
@@ -663,6 +663,9 @@ PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600
PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400
INNER_API_KEY_FOR_PLUGIN=QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1
# Dify Agent backend
AGENT_BACKEND_BASE_URL=http://localhost:5050
# Marketplace configuration
MARKETPLACE_ENABLED=true
MARKETPLACE_API_URL=https://marketplace.dify.ai
+122 -36
View File
@@ -1,10 +1,12 @@
import datetime
import logging
import time
from collections.abc import Callable
from typing import TypedDict
import click
import sqlalchemy as sa
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_database import db
from libs.datetime_utils import naive_utc_now
@@ -12,6 +14,7 @@ from services.clear_free_plan_tenant_expired_logs import ClearFreePlanTenantExpi
from services.retention.conversation.messages_clean_policy import create_message_clean_policy
from services.retention.conversation.messages_clean_service import MessagesCleanService
from services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs import WorkflowRunCleanup
from services.retention.workflow_run.db_retry import run_with_db_retry
from services.retention.workflow_run.tenant_prefix import tenant_prefix_condition
from tasks.remove_app_and_related_data_task import delete_draft_variables_batch
@@ -35,6 +38,12 @@ class WorkflowRunArchiveTenantPlan(TypedDict):
unpaid_tenant_ids: list[str]
class WorkflowRunArchivePrefixStats(TypedDict):
tenant_ids: list[str]
workflow_runs: int
workflow_node_executions: int
def _normalize_utc_datetime(value: datetime.datetime) -> datetime.datetime:
if value.tzinfo is None:
return value.replace(tzinfo=datetime.UTC)
@@ -57,6 +66,7 @@ def _parse_tenant_prefixes(prefixes: str | None) -> list[str]:
def _get_archive_candidate_tenant_ids_by_prefix(
session: Session,
prefix: str,
*,
start_from: datetime.datetime | None,
@@ -75,7 +85,7 @@ def _get_archive_candidate_tenant_ids_by_prefix(
if start_from is not None:
conditions.append(WorkflowRun.created_at >= start_from)
tenant_ids = db.session.scalars(
tenant_ids = session.scalars(
sa.select(WorkflowRun.tenant_id).where(*conditions).distinct().order_by(WorkflowRun.tenant_id)
).all()
return list(tenant_ids)
@@ -102,8 +112,80 @@ def _filter_paid_workflow_archive_tenant_ids(tenant_ids: list[str]) -> tuple[lis
return paid_tenant_ids, unpaid_tenant_ids
def _run_archive_command_db_retry[T](operation_name: str, operation: Callable[[], T]) -> T:
return run_with_db_retry(operation_name, operation, logger=logger)
def _get_archive_candidate_tenant_ids_with_retry(
session_maker: sessionmaker[Session],
prefix: str,
*,
start_from: datetime.datetime | None,
end_before: datetime.datetime,
) -> list[str]:
def fetch_tenant_ids() -> list[str]:
with session_maker() as session:
return _get_archive_candidate_tenant_ids_by_prefix(
session,
prefix,
start_from=start_from,
end_before=end_before,
)
return _run_archive_command_db_retry(f"workflow archive tenant resolve for prefix {prefix}", fetch_tenant_ids)
def _get_archive_plan_prefix_stats(
session_maker: sessionmaker[Session],
prefix: str,
*,
start_from: datetime.datetime | None,
end_before: datetime.datetime,
) -> WorkflowRunArchivePrefixStats:
from graphon.enums import WorkflowExecutionStatus
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun
from services.retention.workflow_run.archive_paid_plan_workflow_run import WorkflowRunArchiver
def fetch_prefix_stats() -> WorkflowRunArchivePrefixStats:
with session_maker() as session:
tenant_ids = _get_archive_candidate_tenant_ids_by_prefix(
session,
prefix,
start_from=start_from,
end_before=end_before,
)
run_conditions = [
WorkflowRun.created_at < end_before,
WorkflowRun.status.in_(WorkflowExecutionStatus.ended_values()),
WorkflowRun.type.in_(WorkflowRunArchiver.ARCHIVED_TYPE),
tenant_prefix_condition(WorkflowRun.tenant_id, prefix),
]
if start_from is not None:
run_conditions.append(WorkflowRun.created_at >= start_from)
workflow_runs = (
session.scalar(sa.select(sa.func.count()).select_from(WorkflowRun).where(*run_conditions)) or 0
)
candidate_runs = sa.select(WorkflowRun.id).where(*run_conditions).subquery()
workflow_node_executions = (
session.scalar(
sa.select(sa.func.count())
.select_from(WorkflowNodeExecutionModel)
.join(candidate_runs, WorkflowNodeExecutionModel.workflow_run_id == candidate_runs.c.id)
)
or 0
)
return WorkflowRunArchivePrefixStats(
tenant_ids=tenant_ids,
workflow_runs=workflow_runs,
workflow_node_executions=workflow_node_executions,
)
return _run_archive_command_db_retry(f"workflow archive plan for prefix {prefix}", fetch_prefix_stats)
def _resolve_archive_tenant_ids_from_plan(
*,
session_maker: sessionmaker[Session],
tenant_ids: str | None,
tenant_prefixes: list[str],
start_from: datetime.datetime | None,
@@ -122,7 +204,8 @@ def _resolve_archive_tenant_ids_from_plan(
requested_tenant_ids = []
for prefix in tenant_prefixes:
requested_tenant_ids.extend(
_get_archive_candidate_tenant_ids_by_prefix(
_get_archive_candidate_tenant_ids_with_retry(
session_maker,
prefix,
start_from=start_from,
end_before=end_before,
@@ -143,6 +226,21 @@ def _resolve_archive_tenant_ids_from_plan(
)
def _safe_remove_scoped_session(context: str) -> None:
try:
db.session.remove()
except Exception:
logger.warning("Ignoring DB scoped-session cleanup error after %s", context, exc_info=True)
try:
db.session.registry.clear()
except Exception:
logger.warning("Ignoring DB scoped-session registry cleanup error after %s", context, exc_info=True)
try:
db.engine.dispose()
except Exception:
logger.warning("Ignoring DB engine dispose error after %s", context, exc_info=True)
def _resolve_archive_time_range(
*,
before_days: int,
@@ -349,10 +447,6 @@ def archive_workflow_runs_plan(
supported workflow types, and the requested created_at window. V2 bundle archive
does not maintain per-run archive logs, so this plan reports source-table volume.
"""
from graphon.enums import WorkflowExecutionStatus
from models.workflow import WorkflowNodeExecutionModel, WorkflowRun
from services.retention.workflow_run.archive_paid_plan_workflow_run import WorkflowRunArchiver
before_days, start_from, end_before = _resolve_archive_time_range(
before_days=before_days,
from_days_ago=from_days_ago,
@@ -364,37 +458,25 @@ def archive_workflow_runs_plan(
if include_archived:
click.echo(click.style("--include-archived is a no-op for V2 bundle archive plans.", fg="yellow"))
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
rows: list[WorkflowRunArchivePlanRow] = []
for prefix in _HEX_PREFIXES:
tenant_ids = _get_archive_candidate_tenant_ids_by_prefix(
prefix,
start_from=start_from,
end_before=plan_end_before,
)
try:
prefix_stats = _get_archive_plan_prefix_stats(
session_maker,
prefix,
start_from=start_from,
end_before=plan_end_before,
)
except Exception as exc:
logger.exception("Failed to build workflow archive plan for prefix %s", prefix)
raise click.ClickException(f"Failed to build workflow archive plan for prefix {prefix}.") from exc
tenant_ids = prefix_stats["tenant_ids"]
workflow_runs = prefix_stats["workflow_runs"]
workflow_node_executions = prefix_stats["workflow_node_executions"]
total_tenants = len(tenant_ids)
paid_tenant_ids, unpaid_tenant_ids = _filter_paid_workflow_archive_tenant_ids(tenant_ids)
run_conditions = [
WorkflowRun.created_at < plan_end_before,
WorkflowRun.status.in_(WorkflowExecutionStatus.ended_values()),
WorkflowRun.type.in_(WorkflowRunArchiver.ARCHIVED_TYPE),
tenant_prefix_condition(WorkflowRun.tenant_id, prefix),
]
if start_from is not None:
run_conditions.append(WorkflowRun.created_at >= start_from)
workflow_runs = (
db.session.scalar(sa.select(sa.func.count()).select_from(WorkflowRun).where(*run_conditions)) or 0
)
candidate_runs = sa.select(WorkflowRun.id).where(*run_conditions).subquery()
workflow_node_executions = (
db.session.scalar(
sa.select(sa.func.count())
.select_from(WorkflowNodeExecutionModel)
.join(candidate_runs, WorkflowNodeExecutionModel.workflow_run_id == candidate_runs.c.id)
)
or 0
)
rows.append(
WorkflowRunArchivePlanRow(
tenant_prefix=prefix,
@@ -574,17 +656,18 @@ def archive_workflow_runs(
)
)
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
try:
tenant_plan = _resolve_archive_tenant_ids_from_plan(
session_maker=session_maker,
tenant_ids=tenant_ids,
tenant_prefixes=parsed_tenant_prefixes,
start_from=start_from,
end_before=plan_end_before,
)
except Exception:
except Exception as exc:
logger.exception("Failed to resolve workflow archive tenant plan")
click.echo(click.style("Failed to resolve workflow archive tenant plan.", fg="red"))
return
raise click.ClickException("Failed to resolve workflow archive tenant plan.") from exc
planned_tenant_ids = tenant_plan["archive_tenant_ids"]
planned_paid_tenant_ids = tenant_plan["paid_tenant_ids"] if planned_tenant_ids is not None else None
@@ -616,7 +699,10 @@ def archive_workflow_runs(
dry_run=dry_run,
delete_after_archive=delete_after_archive,
)
summary = archiver.run()
try:
summary = archiver.run()
finally:
_safe_remove_scoped_session("archive workflow run command")
click.echo(
click.style(
f"Summary: processed={summary.total_runs_processed}, archived={summary.runs_archived}, "
+3 -4
View File
@@ -25,11 +25,10 @@ class AgentBackendConfig(BaseSettings):
AGENT_SHELL_ENABLED: bool = Field(
description=(
"Inject the dify.shell layer (sandboxed bash workspace) into Agent runs. "
"Requires the agent backend to be wired with a shellctl entrypoint; keep it "
"off until shellctl is deployed, otherwise every agent run that includes the "
"shell layer will fail."
"Requires the agent backend to be wired with a shellctl entrypoint before "
"shell-using Agent runs are executed."
),
default=False,
default=True,
)
AGENT_APP_TEXT_DELTA_DEBOUNCE_SECONDS: NonNegativeFloat = Field(
+4 -1
View File
@@ -363,7 +363,10 @@ class FileAccessConfig(BaseSettings):
INTERNAL_FILES_URL: str = Field(
description="Internal base URL for file access within Docker network,"
" used for plugin daemon and internal service communication."
" Falls back to FILES_URL if not specified.",
" Explicit INTERNAL_FILES_URL takes precedence; otherwise SERVER_CONSOLE_API_URL is used,"
" then FILES_URL.",
validation_alias=AliasChoices("INTERNAL_FILES_URL", "SERVER_CONSOLE_API_URL"),
alias_priority=1,
default="",
)
+2 -4
View File
@@ -1,5 +1,3 @@
from mimetypes import guess_extension
from flask import request
from flask_restx import Resource
from flask_restx.api import HTTPStatus
@@ -8,7 +6,7 @@ from werkzeug.exceptions import Forbidden
import services
from core.tools.signature import verify_plugin_file_signature
from core.tools.tool_file_manager import ToolFileManager
from core.tools.tool_file_manager import ToolFileManager, resolve_extension
from core.workflow.file_reference import build_file_reference
from fields.file_fields import FileResponse
@@ -110,7 +108,7 @@ class PluginUploadFileApi(Resource):
conversation_id=args.conversation_id,
)
extension = guess_extension(tool_file.mimetype) or ".bin"
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
preview_url = ToolFileManager.sign_file(tool_file_id=tool_file.id, extension=extension)
# Create a dictionary with all the necessary attributes
@@ -476,6 +476,7 @@ class PluginDownloadFileRequestApi(Resource):
user_from=payload.user_from,
invoke_from=payload.invoke_from,
file_mapping=payload.file.model_dump(mode="python", exclude_none=True),
for_external=payload.for_external,
)
return BaseBackwardsInvocationResponse(
data={
@@ -333,6 +333,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
target=self._generate_worker,
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"session": db.session(),
"context": context,
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
+42 -4
View File
@@ -27,6 +27,7 @@ from clients.agent_backend import (
AgentBackendInternalEventType,
AgentBackendRunClient,
AgentBackendRunEventAdapter,
AgentBackendRunFailedInternalEvent,
AgentBackendRunSucceededInternalEvent,
AgentBackendStreamInternalEvent,
extract_runtime_layer_specs,
@@ -57,6 +58,14 @@ from core.workflow.nodes.agent_v2.ask_human_resume import build_deferred_tool_re
from extensions.ext_database import db
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, PromptMessage, UserPromptMessage
from graphon.model_runtime.errors.invoke import (
InvokeAuthorizationError,
InvokeBadRequestError,
InvokeConnectionError,
InvokeError,
InvokeRateLimitError,
InvokeServerUnavailableError,
)
from models.agent_config_entities import AgentSoulConfig
from models.enums import CreatorUserRole
from models.model import MessageAgentThought
@@ -71,6 +80,22 @@ class _DefaultSessionScopeSnapshotId:
_DEFAULT_SESSION_SCOPE_SNAPSHOT_ID = _DefaultSessionScopeSnapshotId()
_AGENT_BACKEND_INVOKE_ERROR_BY_REASON: Mapping[str, type[InvokeError]] = {
"InvokeAuthorizationError": InvokeAuthorizationError,
"InvokeBadRequestError": InvokeBadRequestError,
"CredentialsValidateFailedError": InvokeBadRequestError,
"InvokeConnectionError": InvokeConnectionError,
"InvokeRateLimitError": InvokeRateLimitError,
"InvokeServerUnavailableError": InvokeServerUnavailableError,
}
def _agent_backend_failure_to_exception(event: AgentBackendRunFailedInternalEvent) -> Exception:
err_cls = _AGENT_BACKEND_INVOKE_ERROR_BY_REASON.get(event.reason or "")
if err_cls is not None:
return err_cls(event.error)
return AgentBackendError(event.error or "Agent backend run did not complete successfully.")
def _prompt_messages_from_query(user_query: str | None) -> list[PromptMessage]:
if not user_query:
@@ -412,12 +437,15 @@ class _AgentProcessRecorder:
def _lookup_tool_thought(self, *, index: int, tool_call_id: str | None) -> str | None:
if tool_call_id and tool_call_id in self._tool_by_call_id:
return self._tool_by_call_id[tool_call_id]
if index < 0:
return None
return self._tool_by_index.get(index)
def _remember_tool_thought(
self, *, index: int, tool_call_id: str | None, tool_name: str | None, thought_id: str
) -> None:
self._tool_by_index[index] = thought_id
if index >= 0:
self._tool_by_index[index] = thought_id
if tool_call_id:
self._tool_by_call_id[tool_call_id] = thought_id
if tool_name:
@@ -433,6 +461,10 @@ class _AgentProcessRecorder:
return None
def _mark_tool_observed(self, thought_id: str) -> None:
self._tool_by_index = {index: value for index, value in self._tool_by_index.items() if value != thought_id}
self._tool_by_call_id = {
tool_call_id: value for tool_call_id, value in self._tool_by_call_id.items() if value != thought_id
}
for open_thought_ids in self._open_tool_by_name.values():
open_thought_ids.discard(thought_id)
@@ -530,7 +562,12 @@ def _event_index(data: dict[str, Any]) -> int:
def _string_or_none(value: Any) -> str | None:
return value if isinstance(value, str) and value else None
if not isinstance(value, str):
return None
normalized = value.strip()
if not normalized or normalized.lower() in {"none", "null"}:
return None
return normalized
def _json_or_text(value: Any) -> str | None:
@@ -652,8 +689,9 @@ class AgentAppRunner:
return
if not isinstance(terminal, AgentBackendRunSucceededInternalEvent):
error = getattr(terminal, "error", None) or "Agent backend run did not complete successfully."
raise AgentBackendError(str(error))
if isinstance(terminal, AgentBackendRunFailedInternalEvent):
raise _agent_backend_failure_to_exception(terminal)
raise AgentBackendError("Agent backend run did not complete successfully.")
answer = self._terminal_output_to_answer(terminal.output)
try:
@@ -205,6 +205,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
target=self._generate_worker,
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"session": db.session(),
"context": context,
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
@@ -8,7 +8,7 @@ from pydantic import JsonValue
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.task_entities import AppBlockingResponse, AppStreamResponse
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from graphon.model_runtime.errors.invoke import InvokeError
from graphon.model_runtime.errors.invoke import InvokeError, InvokeRateLimitError
logger = logging.getLogger(__name__)
@@ -127,6 +127,7 @@ class AppGenerateResponseConverter[TBlockingResponse: AppBlockingResponse](ABC):
},
ModelCurrentlyNotSupportError: {"code": "model_currently_not_support", "status": 400},
InvokeError: {"code": "completion_request_error", "status": 400},
InvokeRateLimitError: {"code": "rate_limit_error", "status": 429},
}
# Determine the response based on the type of exception
@@ -420,6 +420,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
message.total_price = usage.total_price
message.currency = usage.currency
self._task_state.llm_result.usage.latency = message.provider_response_latency
self._task_state.metadata.usage = self._task_state.llm_result.usage
message.message_metadata = self._task_state.metadata.model_dump_json()
if trace_manager:
+1 -1
View File
@@ -27,7 +27,7 @@ class ProviderCredentialsCache:
try:
cached_provider_credentials = cached_provider_credentials.decode("utf-8")
cached_provider_credentials = json.loads(cached_provider_credentials)
except JSONDecodeError:
except (JSONDecodeError, UnicodeDecodeError):
return None
return dict(cached_provider_credentials)
+1 -1
View File
@@ -24,7 +24,7 @@ class ProviderCredentialsCache(ABC):
try:
cached_credentials = cached_credentials.decode("utf-8")
return dict(json.loads(cached_credentials))
except JSONDecodeError:
except (JSONDecodeError, UnicodeDecodeError):
return None
return None
+1 -1
View File
@@ -30,7 +30,7 @@ class ToolParameterCache:
try:
cached_tool_parameter = cached_tool_parameter.decode("utf-8")
cached_tool_parameter = json.loads(cached_tool_parameter)
except JSONDecodeError:
except (JSONDecodeError, UnicodeDecodeError):
return None
return dict(cached_tool_parameter)
+1
View File
@@ -276,6 +276,7 @@ class RequestRequestDownloadFile(BaseModel):
"validation",
]
file: RequestDownloadFileMapping
for_external: bool = True
model_config = ConfigDict(extra="forbid")
+6 -5
View File
@@ -103,16 +103,17 @@ class ApiTool(Tool):
elif not isinstance(credentials["api_key_value"], str):
raise ToolProviderCredentialValidationError("api_key_value must be a string")
api_key_value = credentials["api_key_value"]
if "api_key_header_prefix" in credentials:
api_key_header_prefix = credentials["api_key_header_prefix"]
if api_key_header_prefix == "basic" and credentials["api_key_value"]:
credentials["api_key_value"] = f"Basic {credentials['api_key_value']}"
elif api_key_header_prefix == "bearer" and credentials["api_key_value"]:
credentials["api_key_value"] = f"Bearer {credentials['api_key_value']}"
if api_key_header_prefix == "basic" and api_key_value:
api_key_value = f"Basic {api_key_value}"
elif api_key_header_prefix == "bearer" and api_key_value:
api_key_value = f"Bearer {api_key_value}"
elif api_key_header_prefix == "custom":
pass
headers[api_key_header] = credentials["api_key_value"]
headers[api_key_header] = api_key_value
elif credentials["auth_type"] == "api_key_query":
# For query parameter authentication, we don't add anything to headers
+12 -3
View File
@@ -4,6 +4,7 @@ import hmac
import logging
import os
import time
import urllib.parse
from collections.abc import Generator
from mimetypes import guess_extension, guess_type
from uuid import uuid4
@@ -26,7 +27,7 @@ logger = logging.getLogger(__name__)
class ToolFileManager:
@staticmethod
def _build_graph_file_reference(tool_file: ToolFile) -> File:
extension = guess_extension(tool_file.mimetype) or ".bin"
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
return File(
file_type=get_file_type_by_mime_type(tool_file.mimetype),
transfer_method=FileTransferMethod.TOOL_FILE,
@@ -70,7 +71,7 @@ class ToolFileManager:
mimetype: str,
filename: str | None = None,
) -> ToolFile:
extension = guess_extension(mimetype) or ".bin"
extension = resolve_extension(filename=filename, mimetype=mimetype)
unique_name = uuid4().hex
unique_filename = f"{unique_name}{extension}"
# default just as before
@@ -120,7 +121,8 @@ class ToolFileManager:
or response.headers.get("Content-Type", "").split(";")[0].strip()
or "application/octet-stream"
)
extension = guess_extension(mimetype) or ".bin"
url_filename = os.path.basename(urllib.parse.urlparse(file_url).path)
extension = resolve_extension(filename=url_filename, mimetype=mimetype)
unique_name = uuid4().hex
filename = f"{unique_name}{extension}"
filepath = f"tools/{tenant_id}/{filename}"
@@ -220,4 +222,11 @@ def _factory() -> ToolFileManager:
return ToolFileManager()
def resolve_extension(*, filename: str | None, mimetype: str) -> str:
filename_extension = os.path.splitext(filename or "")[1].lower()
if filename_extension:
return filename_extension
return guess_extension(mimetype) or ".bin"
set_tool_file_manager_factory(_factory)
+5 -4
View File
@@ -3,7 +3,6 @@ import re
from collections.abc import Generator
from datetime import date, datetime
from decimal import Decimal
from mimetypes import guess_extension
from typing import Any
from uuid import UUID
@@ -11,7 +10,7 @@ import numpy as np
import pytz
from core.tools.entities.tool_entities import ToolInvokeMessage
from core.tools.tool_file_manager import ToolFileManager
from core.tools.tool_file_manager import ToolFileManager, resolve_extension
from core.workflow.file_reference import parse_file_reference
from graphon.file import File, FileTransferMethod, FileType
from libs.login import current_user
@@ -91,7 +90,8 @@ class ToolFileMessageTransformer:
conversation_id=conversation_id,
)
url = f"/files/tools/{tool_file.id}{guess_extension(tool_file.mimetype) or '.png'}"
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
url = cls.get_tool_file_url(tool_file_id=tool_file.id, extension=extension)
meta = cls._with_tool_file_meta(
message.meta,
tool_file_id=str(tool_file.id),
@@ -136,7 +136,8 @@ class ToolFileMessageTransformer:
filename=filename,
)
url = cls.get_tool_file_url(tool_file_id=tool_file.id, extension=guess_extension(tool_file.mimetype))
extension = resolve_extension(filename=tool_file.name, mimetype=tool_file.mimetype)
url = cls.get_tool_file_url(tool_file_id=tool_file.id, extension=extension)
meta = cls._with_tool_file_meta(meta, tool_file_id=str(tool_file.id))
# check if file is image
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any, override
from agenton.compositor import CompositorSessionSnapshot
from clients.agent_backend import (
AgentBackendAgentMessageDeltaInternalEvent,
AgentBackendDeferredToolCallInternalEvent,
AgentBackendError,
AgentBackendHTTPError,
@@ -481,6 +482,10 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
if isinstance(internal_event, AgentBackendStreamInternalEvent):
self._record_stream_metadata(metadata, internal_event)
continue
if internal_event.type == AgentBackendInternalEventType.AGENT_MESSAGE_DELTA:
if isinstance(internal_event, AgentBackendAgentMessageDeltaInternalEvent):
self._record_agent_message_delta_metadata(metadata, internal_event)
continue
metadata["agent_backend"] = {
**dict(metadata.get("agent_backend") or {}),
"stream_event_count": stream_event_count,
@@ -734,6 +739,17 @@ class DifyAgentNode(Node[DifyAgentNodeData]):
agent_backend["usage"] = dict(usage)
metadata["agent_backend"] = agent_backend
@staticmethod
def _record_agent_message_delta_metadata(
metadata: dict[str, Any], event: AgentBackendAgentMessageDeltaInternalEvent
) -> None:
agent_backend = dict(metadata.get("agent_backend") or {})
agent_backend["agent_message_delta_count"] = int(agent_backend.get("agent_message_delta_count") or 0) + 1
agent_backend["agent_message_delta_length"] = int(agent_backend.get("agent_message_delta_length") or 0) + len(
event.delta
)
metadata["agent_backend"] = agent_backend
@classmethod
@override
def _extract_variable_selector_to_variable_mapping(
@@ -10,13 +10,13 @@ trustworthy metadata.
from __future__ import annotations
from mimetypes import guess_extension
from uuid import UUID
from sqlalchemy import select
from sqlalchemy.exc import DataError, SQLAlchemyError
from core.db.session_factory import session_factory
from core.tools.tool_file_manager import resolve_extension
from core.workflow.file_reference import build_file_reference
from graphon.file import File, FileTransferMethod, get_file_type_by_mime_type
from models.tools import ToolFile
@@ -46,7 +46,7 @@ def reback_tool_file_output(*, tenant_id: str, tool_file_id: str) -> File | None
return None
mime_type = tool_file.mimetype or ""
extension = guess_extension(mime_type) or ".bin"
extension = resolve_extension(filename=tool_file.name, mimetype=mime_type)
return File(
type=get_file_type_by_mime_type(mime_type),
transfer_method=FileTransferMethod.TOOL_FILE,
+8 -2
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import mimetypes
import os
import uuid
from collections.abc import Mapping, Sequence
from typing import Any, Literal, NotRequired, TypedDict, assert_never, cast
@@ -285,7 +286,7 @@ def _build_from_remote_url(
raise ValueError("Invalid file url")
mime_type, filename, file_size = get_remote_file_info(url)
extension = mimetypes.guess_extension(mime_type) or ("." + filename.split(".")[-1] if "." in filename else ".bin")
extension = os.path.splitext(filename)[1].lower() or mimetypes.guess_extension(mime_type) or ".bin"
detected_file_type = standardize_file_type(extension=extension, mime_type=mime_type)
file_type = _resolve_file_type(
detected_file_type=detected_file_type,
@@ -326,7 +327,12 @@ def _build_from_tool_file(
if tool_file is None:
raise ValueError(f"ToolFile {tool_file_id} not found")
extension = "." + tool_file.file_key.split(".")[-1] if "." in tool_file.file_key else ".bin"
extension = (
os.path.splitext(tool_file.name)[1].lower()
or mimetypes.guess_extension(tool_file.mimetype)
or os.path.splitext(tool_file.file_key)[1].lower()
or ".bin"
)
detected_file_type = standardize_file_type(extension=extension, mime_type=tool_file.mimetype)
file_type = _resolve_file_type(
detected_file_type=detected_file_type,
+11 -1
View File
@@ -1,9 +1,10 @@
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
from uuid import uuid4
from pydantic import Field, field_validator
from pydantic import Field, computed_field, field_validator
from core.entities.execution_extra_content import ExecutionExtraContentDomainModel
from fields.base import ResponseModel
@@ -55,10 +56,19 @@ class MessageListItem(ResponseModel):
created_at: int | None = None
agent_thoughts: list[AgentThought]
message_files: list[MessageFile]
message_tokens: int = 0
answer_tokens: int = 0
provider_response_latency: float = 0
total_price: Decimal | None = None
currency: str | None = None
status: str
error: str | None = None
extra_contents: list[ExecutionExtraContentDomainModel]
@computed_field
def total_tokens(self) -> int:
return self.message_tokens + self.answer_tokens
@field_validator("inputs", mode="before")
@classmethod
def _normalize_inputs(cls, value: JSONValueType) -> JSONValueType:
+6
View File
@@ -17614,19 +17614,25 @@ Built-in tool icons are URL strings; API-based tool icons are provider-defined p
| ---- | ---- | ----------- | -------- |
| agent_thoughts | [ [AgentThought](#agentthought) ] | | Yes |
| answer | string | | Yes |
| answer_tokens | integer | | No |
| conversation_id | string | | Yes |
| created_at | integer | | No |
| currency | string | | No |
| error | string | | No |
| extra_contents | [ [HumanInputContent](#humaninputcontent) ] | | Yes |
| feedback | [SimpleFeedback](#simplefeedback) | | No |
| id | string | | Yes |
| inputs | object | | Yes |
| message_files | [ [MessageFile](#messagefile) ] | | Yes |
| message_tokens | integer | | No |
| metadata | [JSONValueType](#jsonvaluetype) | | No |
| parent_message_id | string | | No |
| provider_response_latency | number | | No |
| query | string | | Yes |
| retriever_resources | [ [RetrieverResource](#retrieverresource) ] | | Yes |
| status | string | | Yes |
| total_price | string | | No |
| total_tokens | integer | | Yes |
#### ExternalApiTemplateListQuery
+6
View File
@@ -3467,18 +3467,24 @@ Model class for i18n object.
| ---- | ---- | ----------- | -------- |
| agent_thoughts | [ [AgentThought](#agentthought) ] | | Yes |
| answer | string | | Yes |
| answer_tokens | integer | | No |
| conversation_id | string | | Yes |
| created_at | integer | | No |
| currency | string | | No |
| error | string | | No |
| extra_contents | [ [HumanInputContent](#humaninputcontent) ] | | Yes |
| feedback | [SimpleFeedback](#simplefeedback) | | No |
| id | string | | Yes |
| inputs | object | | Yes |
| message_files | [ [MessageFile](#messagefile) ] | | Yes |
| message_tokens | integer | | No |
| parent_message_id | string | | No |
| provider_response_latency | number | | No |
| query | string | | Yes |
| retriever_resources | [ [RetrieverResource](#retrieverresource) ] | | Yes |
| status | string | | Yes |
| total_price | string | | No |
| total_tokens | integer | | Yes |
#### MessageListQuery
+6
View File
@@ -1685,19 +1685,25 @@ in form definiton, or a variable while the workflow is running.
| ---- | ---- | ----------- | -------- |
| agent_thoughts | [ [AgentThought](#agentthought) ] | | Yes |
| answer | string | | Yes |
| answer_tokens | integer | | No |
| conversation_id | string | | Yes |
| created_at | integer | | No |
| currency | string | | No |
| error | string | | No |
| extra_contents | [ [HumanInputContent](#humaninputcontent) ] | | Yes |
| feedback | [SimpleFeedback](#simplefeedback) | | No |
| id | string | | Yes |
| inputs | object | | Yes |
| message_files | [ [MessageFile](#messagefile) ] | | Yes |
| message_tokens | integer | | No |
| metadata | [JSONValueType](#jsonvaluetype) | | No |
| parent_message_id | string | | No |
| provider_response_latency | number | | No |
| query | string | | Yes |
| retriever_resources | [ [RetrieverResource](#retrieverresource) ] | | Yes |
| status | string | | Yes |
| total_price | string | | No |
| total_tokens | integer | | Yes |
#### WebModelConfigResponse
@@ -1,5 +1,6 @@
from __future__ import annotations
import logging
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import cast
@@ -259,14 +260,16 @@ def test_get_project_url_success(trace_instance: AliyunDataTrace):
assert trace_instance.get_project_url() == "project-url"
def test_get_project_url_error(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch):
def test_get_project_url_error(
trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
):
monkeypatch.setattr(trace_instance.trace_client, "get_project_url", MagicMock(side_effect=Exception("boom")))
logger_mock = MagicMock()
monkeypatch.setattr(aliyun_trace_module, "logger", logger_mock)
caplog.set_level(logging.INFO, logger=aliyun_trace_module.logger.name)
with pytest.raises(ValueError, match=r"Aliyun get project url failed: boom"):
trace_instance.get_project_url()
logger_mock.info.assert_called()
assert "Aliyun get project url failed: boom" in caplog.text
def test_workflow_trace_adds_workflow_and_node_spans(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch):
@@ -2,6 +2,7 @@
from __future__ import annotations
import logging
import sys
import types
from types import SimpleNamespace
@@ -87,7 +88,6 @@ class PatchedCoreComponents(TypedDict):
tracer: MagicMock
span: MagicMock
tracer_provider: MagicMock
logger: MagicMock
trace_api: Any
@@ -148,9 +148,6 @@ def patch_core_components(monkeypatch: pytest.MonkeyPatch) -> PatchedCoreCompone
resource = MagicMock(name="resource")
monkeypatch.setattr(client_module, "Resource", MagicMock(return_value=resource))
logger_mock = MagicMock(name="tencent_logger")
monkeypatch.setattr(client_module, "logger", logger_mock)
trace_api_stub = SimpleNamespace(
set_span_in_context=MagicMock(name="set_span_in_context", return_value="trace-context"),
NonRecordingSpan=MagicMock(name="non_recording_span", side_effect=lambda ctx: f"non-{ctx}"),
@@ -174,7 +171,6 @@ def patch_core_components(monkeypatch: pytest.MonkeyPatch) -> PatchedCoreCompone
"tracer": tracer,
"span": span,
"tracer_provider": tracer_provider,
"logger": logger_mock,
"trace_api": trace_api_stub,
}
@@ -268,14 +264,15 @@ def test_record_methods_skip_when_histogram_missing() -> None:
client.record_trace_duration(0.5)
def test_record_llm_duration_handles_exceptions(patch_core_components: PatchedCoreComponents) -> None:
def test_record_llm_duration_handles_exceptions(caplog: pytest.LogCaptureFixture) -> None:
client = _build_client()
client.hist_llm_duration = MagicMock(name="hist_llm_duration")
client.hist_llm_duration.record.side_effect = RuntimeError("boom")
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client.record_llm_duration(0.2)
logger = patch_core_components["logger"]
logger.debug.assert_called()
assert "[Tencent APM] Failed to record LLM duration" in caplog.text
def test_create_and_export_span_sets_attributes(patch_core_components: PatchedCoreComponents) -> None:
@@ -328,12 +325,15 @@ def test_create_and_export_span_uses_parent_context(patch_core_components: Patch
trace_api.set_span_in_context.assert_called_once()
def test_create_and_export_span_exception_logs_error(patch_core_components: PatchedCoreComponents) -> None:
def test_create_and_export_span_exception_logs_error(
patch_core_components: PatchedCoreComponents, caplog: pytest.LogCaptureFixture
) -> None:
client = _build_client()
span = patch_core_components["span"]
span.get_span_context.return_value = _make_span_context(span_id=2)
client.tracer.start_span.side_effect = RuntimeError("boom")
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client._create_and_export_span(
SpanData(
trace_id=1,
@@ -346,8 +346,10 @@ def test_create_and_export_span_exception_logs_error(patch_core_components: Patc
end_time=1,
)
)
logger = patch_core_components["logger"]
logger.exception.assert_called_once()
error_records = [record for record in caplog.records if record.levelno == logging.ERROR]
assert len(error_records) == 1
assert error_records[0].getMessage() == "[Tencent APM] Error creating span: span"
def test_api_check_connects_successfully(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -423,23 +425,18 @@ def test_shutdown_flushes_all_components(patch_core_components: PatchedCoreCompo
metric_reader.shutdown.assert_called_once()
def test_shutdown_logs_when_meter_provider_fails(patch_core_components: PatchedCoreComponents) -> None:
def test_shutdown_logs_when_meter_provider_fails(caplog: pytest.LogCaptureFixture) -> None:
client = _build_client()
meter_provider = meter_provider_instances[-1]
meter_provider.shutdown.side_effect = RuntimeError("boom")
assert client.metric_reader is not None
client.metric_reader.shutdown.side_effect = RuntimeError("boom")
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client.shutdown()
logger = patch_core_components["logger"]
logger.debug.assert_any_call(
"[Tencent APM] Error shutting down meter provider",
exc_info=True,
)
logger.debug.assert_any_call(
"[Tencent APM] Error shutting down metric reader",
exc_info=True,
)
assert "[Tencent APM] Error shutting down meter provider" in caplog.text
assert "[Tencent APM] Error shutting down metric reader" in caplog.text
def test_metrics_initialization_failure_sets_histogram_attributes(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -456,10 +453,11 @@ def test_metrics_initialization_failure_sets_histogram_attributes(monkeypatch: p
assert client.metric_reader is None
def test_add_span_logs_exception(monkeypatch: pytest.MonkeyPatch, patch_core_components: PatchedCoreComponents) -> None:
def test_add_span_logs_exception(monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture) -> None:
client = _build_client()
monkeypatch.setattr(client, "_create_and_export_span", MagicMock(side_effect=RuntimeError("boom")))
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
client.add_span(
SpanData(
trace_id=1,
@@ -473,8 +471,9 @@ def test_add_span_logs_exception(monkeypatch: pytest.MonkeyPatch, patch_core_com
)
)
logger = patch_core_components["logger"]
logger.exception.assert_called_once()
error_records = [record for record in caplog.records if record.levelno == logging.ERROR]
assert len(error_records) == 1
assert error_records[0].getMessage() == "[Tencent APM] Failed to create span: span"
def test_create_and_export_span_converts_attribute_types(patch_core_components: PatchedCoreComponents) -> None:
@@ -535,16 +534,20 @@ def test_record_trace_duration_converts_attributes() -> None:
],
)
def test_record_methods_handle_exceptions(
method: str, attr_name: str, args: tuple[object, ...], patch_core_components: PatchedCoreComponents
method: str, attr_name: str, args: tuple[object, ...], caplog: pytest.LogCaptureFixture
) -> None:
client = _build_client()
hist_mock = MagicMock(name=attr_name)
hist_mock.record.side_effect = RuntimeError("boom")
setattr(client, attr_name, hist_mock)
caplog.set_level(logging.DEBUG, logger=client_module.logger.name)
getattr(client, method)(*args)
logger = patch_core_components["logger"]
logger.debug.assert_called()
assert any(
record.levelno == logging.DEBUG and record.getMessage().startswith("[Tencent APM] Failed to record")
for record in caplog.records
)
def test_metrics_initializes_grpc_metric_exporter() -> None:
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-api"
version = "1.15.0"
version = "1.16.0-rc1"
requires-python = "~=3.12.0"
dependencies = [
+2 -1
View File
@@ -45,6 +45,7 @@ class FileRequestService:
user_from: UserFrom | str,
invoke_from: InvokeFrom | str,
file_mapping: Mapping[str, Any],
for_external: bool = True,
) -> DownloadFileRequestResult:
"""Resolve one file mapping into signed download metadata.
@@ -61,7 +62,7 @@ class FileRequestService:
)
with bind_file_access_scope(scope):
file = self._build_file(mapping=file_mapping, tenant_id=tenant_id)
download_url = file_helpers.resolve_file_url(file, for_external=True)
download_url = file_helpers.resolve_file_url(file, for_external=for_external)
if not download_url:
raise ValueError("file does not support signed download")
@@ -29,12 +29,12 @@ import hashlib
import json
import logging
import time
from collections.abc import Sequence
from collections.abc import Callable, Sequence
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from enum import Enum
from threading import Lock
from typing import Any, NotRequired, TypedDict, cast
from typing import Any, NotRequired, TypedDict, TypeVar, cast
import click
import pyarrow as pa
@@ -70,8 +70,10 @@ from services.retention.workflow_run.constants import (
ARCHIVE_BUNDLE_MANIFEST_NAME,
ARCHIVE_BUNDLE_SCHEMA_VERSION,
)
from services.retention.workflow_run.db_retry import is_retryable_db_disconnect, run_with_db_retry
logger = logging.getLogger(__name__)
T = TypeVar("T")
class TableStatsManifestEntry(TypedDict):
@@ -208,6 +210,8 @@ class WorkflowRunArchiver:
"workflow_pause_reasons",
"workflow_trigger_logs",
]
DB_RETRY_ATTEMPTS = 3
DB_RETRY_DELAYS_SECONDS = (1.0, 2.0)
start_from: datetime.datetime | None
end_before: datetime.datetime
@@ -455,16 +459,40 @@ class WorkflowRunArchiver:
"""Fetch a batch of workflow runs to archive."""
repo = self._get_workflow_run_repo()
tenant_ids = list(tenant_scope) if tenant_scope is not None else self.tenant_ids or None
return repo.get_runs_batch_by_time_range(
start_from=self.start_from,
end_before=self.end_before,
last_seen=last_seen,
batch_size=self.batch_size,
run_types=self.ARCHIVED_TYPE,
tenant_ids=tenant_ids,
tenant_prefixes=None if tenant_ids else self.tenant_prefixes or None,
run_shard_index=self.run_shard_index,
run_shard_total=self.run_shard_total,
return self._run_with_db_retry(
"workflow run batch fetch",
lambda: repo.get_runs_batch_by_time_range(
start_from=self.start_from,
end_before=self.end_before,
last_seen=last_seen,
batch_size=self.batch_size,
run_types=self.ARCHIVED_TYPE,
tenant_ids=tenant_ids,
tenant_prefixes=None if tenant_ids else self.tenant_prefixes or None,
run_shard_index=self.run_shard_index,
run_shard_total=self.run_shard_total,
),
)
@staticmethod
def _is_retryable_db_disconnect(exc: BaseException) -> bool:
return is_retryable_db_disconnect(exc)
@staticmethod
def _safe_rollback(session: Session, bundle_id: str) -> None:
try:
session.rollback()
except Exception:
logger.warning("Failed to rollback archive session for bundle %s", bundle_id, exc_info=True)
def _run_with_db_retry(self, operation_name: str, operation: Callable[[], T]) -> T:
return run_with_db_retry(
operation_name,
operation,
logger=logger,
attempts=self.DB_RETRY_ATTEMPTS,
delays_seconds=self.DB_RETRY_DELAYS_SECONDS,
)
def _tenant_scan_scopes(self) -> list[list[str] | None]:
@@ -532,16 +560,14 @@ class WorkflowRunArchiver:
if self.workers == 1 or len(bundle_groups) == 1:
results: list[ArchiveResult] = []
for bundle_runs in bundle_groups:
with session_maker() as session:
results.append(self._archive_bundle(session, storage, bundle_runs))
results.append(self._archive_bundle_with_retry(session_maker, storage, bundle_runs))
return results
results = []
max_workers = min(self.workers, len(bundle_groups))
def archive_in_worker(bundle_runs: Sequence[WorkflowRun]) -> ArchiveResult:
with session_maker() as session:
return self._archive_bundle(session, storage, bundle_runs)
return self._archive_bundle_with_retry(session_maker, storage, bundle_runs)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
futures = [executor.submit(archive_in_worker, bundle_runs) for bundle_runs in bundle_groups]
@@ -549,6 +575,39 @@ class WorkflowRunArchiver:
results.append(future.result())
return results
def _archive_bundle_with_retry(
self,
session_maker: sessionmaker[Session],
storage: ArchiveStorage | None,
runs: Sequence[WorkflowRun],
) -> ArchiveResult:
identity = self._build_bundle_identity(runs)
try:
return self._run_with_db_retry(
f"archive workflow run bundle {identity.bundle_id}",
lambda: self._archive_bundle_once(session_maker, storage, runs),
)
except Exception as exc:
logger.exception("Failed to archive workflow run bundle %s after retries", identity.bundle_id)
return ArchiveResult(
bundle_id=identity.bundle_id,
tenant_id=identity.tenant_id,
object_prefix=identity.object_prefix,
run_count=len(runs),
success=False,
error=str(exc),
)
def _archive_bundle_once(
self,
session_maker: sessionmaker[Session],
storage: ArchiveStorage | None,
runs: Sequence[WorkflowRun],
) -> ArchiveResult:
with session_maker() as session:
return self._archive_bundle(session, storage, runs)
def _archive_bundle(
self,
session: Session,
@@ -645,9 +704,12 @@ class WorkflowRunArchiver:
result.success = True
except Exception as e:
if self._is_retryable_db_disconnect(e):
self._safe_rollback(session, identity.bundle_id)
raise
logger.exception("Failed to archive workflow run bundle %s", identity.bundle_id)
result.error = str(e)
session.rollback()
self._safe_rollback(session, identity.bundle_id)
result.elapsed_time = time.time() - start_time
return result
@@ -794,6 +856,9 @@ class WorkflowRunArchiver:
end_before = self.end_before
if end_before is None:
raise ValueError("archive window end must be set")
formatted_end_before = self._format_window_datetime(end_before)
if formatted_end_before is None:
raise ValueError("archive window end must be set")
return ArchiveManifestDict(
schema_version=ARCHIVE_BUNDLE_SCHEMA_VERSION,
archive_format=ARCHIVE_BUNDLE_FORMAT,
@@ -813,7 +878,7 @@ class WorkflowRunArchiver:
archived_at=datetime.datetime.now(datetime.UTC).isoformat(),
campaign_id=self.campaign_id,
archive_window_start=self._format_window_datetime(self.start_from),
archive_window_end=end_before.isoformat(),
archive_window_end=formatted_end_before,
run_shard=identity.shard,
tables=tables,
run_ids=[run.id for run in sorted_runs],
@@ -0,0 +1,66 @@
import logging
import time
from collections.abc import Callable
from sqlalchemy.exc import DBAPIError
from sqlalchemy.exc import OperationalError as SQLAlchemyOperationalError
DEFAULT_DB_RETRY_ATTEMPTS = 3
DEFAULT_DB_RETRY_DELAYS_SECONDS = (1.0, 2.0)
_DB_DISCONNECT_PATTERNS = (
"server closed the connection unexpectedly",
"connection already closed",
"closed the connection",
"connection not open",
"terminating connection",
"connection reset",
"broken pipe",
"connection invalidated",
)
def is_retryable_db_disconnect(exc: BaseException) -> bool:
if isinstance(exc, DBAPIError) and exc.connection_invalidated:
return True
if not _is_db_operational_error(exc):
return False
original_exception = exc.orig if isinstance(exc, DBAPIError) else None
message = f"{exc} {original_exception or ''}".lower()
return any(pattern in message for pattern in _DB_DISCONNECT_PATTERNS)
def run_with_db_retry[T](
operation_name: str,
operation: Callable[[], T],
*,
logger: logging.Logger,
attempts: int = DEFAULT_DB_RETRY_ATTEMPTS,
delays_seconds: tuple[float, ...] = DEFAULT_DB_RETRY_DELAYS_SECONDS,
) -> T:
for attempt in range(1, attempts + 1):
try:
return operation()
except Exception as exc:
if not is_retryable_db_disconnect(exc) or attempt == attempts:
raise
delay = delays_seconds[min(attempt - 1, len(delays_seconds) - 1)]
logger.warning(
"Retrying %s after retryable DB disconnect (attempt %s/%s, sleep %.1fs)",
operation_name,
attempt,
attempts,
delay,
exc_info=True,
)
time.sleep(delay)
raise RuntimeError(f"{operation_name} did not complete")
def _is_db_operational_error(exc: BaseException) -> bool:
if isinstance(exc, SQLAlchemyOperationalError):
return True
return exc.__class__.__name__ == "OperationalError" and exc.__class__.__module__.startswith(("psycopg", "psycopg2"))
+6 -7
View File
@@ -6,7 +6,7 @@ import uuid
from datetime import UTC, datetime
from typing import TypedDict, cast
from sqlalchemy import select
from sqlalchemy import select, update
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
@@ -912,12 +912,11 @@ class SummaryIndexService:
# Disable summary records (don't delete)
now = naive_utc_now()
for summary in summaries:
summary.enabled = False
summary.disabled_at = now
summary.disabled_by = disabled_by
session.add(summary)
session.execute(
update(DocumentSegmentSummary)
.where(DocumentSegmentSummary.id.in_(s.id for s in summaries))
.values(enabled=False, disabled_at=now, disabled_by=disabled_by)
)
session.commit()
logger.info("Disabled %s summary records for dataset %s", len(summaries), dataset.id)
@@ -6,8 +6,10 @@ from unittest.mock import ANY, MagicMock, patch
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
from sqlalchemy.exc import OperationalError
from services.retention.workflow_run.archive_paid_plan_workflow_run import (
ArchiveResult,
ArchiveSummary,
WorkflowRunArchiver,
)
@@ -32,6 +34,30 @@ class FakeArchiveStorage:
return sorted(key for key in self.objects if key.startswith(prefix))
def _db_disconnect_error() -> OperationalError:
return OperationalError(
"select 1",
{},
RuntimeError("server closed the connection unexpectedly"),
connection_invalidated=True,
)
def _run(run_id: str = "run-1"):
run = MagicMock()
run.id = run_id
run.tenant_id = "tenant-1"
run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0)
return run
def _session_context(session):
context = MagicMock()
context.__enter__.return_value = session
context.__exit__.return_value = False
return context
class TestWorkflowRunArchiverInit:
def test_start_from_without_end_before_raises(self):
with pytest.raises(ValueError, match="start_from and end_before must be provided together"):
@@ -139,6 +165,32 @@ class TestWorkflowRunArchiverInit:
repo.get_runs_batch_by_time_range.assert_called_once()
assert repo.get_runs_batch_by_time_range.call_args.kwargs["tenant_ids"] == ["tenant-b"]
def test_get_runs_batch_retries_retryable_db_disconnect(self):
repo = MagicMock()
repo.get_runs_batch_by_time_range.side_effect = [_db_disconnect_error(), []]
archiver = WorkflowRunArchiver(workflow_run_repo=repo)
with patch("services.retention.workflow_run.db_retry.time.sleep") as sleep:
runs = archiver._get_runs_batch(None)
assert runs == []
assert repo.get_runs_batch_by_time_range.call_count == 2
sleep.assert_called_once_with(1.0)
def test_get_runs_batch_does_not_retry_non_db_broken_pipe_error(self):
repo = MagicMock()
repo.get_runs_batch_by_time_range.side_effect = RuntimeError("broken pipe")
archiver = WorkflowRunArchiver(workflow_run_repo=repo)
with (
patch("services.retention.workflow_run.db_retry.time.sleep") as sleep,
pytest.raises(RuntimeError, match="broken pipe"),
):
archiver._get_runs_batch(None)
repo.get_runs_batch_by_time_range.assert_called_once()
sleep.assert_not_called()
def test_start_message_includes_shard(self):
archiver = WorkflowRunArchiver(tenant_prefixes=["0"], run_shard_index=1, run_shard_total=4)
@@ -351,6 +403,72 @@ class TestDryRunArchive:
assert summary.table_stats["workflow_app_logs"].size_bytes == 32
class TestArchiveDbRetry:
def test_archive_bundle_groups_retries_with_fresh_session(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
session_maker = MagicMock(
side_effect=[
_session_context(MagicMock(name="session-1")),
_session_context(MagicMock(name="session-2")),
]
)
success = ArchiveResult(
bundle_id=archiver._build_bundle_identity([run]).bundle_id,
tenant_id=run.tenant_id,
object_prefix=archiver._build_bundle_identity([run]).object_prefix,
run_count=1,
success=True,
)
with (
patch.object(archiver, "_archive_bundle", side_effect=[_db_disconnect_error(), success]) as archive_bundle,
patch("services.retention.workflow_run.db_retry.time.sleep") as sleep,
):
results = archiver._archive_bundle_groups(session_maker, MagicMock(), [[run]])
assert results == [success]
assert archive_bundle.call_count == 2
assert session_maker.call_count == 2
sleep.assert_called_once_with(1.0)
def test_archive_bundle_groups_returns_failed_result_after_retry_exhaustion(self):
archiver = WorkflowRunArchiver(days=90)
run = _run()
session_maker = MagicMock(
side_effect=[
_session_context(MagicMock(name="session-1")),
_session_context(MagicMock(name="session-2")),
_session_context(MagicMock(name="session-3")),
]
)
with (
patch.object(archiver, "_archive_bundle", side_effect=[_db_disconnect_error()] * 3) as archive_bundle,
patch("services.retention.workflow_run.db_retry.time.sleep") as sleep,
):
results = archiver._archive_bundle_groups(session_maker, MagicMock(), [[run]])
assert len(results) == 1
assert results[0].success is False
assert "server closed the connection unexpectedly" in (results[0].error or "")
assert archive_bundle.call_count == archiver.DB_RETRY_ATTEMPTS
assert session_maker.call_count == archiver.DB_RETRY_ATTEMPTS
assert sleep.call_count == archiver.DB_RETRY_ATTEMPTS - 1
def test_archive_bundle_uses_safe_rollback_when_failure_rolls_back_badly(self):
archiver = WorkflowRunArchiver(days=90, dry_run=True)
session = MagicMock()
session.rollback.side_effect = RuntimeError("rollback failed")
with patch.object(archiver, "_extract_bundle_data", side_effect=RuntimeError("extract failed")):
result = archiver._archive_bundle(session, None, [_run()])
assert result.success is False
assert result.error == "extract failed"
session.rollback.assert_called_once()
class TestArchiveRunIdempotency:
def _index_payload(self, archiver: WorkflowRunArchiver, run_ids: list[str], run) -> tuple[str, bytes]:
identity = archiver._build_bundle_identity([run])
@@ -0,0 +1,139 @@
import datetime
from unittest.mock import MagicMock
import click
import pytest
from sqlalchemy.exc import OperationalError
from commands import retention
def _db_disconnect_error() -> OperationalError:
return OperationalError(
"select 1",
{},
RuntimeError("server closed the connection unexpectedly"),
connection_invalidated=True,
)
def _session_context(session):
context = MagicMock()
context.__enter__.return_value = session
context.__exit__.return_value = False
return context
def test_resolve_archive_tenant_ids_from_plan_uses_explicit_sessions(monkeypatch):
end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC)
sessions = [MagicMock(name="session-a"), MagicMock(name="session-b")]
session_maker = MagicMock(side_effect=[_session_context(sessions[0]), _session_context(sessions[1])])
calls = []
def get_candidate_tenants(session, prefix, *, start_from, end_before):
calls.append((session, prefix, start_from, end_before))
return [f"{prefix}-paid", f"{prefix}-free"]
monkeypatch.setattr(retention, "_get_archive_candidate_tenant_ids_by_prefix", get_candidate_tenants)
monkeypatch.setattr(
retention,
"_filter_paid_workflow_archive_tenant_ids",
lambda tenant_ids: (["a-paid", "b-paid"], ["a-free", "b-free"]),
)
tenant_plan = retention._resolve_archive_tenant_ids_from_plan(
session_maker=session_maker,
tenant_ids=None,
tenant_prefixes=["a", "b"],
start_from=None,
end_before=end_before,
)
assert tenant_plan["archive_tenant_ids"] == ["a-paid", "b-paid"]
assert tenant_plan["paid_tenant_ids"] == ["a-paid", "b-paid"]
assert tenant_plan["unpaid_tenant_ids"] == ["a-free", "b-free"]
assert calls == [
(sessions[0], "a", None, end_before),
(sessions[1], "b", None, end_before),
]
def test_safe_remove_scoped_session_discards_registry_and_disposes_after_remove_error(monkeypatch):
fake_db = MagicMock()
fake_db.session.remove.side_effect = RuntimeError("server closed the connection unexpectedly")
monkeypatch.setattr(retention, "db", fake_db)
retention._safe_remove_scoped_session("archive workflow run command")
fake_db.session.remove.assert_called_once()
fake_db.session.registry.clear.assert_called_once()
fake_db.engine.dispose.assert_called_once()
def test_archive_command_db_retry_retries_retryable_db_disconnect(monkeypatch):
operation = MagicMock(side_effect=[_db_disconnect_error(), "ok"])
sleep = MagicMock()
monkeypatch.setattr("services.retention.workflow_run.db_retry.time.sleep", sleep)
result = retention._run_archive_command_db_retry("archive plan", operation)
assert result == "ok"
assert operation.call_count == 2
sleep.assert_called_once_with(1.0)
def test_archive_plan_prefix_stats_retries_count_query_with_fresh_session(monkeypatch):
end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC)
sessions = [MagicMock(name="session-1"), MagicMock(name="session-2")]
sessions[0].scalar.side_effect = _db_disconnect_error()
sessions[1].scalar.side_effect = [7, 9]
session_maker = MagicMock(side_effect=[_session_context(sessions[0]), _session_context(sessions[1])])
sleep = MagicMock()
monkeypatch.setattr(
retention,
"_get_archive_candidate_tenant_ids_by_prefix",
lambda session, prefix, *, start_from, end_before: [f"{prefix}-tenant"],
)
monkeypatch.setattr("services.retention.workflow_run.db_retry.time.sleep", sleep)
stats = retention._get_archive_plan_prefix_stats(
session_maker,
"a",
start_from=None,
end_before=end_before,
)
assert stats["tenant_ids"] == ["a-tenant"]
assert stats["workflow_runs"] == 7
assert stats["workflow_node_executions"] == 9
assert session_maker.call_count == 2
sleep.assert_called_once_with(1.0)
def test_archive_workflow_runs_raises_click_exception_when_tenant_plan_fails(monkeypatch):
fake_db = MagicMock()
monkeypatch.setattr(retention, "db", fake_db)
monkeypatch.setattr(
retention,
"_resolve_archive_tenant_ids_from_plan",
MagicMock(side_effect=RuntimeError("tenant plan failed")),
)
with pytest.raises(click.ClickException, match="Failed to resolve workflow archive tenant plan"):
retention.archive_workflow_runs.callback(
tenant_ids="tenant-1",
tenant_prefixes=None,
before_days=90,
from_days_ago=None,
to_days_ago=None,
start_from=None,
end_before=None,
batch_size=10000,
workers=1,
run_shard_index=None,
run_shard_total=None,
limit=None,
dry_run=True,
delete_after_archive=False,
)
@@ -75,6 +75,7 @@ def test_dify_config(monkeypatch: pytest.MonkeyPatch):
# default values
assert config.EDITION == "SELF_HOSTED"
assert config.API_COMPRESSION_ENABLED is False
assert config.AGENT_SHELL_ENABLED is True
assert config.SENTRY_TRACES_SAMPLE_RATE == 1.0
assert config.TEMPLATE_TRANSFORM_MAX_LENGTH == 400_000
@@ -110,6 +111,25 @@ def test_http_timeout_defaults(monkeypatch: pytest.MonkeyPatch):
assert config.HTTP_REQUEST_MAX_WRITE_TIMEOUT == 600
def test_internal_files_url_falls_back_to_server_console_api_url(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
monkeypatch.setenv("SERVER_CONSOLE_API_URL", "http://api:5001")
config = DifyConfig(_env_file=None)
assert config.INTERNAL_FILES_URL == "http://api:5001"
def test_internal_files_url_prefers_explicit_value(monkeypatch: pytest.MonkeyPatch):
os.environ.clear()
monkeypatch.setenv("INTERNAL_FILES_URL", "http://files-internal:5001")
monkeypatch.setenv("SERVER_CONSOLE_API_URL", "http://api:5001")
config = DifyConfig(_env_file=None)
assert config.INTERNAL_FILES_URL == "http://files-internal:5001"
# NOTE: If there is a `.env` file in your Workspace, this test might not succeed as expected.
# This is due to `pymilvus` loading all the variables from the `.env` file into `os.environ`.
def test_flask_configs(monkeypatch: pytest.MonkeyPatch):
@@ -49,6 +49,8 @@ def make_message():
msg.query = "hello"
msg.re_sign_file_url_answer = ""
msg.user_feedback = MagicMock(rating=None)
msg.total_price = None
msg.currency = None
msg.status = "normal"
msg.error = None
return msg
@@ -34,11 +34,11 @@ class DummyFile:
class DummyToolFile:
def __init__(self):
def __init__(self, name="test.txt", mimetype="text/plain"):
self.id = "file-id"
self.name = "test.txt"
self.name = name
self.size = 10
self.mimetype = "text/plain"
self.mimetype = mimetype
self.original_url = "http://original"
self.user_id = "user-1"
self.tenant_id = "tenant-1"
@@ -56,7 +56,7 @@ class TestPluginUploadFileApi:
mock_get_user,
mock_verify_signature,
):
dummy_file = DummyFile()
dummy_file = DummyFile(filename="report.docx", mimetype="application/octet-stream")
module.request = fake_request(
{
@@ -71,7 +71,10 @@ class TestPluginUploadFileApi:
)
tool_file_manager_instance = mock_tool_file_manager.return_value
tool_file_manager_instance.create_file_by_raw.return_value = DummyToolFile()
tool_file_manager_instance.create_file_by_raw.return_value = DummyToolFile(
name="report.docx",
mimetype="application/octet-stream",
)
mock_tool_file_manager.sign_file.return_value = "signed-url"
@@ -84,10 +87,12 @@ class TestPluginUploadFileApi:
assert result["id"] == "file-id"
assert result["reference"] == build_file_reference(record_id="file-id")
assert result["preview_url"] == "signed-url"
assert result["extension"] == ".docx"
mock_verify_signature.assert_called_once()
assert mock_verify_signature.call_args.kwargs["conversation_id"] == "conversation-1"
tool_file_manager_instance.create_file_by_raw.assert_called_once()
assert tool_file_manager_instance.create_file_by_raw.call_args.kwargs["conversation_id"] == "conversation-1"
mock_tool_file_manager.sign_file.assert_called_once_with(tool_file_id="file-id", extension=".docx")
def test_missing_file(self):
module.request = fake_request(
@@ -318,6 +318,7 @@ class TestPluginDownloadFileRequestApi:
mock_payload.user_id = "user-id"
mock_payload.user_from = "account"
mock_payload.invoke_from = "debugger"
mock_payload.for_external = False
reference = build_file_reference(record_id="tool-file-1")
mock_payload.file.model_dump.return_value = {
"transfer_method": "tool_file",
@@ -333,6 +334,7 @@ class TestPluginDownloadFileRequestApi:
user_from="account",
invoke_from="debugger",
file_mapping={"transfer_method": "tool_file", "reference": reference},
for_external=False,
)
assert result["data"] == {
"filename": "report.pdf",
@@ -37,6 +37,7 @@ from pydantic_ai.messages import (
from clients.agent_backend import (
AgentBackendError,
AgentBackendRunEventAdapter,
AgentBackendRunFailedInternalEvent,
AgentBackendStreamInternalEvent,
FakeAgentBackendRunClient,
FakeAgentBackendScenario,
@@ -54,6 +55,7 @@ from core.app.entities.queue_entities import (
QueueMessageEndEvent,
)
from core.workflow.nodes.agent_v2.ask_human_resume import AskHumanResumeOutcome
from graphon.model_runtime.errors.invoke import InvokeRateLimitError
from models.agent_config_entities import AgentSoulConfig
from models.model import MessageAgentThought
@@ -1039,6 +1041,130 @@ def test_tool_result_without_call_id_matches_unique_open_tool_name(monkeypatch):
assert rows[0].observation == "Knowledge base search results: browser skill"
def test_repeated_tool_calls_without_call_id_or_index_create_distinct_rows(monkeypatch):
fake_session = _FakeDbSession()
monkeypatch.setattr(app_runner_module.db, "session", fake_session)
qm = _FakeQueueManager()
recorder = app_runner_module._AgentProcessRecorder(
dify_context=_dify_ctx(),
message_id="msg-1",
queue_manager=qm, # type: ignore[arg-type]
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_call",
"part": {
"part_kind": "tool-call",
"tool_name": "shell_run",
"args": {"script": "lookup find"},
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_result",
"part": {
"part_kind": "tool-return",
"tool_name": "shell_run",
"content": "find output",
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_call",
"part": {
"part_kind": "tool-call",
"tool_name": "shell_run",
"args": {"script": "lookup out"},
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_result",
"part": {
"part_kind": "tool-return",
"tool_name": "shell_run",
"content": "out output",
},
},
)
)
rows = sorted(fake_session.rows.values(), key=lambda row: row.position)
assert len(rows) == 2
assert rows[0].tool == "shell_run"
assert rows[0].tool_input == '{"script": "lookup find"}'
assert rows[0].observation == "find output"
assert rows[1].tool == "shell_run"
assert rows[1].tool_input == '{"script": "lookup out"}'
assert rows[1].observation == "out output"
def test_repeated_tool_calls_with_placeholder_call_id_and_reused_index_create_distinct_rows(monkeypatch):
fake_session = _FakeDbSession()
monkeypatch.setattr(app_runner_module.db, "session", fake_session)
qm = _FakeQueueManager()
recorder = app_runner_module._AgentProcessRecorder(
dify_context=_dify_ctx(),
message_id="msg-1",
queue_manager=qm, # type: ignore[arg-type]
)
for script, output in (("lookup find", "find output"), ("lookup out", "out output")):
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_call",
"index": 0,
"part": {
"part_kind": "tool-call",
"tool_name": "shell_run",
"tool_call_id": "None",
"args": {"script": script},
},
},
)
)
recorder.handle_stream_event(
AgentBackendStreamInternalEvent(
run_id="run-1",
data={
"event_kind": "function_tool_result",
"part": {
"part_kind": "tool-return",
"tool_name": "shell_run",
"tool_call_id": "None",
"content": output,
},
},
)
)
rows = sorted(fake_session.rows.values(), key=lambda row: row.position)
assert len(rows) == 2
assert rows[0].tool == "shell_run"
assert rows[0].tool_input == '{"script": "lookup find"}'
assert rows[0].observation == "find output"
assert rows[1].tool == "shell_run"
assert rows[1].tool_input == '{"script": "lookup out"}'
assert rows[1].observation == "out output"
def test_prior_session_snapshot_is_threaded_into_request():
prior = CompositorSessionSnapshot(layers=[])
client = FakeAgentBackendRunClient()
@@ -1088,6 +1214,19 @@ def test_failed_run_raises_agent_backend_error():
assert store.saved == []
def test_agent_backend_failure_to_exception_maps_rate_limit_reason():
err = app_runner_module._agent_backend_failure_to_exception(
AgentBackendRunFailedInternalEvent(
run_id="run-1",
error="quota exceeded",
reason="InvokeRateLimitError",
)
)
assert isinstance(err, InvokeRateLimitError)
assert str(err) == "quota exceeded"
def test_stopped_task_cancels_agent_backend_run_and_skips_session_save():
client = _RecordingFakeAgentBackendRunClient()
store = _FakeSessionStore()
@@ -116,10 +116,12 @@ class TestAgentChatAppGeneratorGenerate:
)
thread_obj = mocker.MagicMock()
mocker.patch(
thread_constructor = mocker.patch(
"core.app.apps.agent_chat.app_generator.threading.Thread",
return_value=thread_obj,
)
session = mocker.MagicMock()
mocker.patch("core.app.apps.agent_chat.app_generator.db.session", return_value=session)
mocker.patch(
"core.app.apps.agent_chat.app_generator.AgentChatAppGenerateResponseConverter.convert",
@@ -144,6 +146,7 @@ class TestAgentChatAppGeneratorGenerate:
assert result == {"result": "ok"}
assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1"
assert thread_constructor.call_args.kwargs["kwargs"]["session"] is session
thread_obj.start.assert_called_once()
def test_generate_without_file_config(self, generator, mocker: MockerFixture):
@@ -3,10 +3,11 @@ from unittest.mock import Mock
import pytest
from core.app.apps.base_app_generate_response_converter import AppGenerateResponseConverter
from core.app.entities.queue_entities import QueueErrorEvent
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
from core.errors.error import QuotaExceededError
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError, InvokeRateLimitError
from models.enums import MessageStatus
@@ -68,6 +69,11 @@ class TestBasedGenerateTaskPipeline:
assert error_response.task_id == "task-1"
assert ping_response.task_id == "task-1"
def test_stream_converter_maps_invoke_rate_limit_error(self):
data = AppGenerateResponseConverter._error_to_stream_response(InvokeRateLimitError("quota exceeded"))
assert data == {"code": "rate_limit_error", "status": 429, "message": "quota exceeded"}
def test_handle_output_moderation_when_flagged(self, pipeline):
handler = Mock()
handler.moderation_completion.return_value = ("filtered", True)
@@ -0,0 +1,32 @@
import json
from pytest_mock import MockerFixture
from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType
def test_model_provider_credentials_cache_get_returns_decoded_dict(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.model_provider_cache.redis_client")
cache = ProviderCredentialsCache(
tenant_id="tenant",
identity_id="identity",
cache_type=ProviderCredentialsCacheType.PROVIDER,
)
payload = {"api_key": "secret"}
redis_client_mock.get.return_value = json.dumps(payload).encode("utf-8")
assert cache.get() == payload
def test_model_provider_credentials_cache_get_returns_none_for_invalid_utf8(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.model_provider_cache.redis_client")
cache = ProviderCredentialsCache(
tenant_id="tenant",
identity_id="identity",
cache_type=ProviderCredentialsCacheType.PROVIDER,
)
redis_client_mock.get.return_value = b"\xff"
assert cache.get() is None
@@ -0,0 +1,24 @@
import json
from pytest_mock import MockerFixture
from core.helper.provider_cache import ToolProviderCredentialsCache
def test_provider_credentials_cache_get_returns_decoded_dict(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.provider_cache.redis_client")
cache = ToolProviderCredentialsCache(tenant_id="tenant", provider="provider", credential_id="credential")
payload = {"api_key": "secret"}
redis_client_mock.get.return_value = json.dumps(payload).encode("utf-8")
assert cache.get() == payload
def test_provider_credentials_cache_get_returns_none_for_invalid_utf8(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.provider_cache.redis_client")
cache = ToolProviderCredentialsCache(tenant_id="tenant", provider="provider", credential_id="credential")
redis_client_mock.get.return_value = b"\xff"
assert cache.get() is None
@@ -38,6 +38,21 @@ def test_tool_parameter_cache_get_returns_none_for_invalid_json(mocker: MockerFi
assert cache.get() is None
def test_tool_parameter_cache_get_returns_none_for_invalid_utf8(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.tool_parameter_cache.redis_client")
cache = ToolParameterCache(
tenant_id="tenant",
provider="provider",
tool_name="tool",
cache_type=ToolParameterCacheType.PARAMETER,
identity_id="identity",
)
redis_client_mock.get.return_value = b"\xff"
assert cache.get() is None
def test_tool_parameter_cache_get_returns_none_when_key_is_missing(mocker: MockerFixture) -> None:
redis_client_mock = mocker.patch("core.helper.tool_parameter_cache.redis_client")
cache = ToolParameterCache(
@@ -22,6 +22,7 @@ def test_request_download_file_accepts_tool_file_reference() -> None:
assert payload.file.transfer_method == "tool_file"
assert payload.file.reference == reference
assert payload.for_external is True
def test_request_download_file_accepts_remote_url() -> None:
@@ -42,6 +43,25 @@ def test_request_download_file_accepts_remote_url() -> None:
assert payload.file.url == "https://example.com/report.pdf"
def test_request_download_file_accepts_internal_download_request() -> None:
reference = build_file_reference(record_id="tool-file-1")
payload = RequestRequestDownloadFile.model_validate(
{
"tenant_id": "tenant-1",
"user_id": "user-1",
"user_from": "account",
"invoke_from": "debugger",
"file": {
"transfer_method": "tool_file",
"reference": reference,
},
"for_external": False,
}
)
assert payload.for_external is False
def test_request_download_file_rejects_remote_url_without_url() -> None:
with pytest.raises(ValidationError, match="url is required"):
_ = RequestRequestDownloadFile.model_validate(
@@ -85,6 +85,10 @@ def test_assembling_request_auth_header_assembly():
assert headers["Authorization"] == "Bearer abc"
tool.runtime.credentials = {"auth_type": "api_key_header", "api_key_header_prefix": "basic", "api_key_value": "abc"}
headers = tool.assembling_request(parameters={})
assert headers["Authorization"] == "Basic abc"
assert tool.runtime.credentials["api_key_value"] == "abc"
headers = tool.assembling_request(parameters={})
assert headers["Authorization"] == "Basic abc"
@@ -60,6 +60,37 @@ def test_create_file_by_raw_stores_file_and_persists_record() -> None:
session.refresh.assert_called_once_with(file_model)
def test_create_file_by_raw_prefers_filename_extension_over_mimetype() -> None:
manager = ToolFileManager()
session = Mock()
session.refresh.side_effect = lambda model: setattr(model, "id", "tf-docx")
def tool_file_factory(**kwargs):
return SimpleNamespace(**kwargs)
with (
patch("core.tools.tool_file_manager.storage") as storage,
patch("core.tools.tool_file_manager.ToolFile", side_effect=tool_file_factory),
patch("core.tools.tool_file_manager.uuid4", return_value=SimpleNamespace(hex="abc")),
_patch_session_factory(session),
):
file_model = manager.create_file_by_raw(
user_id="u1",
tenant_id="t1",
conversation_id="c1",
file_binary=b"docx",
mimetype="application/octet-stream",
filename="report.docx",
)
assert file_model.name == "report.docx"
assert file_model.file_key == "tools/t1/abc.docx"
storage.save.assert_called_once_with("tools/t1/abc.docx", b"docx")
session.add.assert_called_once_with(file_model)
session.commit.assert_called_once()
session.refresh.assert_called_once_with(file_model)
def test_create_file_by_url_downloads_and_persists_record() -> None:
manager = ToolFileManager()
response = Mock()
@@ -88,6 +119,32 @@ def test_create_file_by_url_downloads_and_persists_record() -> None:
session.refresh.assert_called_once_with(file_model)
def test_create_file_by_url_prefers_url_extension_over_mimetype() -> None:
manager = ToolFileManager()
response = Mock()
response.content = b"docx"
response.headers = {"Content-Type": "application/octet-stream"}
response.raise_for_status.return_value = None
session = Mock()
def tool_file_factory(**kwargs):
return SimpleNamespace(**kwargs)
session.refresh.side_effect = lambda model: setattr(model, "id", "tf-docx")
with (
patch("core.tools.tool_file_manager.storage") as storage,
patch("core.tools.tool_file_manager.ToolFile", side_effect=tool_file_factory),
patch("core.tools.tool_file_manager.uuid4", return_value=SimpleNamespace(hex="urlabc")),
_patch_session_factory(session),
patch("core.tools.tool_file_manager.remote_fetcher.make_request", return_value=response),
):
file_model = manager.create_file_by_url("u1", "t1", "https://example.com/report.docx?download=1", "c1")
assert file_model.file_key == "tools/t1/urlabc.docx"
assert file_model.name == "urlabc.docx"
storage.save.assert_called_once_with("tools/t1/urlabc.docx", b"docx")
def test_create_file_by_url_raises_on_timeout() -> None:
manager = ToolFileManager()
@@ -7,9 +7,10 @@ from core.tools.entities.tool_entities import ToolInvokeMessage
class _FakeToolFile:
def __init__(self, mimetype: str):
def __init__(self, mimetype: str, name: str | None):
self.id = "fake-tool-file-id"
self.mimetype = mimetype
self.name = name or "fake-tool-file.bin"
class _FakeToolFileManager:
@@ -38,7 +39,7 @@ class _FakeToolFileManager:
"mimetype": mimetype,
"filename": filename,
}
return _FakeToolFile(mimetype)
return _FakeToolFile(mimetype, filename)
@pytest.fixture(autouse=True)
@@ -89,6 +90,29 @@ def test_transform_tool_invoke_messages_mimetype_key_present_but_none():
assert o.meta["tool_file_id"] == "fake-tool-file-id"
def test_transform_tool_invoke_messages_prefers_filename_extension_over_mimetype():
msg = ToolInvokeMessage(
type=ToolInvokeMessage.MessageType.BLOB,
message=ToolInvokeMessage.BlobMessage(blob=b"docx"),
meta={"mime_type": "application/octet-stream", "filename": "report.docx"},
)
out = list(
mt.ToolFileMessageTransformer.transform_tool_invoke_messages(
messages=_gen([msg]),
user_id="u1",
tenant_id="t1",
conversation_id="c1",
)
)
assert _FakeToolFileManager.last_call is not None
assert _FakeToolFileManager.last_call["filename"] == "report.docx"
assert len(out) == 1
assert isinstance(out[0].message, ToolInvokeMessage.TextMessage)
assert out[0].message.text.endswith(".docx")
def test_transform_tool_invoke_messages_parses_existing_tool_file_link_meta():
msg = ToolInvokeMessage(
type=ToolInvokeMessage.MessageType.IMAGE_LINK,
@@ -1,10 +1,12 @@
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import cast
from unittest.mock import MagicMock, patch
from agenton.compositor import CompositorSessionSnapshot
from dify_agent.layers.ask_human import AskHumanToolResult
from dify_agent.protocol import RunStartedEvent, RunSucceededEvent, RunSucceededEventData
from dify_agent.protocol import PydanticAIStreamRunEvent, RunStartedEvent, RunSucceededEvent, RunSucceededEventData
from pydantic_ai.messages import PartDeltaEvent, TextPartDelta
from clients.agent_backend import (
AgentBackendRunEventAdapter,
@@ -190,6 +192,30 @@ class FileOutputBackendClient(FakeAgentBackendRunClient):
)
class AgentMessageDeltaBackendClient(FakeAgentBackendRunClient):
def _events(self, run_id: str):
created_at = datetime(2026, 1, 1, tzinfo=UTC)
return (
RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at),
PydanticAIStreamRunEvent(
id="2-0",
run_id=run_id,
created_at=created_at,
data=PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="hello ")),
agent_message_delta="hello ",
),
RunSucceededEvent(
id="3-0",
run_id=run_id,
created_at=created_at,
data=RunSucceededEventData(
output={"text": "hello agent"},
session_snapshot=CompositorSessionSnapshot(layers=[]),
),
),
)
def _node(
*,
scenario: FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS,
@@ -277,6 +303,19 @@ def test_agent_node_run_maps_successful_agent_backend_run_to_node_result():
assert layers["llm"]["config"]["credentials"] == "[REDACTED]"
def test_agent_node_run_ignores_agent_message_delta_until_terminal_result():
events = list(_node(agent_backend_client=AgentMessageDeltaBackendClient())._run())
assert len(events) == 1
result = cast(StreamCompletedEvent, events[0]).node_run_result
assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED
assert result.outputs == {"text": "hello agent"}
agent_backend = result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"]
assert agent_backend["status"] == "succeeded"
assert agent_backend["agent_message_delta_count"] == 1
assert agent_backend["agent_message_delta_length"] == len("hello ")
def test_agent_node_run_normalizes_declared_file_output_with_canonical_mapping():
tool_reference = build_file_reference(record_id="tool-file-1")
with patch(
@@ -57,6 +57,17 @@ def test_reback_resolves_tenant_tool_file_to_file():
assert file.extension == ".png"
def test_reback_prefers_filename_extension_over_mimetype():
tf = _seed(mimetype="application/octet-stream", name="report.docx", size=99)
file = reback_tool_file_output(tenant_id=TENANT, tool_file_id=tf)
assert file is not None
assert file.filename == "report.docx"
assert file.mime_type == "application/octet-stream"
assert file.extension == ".docx"
assert file.type == FileType.CUSTOM
def test_reback_other_tenant_returns_none():
tf = _seed()
assert reback_tool_file_output(tenant_id="33333333-3333-3333-3333-333333333333", tool_file_id=tf) is None
@@ -165,6 +165,20 @@ def test_build_from_mapping_accepts_opaque_related_id_for_tool_file(mock_tool_fi
assert file.storage_key == "tool_file.pdf"
def test_build_from_mapping_prefers_tool_filename_extension_over_mimetype(mock_tool_file):
mock_tool_file.name = "report.docx"
mock_tool_file.file_key = "tools/test_tenant_id/file.bin"
mock_tool_file.mimetype = "application/octet-stream"
mapping = tool_file_mapping(file_type="document")
file = build_from_mapping(mapping=mapping, tenant_id=TEST_TENANT_ID)
assert file.extension == ".docx"
assert file.filename == "report.docx"
assert file.mime_type == "application/octet-stream"
assert file.storage_key == "tools/test_tenant_id/file.bin"
@pytest.mark.parametrize(
("file_type", "should_pass", "expected_error"),
[
@@ -213,6 +227,25 @@ def test_build_from_remote_url(mock_http_head):
assert file.size == 2048
def test_build_from_remote_url_prefers_filename_extension_over_mimetype():
mapping = {
"transfer_method": "remote_url",
"url": TEST_REMOTE_URL,
"type": "document",
}
with patch(
"factories.file_factory.builders.get_remote_file_info",
return_value=("application/octet-stream", "report.docx", 99),
):
file = build_from_mapping(mapping=mapping, tenant_id=TEST_TENANT_ID)
assert file.filename == "report.docx"
assert file.extension == ".docx"
assert file.mime_type == "application/octet-stream"
assert file.size == 99
@pytest.mark.parametrize(
("file_type", "should_pass", "expected_error"),
[
@@ -1,4 +1,6 @@
from fields.message_fields import ExploreMessageListItem, MessageListItem
from decimal import Decimal
from fields.message_fields import ExploreMessageListItem, MessageListItem, WebMessageListItem
def _base_kwargs():
@@ -34,3 +36,33 @@ class TestExploreMessageListItem:
# Guard the public service-API contract: the base item must not leak metadata.
payload = MessageListItem(**_base_kwargs()).model_dump(mode="json")
assert "metadata" not in payload
def test_message_list_item_exposes_usage_fields(self):
payload = MessageListItem(
**_base_kwargs(),
message_tokens=7,
answer_tokens=11,
provider_response_latency=1.25,
total_price=Decimal("0.0001234"),
currency="USD",
).model_dump(mode="json")
assert payload["message_tokens"] == 7
assert payload["answer_tokens"] == 11
assert payload["total_tokens"] == 18
assert payload["provider_response_latency"] == 1.25
assert payload["total_price"] == "0.0001234"
assert payload["currency"] == "USD"
def test_web_message_list_item_exposes_usage_and_metadata(self):
payload = WebMessageListItem(
**_base_kwargs(),
metadata={"usage": {"total_tokens": 18}},
message_tokens=7,
answer_tokens=11,
).model_dump(mode="json")
assert payload["metadata"] == {"usage": {"total_tokens": 18}}
assert payload["message_tokens"] == 7
assert payload["answer_tokens"] == 11
assert payload["total_tokens"] == 18
@@ -59,6 +59,31 @@ def test_request_download_url_builds_file_under_bound_scope(
assert result.download_url == "https://files.example.com/x"
def test_request_download_url_supports_internal_download_urls() -> None:
fake_file = MagicMock(filename="report.pdf", mime_type="application/pdf", size=123)
service = FileRequestService(access_controller=MagicMock())
with (
patch("services.file_request_service.bind_file_access_scope", return_value=nullcontext()),
patch.object(service, "_build_file", return_value=fake_file),
patch(
"services.file_request_service.file_helpers.resolve_file_url",
return_value="http://internal-files/report.pdf",
) as resolve_file_url,
):
result = service.request_download_url(
tenant_id="tenant-1",
user_id="user-1",
user_from="account",
invoke_from="debugger",
file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:tool-file-1"},
for_external=False,
)
resolve_file_url.assert_called_once_with(fake_file, for_external=False)
assert result.download_url == "http://internal-files/report.pdf"
def test_request_download_url_rejects_unsupported_files() -> None:
service = FileRequestService(access_controller=MagicMock())
@@ -10,9 +10,12 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy import create_engine, select
from sqlalchemy.orm import sessionmaker
import services.summary_index_service as summary_module
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from models.dataset import DocumentSegmentSummary
from models.enums import SegmentStatus, SummaryStatus
from services.summary_index_service import SummaryIndexService
@@ -653,32 +656,48 @@ def test_generate_summaries_for_document_applies_segment_ids_and_only_parent_chu
session.scalars.assert_called()
def test_disable_summaries_for_segments_handles_vector_delete_error(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
summary1 = _summary_record(summary_content="s", node_id="n1")
summary2 = _summary_record(summary_content="s", node_id=None)
def test_disable_summaries_for_segments_updates_sqlite_records() -> None:
dataset = SimpleNamespace(id="dataset-1", indexing_technique=IndexTechniqueType.ECONOMY)
engine = create_engine("sqlite+pysqlite:///:memory:")
DocumentSegmentSummary.__table__.create(engine)
summary_rows = [
{
"id": "sum-1",
"dataset_id": dataset.id,
"document_id": "doc-1",
"chunk_id": "seg-1",
"summary_content": "s",
"summary_index_node_id": "n1",
"status": SummaryStatus.COMPLETED,
"enabled": True,
},
{
"id": "sum-2",
"dataset_id": dataset.id,
"document_id": "doc-1",
"chunk_id": "seg-1",
"summary_content": "s",
"summary_index_node_id": None,
"status": SummaryStatus.COMPLETED,
"enabled": True,
},
]
with engine.begin() as connection:
connection.execute(DocumentSegmentSummary.__table__.insert(), summary_rows)
session = MagicMock()
session.scalars.return_value.all.return_value = [summary1, summary2]
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
monkeypatch.setattr(
summary_module,
"Vector",
MagicMock(return_value=MagicMock(delete_by_ids=MagicMock(side_effect=RuntimeError("boom")))),
)
monkeypatch.setitem(
sys.modules, "libs.datetime_utils", SimpleNamespace(naive_utc_now=MagicMock(return_value=datetime(2024, 1, 1)))
)
session_maker = sessionmaker(bind=engine, expire_on_commit=False)
summary_module.session_factory.configure(engine, expire_on_commit=False)
SummaryIndexService.disable_summaries_for_segments(dataset, segment_ids=["seg-1"], disabled_by="u")
assert summary1.enabled is False
assert summary1.disabled_by == "u"
session.commit.assert_called_once()
with session_maker() as session:
summaries = session.scalars(select(DocumentSegmentSummary).order_by(DocumentSegmentSummary.id)).all()
assert [(summary.id, summary.enabled, summary.disabled_by) for summary in summaries] == [
("sum-1", False, "u"),
("sum-2", False, "u"),
]
assert all(summary.disabled_at is not None for summary in summaries)
def test_disable_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.MonkeyPatch) -> None:
Generated
+24 -17
View File
@@ -322,16 +322,15 @@ wheels = [
[[package]]
name = "anyio"
version = "4.11.0"
version = "4.14.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "idna" },
{ name = "sniffio" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/c6/78/7d432127c41b50bccba979505f272c16cbcadcc33645d5fa3a738110ae75/anyio-4.11.0.tar.gz", hash = "sha256:82a8d0b81e318cc5ce71a5f1f8b5c4e63619620b63141ef8c995fa0db95a57c4", size = 219094, upload-time = "2025-09-23T09:19:12.58Z" }
sdist = { url = "https://files.pythonhosted.org/packages/3b/72/5562aabb8dd7181e8e860622a38bea08d17842b99ecd4c91f84ac95251b0/anyio-4.14.1.tar.gz", hash = "sha256:8d648a3544c1a700e3ff78615cd679e4c5c3f149904287e73687b2596963629e", size = 254831, upload-time = "2026-06-24T20:56:06.017Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/15/b3/9b1a8074496371342ec1e796a96f99c82c945a339cd81a8e73de28b4cf9e/anyio-4.11.0-py3-none-any.whl", hash = "sha256:0287e96f4d26d4149305414d4e3bc32f0dcd0862365a4bddea19d7a1ec38c4fc", size = 109097, upload-time = "2025-09-23T09:19:10.601Z" },
{ url = "https://files.pythonhosted.org/packages/b0/7b/90df4a0a816d98d6ea26f559d87836d494a2cf1fcf063be67df50a7bcc30/anyio-4.14.1-py3-none-any.whl", hash = "sha256:4e5533c5b8ff0a24f5d7a176cbe6877129cd183893f66b537f8f227d10527d72", size = 124875, upload-time = "2026-06-24T20:56:04.413Z" },
]
[[package]]
@@ -1281,10 +1280,12 @@ wheels = [
[[package]]
name = "dify-agent"
version = "0.1.0"
version = "1.16.0rc1"
source = { editable = "../dify-agent" }
dependencies = [
{ name = "anyio" },
{ name = "httpx" },
{ name = "httpx2" },
{ name = "pydantic" },
{ name = "pydantic-ai-slim" },
{ name = "typer" },
@@ -1293,10 +1294,14 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "aiosqlite", marker = "extra == 'shellctl-server'", specifier = ">=0.21.0,<1.0.0" },
{ name = "anyio", specifier = ">=4.12.1,<5.0.0" },
{ name = "fastapi", marker = "extra == 'server'", specifier = "==0.136.0" },
{ name = "fastapi", marker = "extra == 'shellctl-server'", specifier = "==0.136.0" },
{ name = "graphon", marker = "extra == 'server'", specifier = "==0.5.2" },
{ name = "grpclib", extras = ["protobuf"], marker = "extra == 'grpc'", specifier = ">=0.4.9,<0.5.0" },
{ name = "httpx", specifier = "==0.28.1" },
{ name = "httpx2", specifier = ">=2.5.0,<3.0.0" },
{ name = "jsonschema", marker = "extra == 'server'", specifier = ">=4.23.0,<5.0.0" },
{ name = "jwcrypto", marker = "extra == 'server'", specifier = ">=1.5.6,<2" },
{ name = "logfire", extras = ["fastapi", "httpx", "redis"], marker = "extra == 'server'", specifier = ">=4.37.0,<5.0.0" },
@@ -1306,12 +1311,13 @@ requires-dist = [
{ name = "pydantic-ai-slim", extras = ["anthropic", "google", "openai"], marker = "extra == 'server'", specifier = ">=1.85.1,<2.0.0" },
{ name = "pydantic-settings", marker = "extra == 'server'", specifier = ">=2.12.0,<3.0.0" },
{ name = "redis", marker = "extra == 'server'", specifier = ">=7.4.0,<8.0.0" },
{ name = "shell-session-manager", marker = "extra == 'server'", specifier = "==2.4.0" },
{ name = "sqlmodel", marker = "extra == 'shellctl-server'", specifier = ">=0.0.24,<0.1.0" },
{ name = "typer", specifier = ">=0.16.1,<0.17" },
{ name = "typing-extensions", specifier = ">=4.12.2,<5.0.0" },
{ name = "uvicorn", extras = ["standard"], marker = "extra == 'server'", specifier = "==0.46.0" },
{ name = "uvicorn", extras = ["standard"], marker = "extra == 'shellctl-server'", specifier = "==0.46.0" },
]
provides-extras = ["grpc", "server"]
provides-extras = ["grpc", "server", "shellctl-server"]
[package.metadata.requires-dev]
dev = [
@@ -1332,7 +1338,7 @@ docs = [
[[package]]
name = "dify-api"
version = "1.15.0"
version = "1.16.0rc1"
source = { virtual = "." }
dependencies = [
{ name = "aliyun-log-python-sdk" },
@@ -3283,15 +3289,15 @@ wheels = [
[[package]]
name = "httpcore2"
version = "2.3.0"
version = "2.5.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "h11" },
{ name = "truststore" },
]
sdist = { url = "https://files.pythonhosted.org/packages/e6/34/18f1c596e677962f040284246f393b10a1f8ce440b3a7e69c637d0f1c7ad/httpcore2-2.3.0.tar.gz", hash = "sha256:07327e251560960eea8e969d92d4c6a325feb13cca39e25340731336c3baf924", size = 64300, upload-time = "2026-06-01T13:15:02.998Z" }
sdist = { url = "https://files.pythonhosted.org/packages/47/06/5c12df521b5322fb1114a83d46911b2fbcb8855ddb3a635f11c01a214af5/httpcore2-2.5.0.tar.gz", hash = "sha256:88aa170137c17328d5ac44234f9fd10706466d5fb347f3edac4d39b91137b09d", size = 64808, upload-time = "2026-06-25T14:16:56.472Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c2/dd/3357218c69360d1cecc196c230c9a1d5c9afd5dba362056e23e60a5e64e5/httpcore2-2.3.0-py3-none-any.whl", hash = "sha256:477e9e334f74e5240dcac002e890580f36a57d40ff0fb14cc9655731d23b8415", size = 80024, upload-time = "2026-06-01T13:15:00.001Z" },
{ url = "https://files.pythonhosted.org/packages/c9/a1/7564199d1a8728fe737b0a72e5b3f8d92dfe085a74ddf7cdd83bce5f206d/httpcore2-2.5.0-py3-none-any.whl", hash = "sha256:5ce35188de461d31e8d000bfb8ef8bf22c6c16587a211e5571deaa5e9bdf842a", size = 80330, upload-time = "2026-06-25T14:16:53.634Z" },
]
[[package]]
@@ -3355,17 +3361,18 @@ wheels = [
[[package]]
name = "httpx2"
version = "2.3.0"
version = "2.5.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "anyio" },
{ name = "httpcore2" },
{ name = "idna" },
{ name = "truststore" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/9f/9a/cca0b9145f13d8ae34b885ae28d403a1469a433abc78e0f94f4ce94e650b/httpx2-2.3.0.tar.gz", hash = "sha256:227e7c41d95a76d4077a52640564132777215fc3394e07b66a3116c33d668fa9", size = 81115, upload-time = "2026-06-01T13:15:04.324Z" }
sdist = { url = "https://files.pythonhosted.org/packages/d0/e2/b5dedc0cf35aa65de5f541ccd30d2bc1fd7f1d43c9ab09f8ed9a7342317b/httpx2-2.5.0.tar.gz", hash = "sha256:e2df9cb4611021527ff8a675b1c320b610a2ec397acc8d6fe6e91df2d9b33c29", size = 83121, upload-time = "2026-06-25T14:16:57.491Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/87/ce/ae2911859847f9ba1d6b23027e53481cbeb50b93234f355a968d300ca2cb/httpx2-2.3.0-py3-none-any.whl", hash = "sha256:6f393663bdf6dbe7fe90118e3eb5b2bd024a675cae0390ac08cec9198812d8b7", size = 74538, upload-time = "2026-06-01T13:15:01.566Z" },
{ url = "https://files.pythonhosted.org/packages/31/22/859d8252dad9bc9adee34b52e62cde621ece07b042ccb2ab4da1be46695f/httpx2-2.5.0-py3-none-any.whl", hash = "sha256:3d2d4d9cf4b61f1a1f46a95947cfdb47e80cb56a2f91c6256ac8f58e4891df41", size = 76652, upload-time = "2026-06-25T14:16:55.23Z" },
]
[[package]]
@@ -3423,11 +3430,11 @@ wheels = [
[[package]]
name = "idna"
version = "3.11"
version = "3.18"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/6f/6d/0703ccc57f3a7233505399edb88de3cbd678da106337b9fcde432b65ed60/idna-3.11.tar.gz", hash = "sha256:795dafcc9c04ed0c1fb032c2aa73654d8e8c5023a7df64a53f39190ada629902", size = 194582, upload-time = "2025-10-12T14:55:20.501Z" }
sdist = { url = "https://files.pythonhosted.org/packages/cd/63/9496c57188a2ee585e0f1db071d75089a11e98aa86eb99d9d7618fc1edce/idna-3.18.tar.gz", hash = "sha256:ffb385a7e039654cef1ab9ef32c6fafe283c0c0467bba1d9029738ce4a14a848", size = 196711, upload-time = "2026-06-02T14:34:07.794Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0e/61/66938bbb5fc52dbdf84594873d5b51fb1f7c7794e9c0f5bd885f30bc507b/idna-3.11-py3-none-any.whl", hash = "sha256:771a87f49d9defaf64091e6e6fe9c18d4833f140bd19464795bc32d966ca37ea", size = 71008, upload-time = "2025-10-12T14:55:18.883Z" },
{ url = "https://files.pythonhosted.org/packages/1e/5e/d4e9f1a599fb8e573b7b87160658329fbf28d19eac2718f51fc3def3aa5a/idna-3.18-py3-none-any.whl", hash = "sha256:7f952cbe720b688055e3f87de14f5c3e5fdaa8bc3928985c4077ca689de849a2", size = 65455, upload-time = "2026-06-02T14:34:06.319Z" },
]
[[package]]
+12 -4
View File
@@ -20,6 +20,12 @@ DIFY_AGENT_PLUGIN_DAEMON_URL=http://localhost:5002
# API key sent to the Dify plugin daemon.
DIFY_AGENT_PLUGIN_DAEMON_API_KEY=
# Dify API inner endpoints
# Base URL for Dify API inner endpoints used by Agent Stub config/file/drive requests.
DIFY_AGENT_INNER_API_URL=http://localhost:5001
# Must match API/worker INNER_API_KEY_FOR_PLUGIN, not the generic INNER_API_KEY.
DIFY_AGENT_INNER_API_KEY=
# Shell layer
# Base URL for the shellctl server used by the dify.shell layer. Leave empty to disable shell layer use.
DIFY_AGENT_SHELLCTL_ENTRYPOINT=
@@ -30,12 +36,14 @@ DIFY_AGENT_SHELLCTL_AUTH_TOKEN=
# Public Agent Stub URL reachable from shellctl-managed remote machines.
# Use http(s)://.../agent-stub for HTTP or grpc://host:port for gRPC.
# Leave empty to avoid injecting DIFY_AGENT_STUB_* into shell.run jobs.
DIFY_AGENT_STUB_URL=
# Optional bind override used only when DIFY_AGENT_STUB_URL uses grpc://.
DIFY_AGENT_STUB_API_BASE_URL=http://localhost:5050/agent-stub
# Optional bind override used only when DIFY_AGENT_STUB_API_BASE_URL uses grpc://.
DIFY_AGENT_STUB_GRPC_BIND_ADDRESS=
# Server-wide root secret used to derive Agent Stub JWE keys.
# Required when DIFY_AGENT_STUB_URL is set; must be unpadded base64url for 32 bytes.
DIFY_AGENT_SERVER_SECRET_KEY=
# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens.
# Replace this development default in production.
# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))'
DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY
# Shared plugin-daemon HTTP client timeouts and limits.
# Plugin-daemon HTTP connect timeout in seconds.
+2 -2
View File
@@ -6,8 +6,8 @@
# cd /app/api && .venv/bin/uvicorn dify_agent.server.app:app --host 0.0.0.0 --port 5050
#
# Unlike the dify-api image (which only installs the base `dify-agent`
# dependency), this image installs the `[server]` extra, so jwcrypto,
# shell-session-manager, fastapi, uvicorn, etc. are present and the server can
# dependency), this image installs the `[server]` extra, so jwcrypto, fastapi,
# uvicorn, etc. are present and the server can
# actually start. dify-api is intentionally left lean.
# base image
+10 -13
View File
@@ -12,19 +12,15 @@ FROM python:3.12-slim-bookworm AS base
ARG NODE_VERSION=22.22.1
ARG PNPM_VERSION=11.9.0
ARG UV_VERSION=0.8.9
ARG DIFY_AGENT_TOOL_SPEC=.[grpc]
ARG SHELL_SESSION_MANAGER_TOOL_SPEC=shell-session-manager==2.4.0
ARG DIFY_AGENT_TOOL_SPEC=.[grpc,shellctl-server]
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1 \
DIFY_AGENT_STUB_DRIVE_BASE=/mnt/drive \
UV_TOOL_DIR=/opt/dify-agent-tools/envs \
UV_TOOL_BIN_DIR=/opt/dify-agent-tools/bin
ENV PATH="${UV_TOOL_BIN_DIR}:${PATH}"
PIP_NO_CACHE_DIR=1
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
bash \
ca-certificates \
curl \
file \
@@ -62,24 +58,25 @@ WORKDIR /opt/dify-agent
FROM base AS tools
ARG DIFY_AGENT_TOOL_SPEC
ARG SHELL_SESSION_MANAGER_TOOL_SPEC
COPY pyproject.toml uv.lock README.md ./
COPY src ./src
RUN uv export --frozen --no-dev --all-extras --no-emit-project --no-hashes \
> /tmp/dify-agent-constraints.txt \
&& uv tool install --force --python /usr/local/bin/python --no-python-downloads \
&& UV_TOOL_DIR=/opt/dify-agent-tools/envs \
UV_TOOL_BIN_DIR=/opt/dify-agent-tools/bin \
uv tool install --force --python /usr/local/bin/python --no-python-downloads \
--constraints /tmp/dify-agent-constraints.txt --link-mode=copy "${DIFY_AGENT_TOOL_SPEC}" \
&& uv tool install --force --python /usr/local/bin/python --no-python-downloads \
--constraints /tmp/dify-agent-constraints.txt --link-mode=copy "${SHELL_SESSION_MANAGER_TOOL_SPEC}" \
&& rm -f /tmp/dify-agent-constraints.txt
FROM base AS production
COPY --from=tools ${UV_TOOL_DIR} ${UV_TOOL_DIR}
COPY --from=tools ${UV_TOOL_BIN_DIR} ${UV_TOOL_BIN_DIR}
ENV PATH="/opt/dify-agent-tools/bin:${PATH}"
COPY --from=tools /opt/dify-agent-tools/envs /opt/dify-agent-tools/envs
COPY --from=tools /opt/dify-agent-tools/bin /opt/dify-agent-tools/bin
RUN useradd --create-home --shell /bin/sh dify \
&& mkdir -p /mnt/drive \
@@ -75,8 +75,9 @@ See `.example.env` for the full server settings template.
If you plan to run `dify.shell`, also configure `DIFY_AGENT_SHELLCTL_ENTRYPOINT`
and, when shell jobs need to call back with the `dify-agent` command, set
`DIFY_AGENT_STUB_API_BASE_URL` plus a 32-byte base64url
`DIFY_AGENT_SERVER_SECRET_KEY` as documented in `.example.env`.
`DIFY_AGENT_STUB_API_BASE_URL`. The supplied default configs include a
development `DIFY_AGENT_SERVER_SECRET_KEY`, but production deployments should
override it with a unique 32-byte base64url value as documented in `.example.env`.
## Start the Dify Agent server
+5 -3
View File
@@ -42,7 +42,7 @@ also reads `.env` and `dify-agent/.env` when present.
| `DIFY_AGENT_SHELLCTL_AUTH_TOKEN` | empty | Optional bearer token sent to the shellctl server. |
| `DIFY_AGENT_STUB_API_BASE_URL` | empty | Public Agent Stub API base URL reachable from shellctl-managed remote machines. HTTP may be the service root or `/agent-stub`; gRPC must be `grpc://host:port`. Enables `DIFY_AGENT_STUB_*` env injection for user `shell.run` jobs. |
| `DIFY_AGENT_STUB_GRPC_BIND_ADDRESS` | empty | Optional `host:port` bind override used only when `DIFY_AGENT_STUB_API_BASE_URL` uses `grpc://`. |
| `DIFY_AGENT_SERVER_SECRET_KEY` | empty | Server-wide root secret used to derive Agent Stub JWE keys; required when `DIFY_AGENT_STUB_API_BASE_URL` is set and must be unpadded base64url for 32 bytes. |
| `DIFY_AGENT_SERVER_SECRET_KEY` | empty | Security-sensitive server-wide root secret used to derive the JWE encryption key for Agent Stub bearer tokens; required when `DIFY_AGENT_STUB_API_BASE_URL` is set. The supplied default config uses a development value; set a unique unpadded base64url 32-byte secret in production. |
| `DIFY_AGENT_PLUGIN_DAEMON_CONNECT_TIMEOUT` | `10` | Plugin-daemon HTTP connect timeout in seconds. |
| `DIFY_AGENT_PLUGIN_DAEMON_READ_TIMEOUT` | `600` | Plugin-daemon HTTP read timeout in seconds. |
| `DIFY_AGENT_PLUGIN_DAEMON_WRITE_TIMEOUT` | `30` | Plugin-daemon HTTP write timeout in seconds. |
@@ -64,9 +64,11 @@ DIFY_AGENT_INNER_API_URL=http://localhost:5001
DIFY_AGENT_INNER_API_KEY=replace-with-dify-inner-api-key-for-plugin
DIFY_AGENT_SHELLCTL_ENTRYPOINT=http://127.0.0.1:5004
DIFY_AGENT_SHELLCTL_AUTH_TOKEN=replace-with-shellctl-token
# Generate with: python -c 'import base64, secrets; print(base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode())'
DIFY_AGENT_STUB_API_BASE_URL=https://agent.example.com/agent-stub
DIFY_AGENT_SERVER_SECRET_KEY=replace-with-base64url-32-byte-secret
# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens.
# Replace this development default in production.
# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))'
DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY
```
Run records and event streams use the same retention. Status writes refresh the
@@ -56,18 +56,22 @@ with `dify-agent ...`, also enable the Agent Stub:
```env
DIFY_AGENT_STUB_API_BASE_URL=https://agent.example.com/agent-stub
DIFY_AGENT_SERVER_SECRET_KEY=replace-with-base64url-32-byte-secret
# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens.
# Replace this development default in production.
# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))'
DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY
```
HTTP `DIFY_AGENT_STUB_API_BASE_URL` may be either the service root or the
explicit `/agent-stub` API root; the server normalizes the service root to
`/agent-stub`. Other HTTP paths are rejected at startup.
`DIFY_AGENT_SERVER_SECRET_KEY` must be unpadded base64url text for exactly 32
decoded bytes. One way to generate it is:
The supplied Docker and `.example.env` configs use a development
`DIFY_AGENT_SERVER_SECRET_KEY`. Override it in production with unpadded base64url
text for exactly 32 decoded bytes. One way to generate it is:
```bash
python -c 'import base64, secrets; print(base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode())'
python -c 'import secrets; print(secrets.token_urlsafe(32))'
```
## Client request shape
@@ -230,12 +234,11 @@ The provided `docker/local-sandbox/Dockerfile` installs:
- `tmux`, required by `shellctl` to manage shell jobs;
- common shell workspace tools: `git`, `openssh-client`, `jq`, `ripgrep`,
`unzip`, `zip`, `file`, `procps`, and `less`;
- `shell-session-manager==2.3.1` as a standalone uv tool, which provides the
`shellctl` CLI/server;
- `dify-agent[grpc,shellctl-server]` as a standalone uv tool, which provides
both the Agent Stub client CLI and the built-in `shellctl` CLI/server;
- `uv`, so uv shebang scripts with PEP 723 metadata can run inside the shell
workspace and Python CLI tools can be installed with isolated tool
environments;
- `node==22.22.1` and `pnpm==11.9.0`, so JavaScript and TypeScript tooling can
run inside the shell workspace without per-job installation;
- the `dify-agent[grpc]` Agent Stub client CLI as a standalone uv tool;
- a non-root default user named `dify`.
@@ -36,6 +36,7 @@ message FileMapping {
message FileDownloadRequest {
FileMapping file = 1;
optional bool for_external = 2;
}
message FileDownloadResponse {
+13 -3
View File
@@ -1,11 +1,13 @@
[project]
name = "dify-agent"
version = "0.1.0"
version = "1.16.0-rc1"
description = "Add your description here"
readme = "README.md"
requires-python = ">=3.12,<4.0"
dependencies = [
"anyio>=4.12.1,<5.0.0",
"httpx==0.28.1",
"httpx2>=2.5.0,<3.0.0",
"pydantic>=2.12.5,<2.13",
"pydantic-ai-slim>=1.102.0,<2.0.0",
"typer>=0.16.1,<0.17",
@@ -15,6 +17,9 @@ dependencies = [
[project.scripts]
dify-agent = "dify_agent.agent_stub.cli.main:main"
dify-agent-stub-server = "dify_agent.agent_stub.server.cli:main"
shellctl = "shellctl.cli:main"
shellctl-sanitize-pty = "shellctl_runtime.sanitize:main"
shellctl-runner-exit = "shellctl_runtime.runner_exit:main"
[project.optional-dependencies]
grpc = ["grpclib[protobuf]>=0.4.9,<0.5.0", "protobuf>=6.33.5,<7.0.0"]
@@ -27,13 +32,18 @@ server = [
"pydantic-ai-slim[anthropic,google,openai]>=1.85.1,<2.0.0",
"pydantic-settings>=2.12.0,<3.0.0",
"redis>=7.4.0,<8.0.0",
"shell-session-manager==2.4.0",
"uvicorn[standard]==0.46.0",
]
shellctl-server = [
"aiosqlite>=0.21.0,<1.0.0",
"fastapi==0.136.0",
"sqlmodel>=0.0.24,<0.1.0",
"uvicorn[standard]==0.46.0",
]
[tool.setuptools.packages.find]
where = ["src"]
include = ["agenton*", "agenton_collections*", "dify_agent*"]
include = ["agenton*", "agenton_collections*", "dify_agent*", "shellctl*", "shellctl_runtime*"]
[tool.pyright]
include = ["src", "examples", "tests"]
@@ -205,6 +205,7 @@ class DifyStreamedResponse(StreamedResponse):
@override
async def _get_event_iterator(self) -> AsyncIterator[ModelResponseStreamEvent]:
chunk_sequence = 0
async for chunk in self.chunks:
if chunk.delta.usage is not None:
self._usage: RequestUsage = _map_usage(chunk.delta.usage)
@@ -216,8 +217,10 @@ class DifyStreamedResponse(StreamedResponse):
chunk,
self.provider_name_value,
self._embedded_thinking_parser,
chunk_sequence,
):
yield event
chunk_sequence += 1
for event in self._embedded_thinking_parser.flush(self._parts_manager, self.provider_name_value):
yield event
@@ -551,11 +554,21 @@ def _normalize_finish_reason(finish_reason: str) -> FinishReason:
return "error"
def _normalize_tool_call_id(tool_call_id: str | None) -> str | None:
if tool_call_id is None:
return None
normalized = tool_call_id.strip()
if not normalized or normalized.lower() in {"none", "null"}:
return None
return normalized
def _chunk_to_stream_events(
parts_manager: ModelResponsePartsManager,
chunk: LLMResultChunk,
provider_name: str,
embedded_thinking_parser: "_EmbeddedThinkingParser",
chunk_sequence: int,
) -> list[ModelResponseStreamEvent]:
events: list[ModelResponseStreamEvent] = []
message = chunk.delta.message
@@ -571,13 +584,14 @@ def _chunk_to_stream_events(
events.append(parts_manager.handle_part(vendor_part_id=None, part=part))
for index, tool_call in enumerate(message.tool_calls):
vendor_id = tool_call.id or f"chunk-{chunk.delta.index}-tool-{index}"
tool_call_id = _normalize_tool_call_id(tool_call.id)
vendor_id = tool_call_id or f"chunk-{chunk_sequence}-tool-{index}"
events.append(
parts_manager.handle_tool_call_part(
vendor_part_id=vendor_id,
tool_name=tool_call.function.name,
args=tool_call.function.arguments,
tool_call_id=tool_call.id,
tool_call_id=tool_call_id or vendor_id,
provider_name=provider_name,
)
)
@@ -1,6 +1,6 @@
"""Shellctl-backed shell provider adapter for dify-agent.
The shell-session-manager SDK owns the HTTP timeout policy for long-polling
The built-in shellctl SDK owns the HTTP timeout policy for long-polling
shellctl requests. This adapter stays narrowly focused on translating SDK and
transport failures into ``ShellProviderError`` so the shell layer can return
tool observations instead of aborting the agent loop.
@@ -16,9 +16,9 @@ import time
from collections.abc import Awaitable
from collections.abc import Callable
from dataclasses import dataclass
from typing import Protocol, TypeVar
from typing import Protocol, TypeVar, cast
import httpx
import httpx2 as httpx
from dify_agent.adapters.shell.protocols import (
ShellCommandProtocol,
@@ -228,8 +228,16 @@ class ShellctlFileTransfer(ShellFileTransferProtocol):
@dataclass(slots=True)
class ShellctlResource(ShellResourceProtocol):
client: ShellctlClientProtocol
commands: ShellCommandProtocol
files: ShellFileTransferProtocol
_commands: ShellCommandProtocol
_files: ShellFileTransferProtocol
@property
def commands(self) -> ShellCommandProtocol:
return self._commands
@property
def files(self) -> ShellFileTransferProtocol:
return self._files
async def close(self) -> None:
try:
@@ -257,8 +265,8 @@ class ShellctlProvider(ShellProviderProtocol):
)
return ShellctlResource(
client=client,
commands=ShellctlCommands(client=client),
files=ShellctlFileTransfer(client=client),
_commands=ShellctlCommands(client=client),
_files=ShellctlFileTransfer(client=client),
)
@@ -269,9 +277,15 @@ def create_default_shellctl_client_factory(
output_limit: int = _SHELLCTL_OUTPUT_LIMIT_BYTES,
) -> ShellctlClientFactory:
def factory() -> ShellctlClientProtocol:
from shell_session_manager.shellctl.client import ShellctlClient
from shellctl.client import ShellctlClient
return ShellctlClient(entrypoint, token=token, output_limit=output_limit)
return cast(
ShellctlClientProtocol,
cast(
object,
ShellctlClient(entrypoint, token=token, output_limit=output_limit),
),
)
return factory
@@ -143,6 +143,7 @@ def download_file_from_environment(
url=environment.url,
auth_jwe=environment.auth_jwe,
file=file_mapping,
for_external=False,
)
if not hasattr(download_request, "filename") or not isinstance(download_request.filename, str):
raise AgentStubTransferError("signed file download response is missing filename")
@@ -207,8 +208,10 @@ def _request_uploaded_tool_file_download_url(*, url: str, auth_jwe: str, referen
file=AgentStubFileMapping(transfer_method="tool_file", reference=reference),
),
)
if not hasattr(download_request, "download_url") or not isinstance(download_request.download_url, str):
raise AgentStubTransferError("signed file download response is missing download_url")
download_url = download_request.download_url
if not isinstance(download_url, str) or not download_url:
if not download_url:
raise AgentStubTransferError("signed file download response is missing download_url")
return download_url
@@ -133,6 +133,26 @@ def config_skills_push(
Pass a directory such as ./skills/researcher that contains SKILL.md. Other files in that directory are
archived with the skill. Pushing a skill with an existing name replaces that config skill.
Skill directory requirements:
- Each PATH must be one skill directory; the directory basename is the config skill name.
- The directory must contain a top-level SKILL.md.
- SKILL.md must be non-empty UTF-8 Markdown.
- SKILL.md must start with YAML frontmatter matching this schema:
\b
---
name: <non-empty string>
description: <string>
---
- Symlinked files are rejected.
- Dependency/cache folders such as .git, __pycache__, .venv and node_modules should be manually cleared before push.
"""
_run_config_skills_push(paths=paths)
@@ -97,6 +97,7 @@ def request_agent_stub_file_download_sync(
url: str,
auth_jwe: str,
file: AgentStubFileMapping,
for_external: bool = True,
timeout: float | httpx.Timeout = 30.0,
sync_http_client: httpx.Client | None = None,
):
@@ -109,12 +110,14 @@ def request_agent_stub_file_download_sync(
url=endpoint.url,
auth_jwe=auth_jwe,
file=file,
for_external=for_external,
timeout=timeout,
)
return request_agent_stub_file_download_http_sync(
base_url=endpoint.url,
auth_jwe=auth_jwe,
file=file,
for_external=for_external,
timeout=timeout,
sync_http_client=sync_http_client,
)
@@ -124,6 +124,7 @@ def request_agent_stub_file_download_grpc_sync(
url: str,
auth_jwe: str,
file: AgentStubFileMapping,
for_external: bool = True,
timeout: float | httpx.Timeout = 30.0,
):
"""Request one signed download URL through the gRPC Agent Stub endpoint.
@@ -144,7 +145,7 @@ def request_agent_stub_file_download_grpc_sync(
auth_jwe=auth_jwe,
method_name="CreateFileDownloadRequest",
request_factory=lambda runtime: _require_conversions().proto_file_download_request(
runtime.agent_stub_pb2, file=file
runtime.agent_stub_pb2, file=file, for_external=for_external
),
response_parser=lambda response: _require_conversions().file_download_response_from_proto(response),
timeout=timeout,
@@ -117,13 +117,14 @@ def request_agent_stub_file_download_http_sync(
base_url: str,
auth_jwe: str,
file: AgentStubFileMapping,
for_external: bool = True,
timeout: float | httpx.Timeout = 30.0,
sync_http_client: httpx.Client | None = None,
) -> AgentStubFileDownloadResponse:
"""Request one signed download URL from the HTTP Agent Stub endpoint."""
try:
request_model = AgentStubFileDownloadRequest(file=file)
request_model = AgentStubFileDownloadRequest(file=file, for_external=for_external)
except ValidationError as exc:
raise AgentStubValidationError("invalid Agent Stub file download request") from exc
response = _post_agent_stub_json(
@@ -131,7 +132,7 @@ def request_agent_stub_file_download_http_sync(
auth_jwe=auth_jwe,
endpoint_name="file download request",
endpoint_url_factory=agent_stub_file_download_request_url,
request_body=request_model.model_dump_json(exclude_none=True),
request_body=request_model.model_dump_json(exclude_none=True, exclude_defaults=True),
timeout=timeout,
sync_http_client=sync_http_client,
)
@@ -1,6 +1,6 @@
# pyright: reportAttributeAccessIssue=false
# -*- coding: utf-8 -*-
# Generated by the protocol buffer compiler. DO NOT EDIT!
# NO CHECKED-IN PROTOBUF GENCODE
# source: dify/agent/stub/v1/agent_stub.proto
# Protobuf Python Version: 6.33.5
"""Generated protocol buffer code."""
@@ -12,12 +12,7 @@ from google.protobuf import symbol_database as _symbol_database
from google.protobuf.internal import builder as _builder
_runtime_version.ValidateProtobufRuntimeVersion(
_runtime_version.Domain.PUBLIC,
6,
33,
5,
"",
"dify/agent/stub/v1/agent_stub.proto",
_runtime_version.Domain.PUBLIC, 6, 33, 5, "", "dify/agent/stub/v1/agent_stub.proto"
)
# @@protoc_insertion_point(imports)
@@ -25,12 +20,12 @@ _sym_db = _symbol_database.Default()
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(
b'\n#dify/agent/stub/v1/agent_stub.proto\x12\x12\x64ify.agent.stub.v1"O\n\x0e\x43onnectRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\x05\x12\x0c\n\x04\x61rgv\x18\x02 \x03(\t\x12\x15\n\rmetadata_json\x18\x03 \x01(\t"8\n\x0f\x43onnectResponse\x12\x15\n\rconnection_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t"7\n\x11\x46ileUploadRequest\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x10\n\x08mimetype\x18\x02 \x01(\t"(\n\x12\x46ileUploadResponse\x12\x12\n\nupload_url\x18\x01 \x01(\t"f\n\x0b\x46ileMapping\x12\x17\n\x0ftransfer_method\x18\x01 \x01(\t\x12\x16\n\treference\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x10\n\x03url\x18\x03 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_referenceB\x06\n\x04_url"D\n\x13\x46ileDownloadRequest\x12-\n\x04\x66ile\x18\x01 \x01(\x0b\x32\x1f.dify.agent.stub.v1.FileMapping"r\n\x14\x46ileDownloadResponse\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x16\n\tmime_type\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x0c\n\x04size\x18\x03 \x01(\x03\x12\x14\n\x0c\x64ownload_url\x18\x04 \x01(\tB\x0c\n\n_mime_type2\xc0\x02\n\x10\x41gentStubService\x12R\n\x07\x43onnect\x12".dify.agent.stub.v1.ConnectRequest\x1a#.dify.agent.stub.v1.ConnectResponse\x12h\n\x17\x43reateFileUploadRequest\x12%.dify.agent.stub.v1.FileUploadRequest\x1a&.dify.agent.stub.v1.FileUploadResponse\x12n\n\x19\x43reateFileDownloadRequest\x12\'.dify.agent.stub.v1.FileDownloadRequest\x1a(.dify.agent.stub.v1.FileDownloadResponseb\x06proto3'
b'\n#dify/agent/stub/v1/agent_stub.proto\x12\x12\x64ify.agent.stub.v1"O\n\x0e\x43onnectRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\x05\x12\x0c\n\x04\x61rgv\x18\x02 \x03(\t\x12\x15\n\rmetadata_json\x18\x03 \x01(\t"8\n\x0f\x43onnectResponse\x12\x15\n\rconnection_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t"7\n\x11\x46ileUploadRequest\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x10\n\x08mimetype\x18\x02 \x01(\t"(\n\x12\x46ileUploadResponse\x12\x12\n\nupload_url\x18\x01 \x01(\t"f\n\x0b\x46ileMapping\x12\x17\n\x0ftransfer_method\x18\x01 \x01(\t\x12\x16\n\treference\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x10\n\x03url\x18\x03 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_referenceB\x06\n\x04_url"p\n\x13\x46ileDownloadRequest\x12-\n\x04\x66ile\x18\x01 \x01(\x0b\x32\x1f.dify.agent.stub.v1.FileMapping\x12\x19\n\x0c\x66or_external\x18\x02 \x01(\x08H\x00\x88\x01\x01\x42\x0f\n\r_for_external"r\n\x14\x46ileDownloadResponse\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x16\n\tmime_type\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x0c\n\x04size\x18\x03 \x01(\x03\x12\x14\n\x0c\x64ownload_url\x18\x04 \x01(\tB\x0c\n\n_mime_type2\xc0\x02\n\x10\x41gentStubService\x12R\n\x07\x43onnect\x12".dify.agent.stub.v1.ConnectRequest\x1a#.dify.agent.stub.v1.ConnectResponse\x12h\n\x17\x43reateFileUploadRequest\x12%.dify.agent.stub.v1.FileUploadRequest\x1a&.dify.agent.stub.v1.FileUploadResponse\x12n\n\x19\x43reateFileDownloadRequest\x12\'.dify.agent.stub.v1.FileDownloadRequest\x1a(.dify.agent.stub.v1.FileDownloadResponseb\x06proto3'
)
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "dify_agent.agent_stub.grpc._generated.agent_stub_pb2", _globals)
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "dify.agent.stub.v1.agent_stub_pb2", _globals)
if not _descriptor._USE_C_DESCRIPTORS:
DESCRIPTOR._loaded_options = None
_globals["_CONNECTREQUEST"]._serialized_start = 59
@@ -44,9 +39,9 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals["_FILEMAPPING"]._serialized_start = 297
_globals["_FILEMAPPING"]._serialized_end = 399
_globals["_FILEDOWNLOADREQUEST"]._serialized_start = 401
_globals["_FILEDOWNLOADREQUEST"]._serialized_end = 469
_globals["_FILEDOWNLOADRESPONSE"]._serialized_start = 471
_globals["_FILEDOWNLOADRESPONSE"]._serialized_end = 585
_globals["_AGENTSTUBSERVICE"]._serialized_start = 588
_globals["_AGENTSTUBSERVICE"]._serialized_end = 908
_globals["_FILEDOWNLOADREQUEST"]._serialized_end = 513
_globals["_FILEDOWNLOADRESPONSE"]._serialized_start = 515
_globals["_FILEDOWNLOADRESPONSE"]._serialized_end = 629
_globals["_AGENTSTUBSERVICE"]._serialized_start = 632
_globals["_AGENTSTUBSERVICE"]._serialized_end = 952
# @@protoc_insertion_point(module_scope)
@@ -1,71 +1,69 @@
from __future__ import annotations
from google.protobuf.internal import containers as _containers
from google.protobuf import descriptor as _descriptor
from google.protobuf import message as _message
from collections.abc import Iterable as _Iterable, Mapping as _Mapping
from typing import ClassVar as _ClassVar, Optional as _Optional, Union as _Union
from collections.abc import Iterable
DESCRIPTOR: _descriptor.FileDescriptor
from google.protobuf.message import Message
class ConnectRequest(Message):
class ConnectRequest(_message.Message):
__slots__ = ("protocol_version", "argv", "metadata_json")
PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int]
ARGV_FIELD_NUMBER: _ClassVar[int]
METADATA_JSON_FIELD_NUMBER: _ClassVar[int]
protocol_version: int
argv: list[str]
argv: _containers.RepeatedScalarFieldContainer[str]
metadata_json: str
def __init__(self, protocol_version: _Optional[int] = ..., argv: _Optional[_Iterable[str]] = ..., metadata_json: _Optional[str] = ...) -> None: ...
def __init__(
self,
*,
protocol_version: int = ...,
argv: Iterable[str] = ...,
metadata_json: str = ...,
) -> None: ...
class ConnectResponse(Message):
class ConnectResponse(_message.Message):
__slots__ = ("connection_id", "status")
CONNECTION_ID_FIELD_NUMBER: _ClassVar[int]
STATUS_FIELD_NUMBER: _ClassVar[int]
connection_id: str
status: str
def __init__(self, connection_id: _Optional[str] = ..., status: _Optional[str] = ...) -> None: ...
def __init__(self, *, connection_id: str = ..., status: str = ...) -> None: ...
class FileUploadRequest(Message):
class FileUploadRequest(_message.Message):
__slots__ = ("filename", "mimetype")
FILENAME_FIELD_NUMBER: _ClassVar[int]
MIMETYPE_FIELD_NUMBER: _ClassVar[int]
filename: str
mimetype: str
def __init__(self, filename: _Optional[str] = ..., mimetype: _Optional[str] = ...) -> None: ...
def __init__(self, *, filename: str = ..., mimetype: str = ...) -> None: ...
class FileUploadResponse(Message):
class FileUploadResponse(_message.Message):
__slots__ = ("upload_url",)
UPLOAD_URL_FIELD_NUMBER: _ClassVar[int]
upload_url: str
def __init__(self, upload_url: _Optional[str] = ...) -> None: ...
def __init__(self, *, upload_url: str = ...) -> None: ...
class FileMapping(Message):
class FileMapping(_message.Message):
__slots__ = ("transfer_method", "reference", "url")
TRANSFER_METHOD_FIELD_NUMBER: _ClassVar[int]
REFERENCE_FIELD_NUMBER: _ClassVar[int]
URL_FIELD_NUMBER: _ClassVar[int]
transfer_method: str
reference: str
url: str
def __init__(self, transfer_method: _Optional[str] = ..., reference: _Optional[str] = ..., url: _Optional[str] = ...) -> None: ...
def __init__(self, *, transfer_method: str = ..., reference: str = ..., url: str = ...) -> None: ...
def HasField(self, field_name: str) -> bool: ...
class FileDownloadRequest(Message):
class FileDownloadRequest(_message.Message):
__slots__ = ("file", "for_external")
FILE_FIELD_NUMBER: _ClassVar[int]
FOR_EXTERNAL_FIELD_NUMBER: _ClassVar[int]
file: FileMapping
for_external: bool
def __init__(self, file: _Optional[_Union[FileMapping, _Mapping]] = ..., for_external: _Optional[bool] = ...) -> None: ...
def __init__(self, *, file: FileMapping | None = ...) -> None: ...
class FileDownloadResponse(Message):
class FileDownloadResponse(_message.Message):
__slots__ = ("filename", "mime_type", "size", "download_url")
FILENAME_FIELD_NUMBER: _ClassVar[int]
MIME_TYPE_FIELD_NUMBER: _ClassVar[int]
SIZE_FIELD_NUMBER: _ClassVar[int]
DOWNLOAD_URL_FIELD_NUMBER: _ClassVar[int]
filename: str
mime_type: str
size: int
download_url: str
def __init__(
self,
*,
filename: str = ...,
mime_type: str = ...,
size: int = ...,
download_url: str = ...,
) -> None: ...
def HasField(self, field_name: str) -> bool: ...
def __init__(self, filename: _Optional[str] = ..., mime_type: _Optional[str] = ..., size: _Optional[int] = ..., download_url: _Optional[str] = ...) -> None: ...
@@ -98,13 +98,19 @@ def file_download_request_from_proto(message: agent_stub_pb2.FileDownloadRequest
"reference": message.file.reference if message.file.HasField("reference") else None,
"url": message.file.url if message.file.HasField("url") else None,
}
return AgentStubFileDownloadRequest.model_validate({"file": file_mapping_kwargs})
return AgentStubFileDownloadRequest.model_validate(
{
"file": file_mapping_kwargs,
"for_external": message.for_external if message.HasField("for_external") else True,
}
)
def proto_file_download_request(
pb2_module,
*,
file: AgentStubFileMapping,
for_external: bool = True,
) -> agent_stub_pb2.FileDownloadRequest:
"""Build one protobuf file-download request from the public DTO."""
mapping = pb2_module.FileMapping(transfer_method=file.transfer_method)
@@ -112,7 +118,9 @@ def proto_file_download_request(
mapping.reference = file.reference
if file.url is not None:
mapping.url = file.url
return pb2_module.FileDownloadRequest(file=mapping)
request = pb2_module.FileDownloadRequest(file=mapping)
request.for_external = for_external
return request
def file_download_response_from_proto(message: agent_stub_pb2.FileDownloadResponse) -> AgentStubFileDownloadResponse:
@@ -249,6 +249,7 @@ class AgentStubFileDownloadRequest(BaseModel):
"""Request body for one signed download URL allocation."""
file: AgentStubFileMapping
for_external: bool = True
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid")
@@ -161,6 +161,8 @@ class DifyApiAgentStubFileRequestHandler:
"invoke_from": execution_context.invoke_from,
"file": request.file.model_dump(mode="json", exclude_none=True),
}
if request.for_external is False:
payload["for_external"] = False
data = await self._post_inner_api("/inner/api/download/file/request", payload)
try:
return AgentStubFileDownloadResponse.model_validate(data)
@@ -141,7 +141,18 @@ shell_run script rules:
Tips:
- When using Python, prefer a uv script with a PEP 723 dependency header.
- Python 3.12, uv, pip, Node.js, pnpm, and pnx are preinstalled in the local sandbox.
- For one-off Python dependencies, prefer a uv script with a PEP 723 dependency header or:
`uv run --with <package> python <script-or--c>`.
- For reusable Python CLI tools, use `uv tool install <tool>`; installed commands land in `$HOME/.local/bin`.
Run them by full path or add `$HOME/.local/bin` to PATH in the command that needs them.
- `python3 -m pip install --user <package>` also installs into `$HOME/.local`; add `$HOME/.local/bin` to PATH
when you need console scripts.
- For reusable Node.js CLIs, use user-level global installs:
`PNPM_HOME=$HOME/.local/share/pnpm PATH=$HOME/.local/share/pnpm/bin:$PATH pnpm add -g <package>`.
Installed commands land in `$PNPM_HOME/bin`; run them by full path or with the same PATH prefix.
- For one-off Node.js CLIs, prefer `pnx <command> [args]`.
- Do not install new packages into system or image tool paths such as `/usr/local`, `/usr`, or `/opt/dify-agent-tools`.
- If you need MCP, install the MCP server in the shell environment and start that server when you use it.
Example shell_run script:
+26 -3
View File
@@ -31,13 +31,14 @@ both the JSON-safe final output or deferred tool call and the session snapshot;
there are no separate output or snapshot events to correlate.
"""
from collections.abc import AsyncIterable, Callable
from collections.abc import AsyncIterable, Callable, Mapping
from collections import Counter
from dataclasses import dataclass
from typing import Any, Literal, Protocol, cast, runtime_checkable
import httpx
from pydantic import JsonValue, TypeAdapter
from pydantic_ai.exceptions import ModelHTTPError
from pydantic_ai.messages import AgentStreamEvent, PartDeltaEvent, PartStartEvent, TextPart, TextPartDelta
from pydantic_ai.output import OutputSpec
from pydantic_ai.tools import DeferredToolRequests, DeferredToolResults
@@ -104,6 +105,28 @@ class AgentRunValidationError(ValueError):
"""Raised when a run request is valid JSON but cannot execute."""
def _run_failed_error_payload(exc: Exception) -> tuple[str, str | None]:
"""Return the public failed-run error text and structured reason."""
message = str(exc) or type(exc).__name__
reason: str | None = None
if isinstance(exc, ModelHTTPError):
body = exc.body
if isinstance(body, Mapping):
body_message = body.get("message")
if isinstance(body_message, str) and body_message:
message = body_message
error_type = body.get("error_type")
if isinstance(error_type, str) and error_type:
reason = error_type
if reason is None and exc.status_code == 429:
reason = "InvokeRateLimitError"
return message, reason
def _has_model_layer(request: CreateRunRequest) -> bool:
"""Return whether the public composition includes the reserved model layer."""
return any(layer.name == DIFY_AGENT_MODEL_LAYER_ID for layer in request.composition.layers)
@@ -165,8 +188,8 @@ class AgentRunRunner:
try:
outcome = await self._run_agent()
except Exception as exc:
message = str(exc) or type(exc).__name__
_ = await emit_run_failed(self.sink, run_id=self.run_id, error=message)
message, reason = _run_failed_error_payload(exc)
_ = await emit_run_failed(self.sink, run_id=self.run_id, error=message, reason=reason)
await self.sink.update_status(self.run_id, "failed", message)
raise
+124
View File
@@ -0,0 +1,124 @@
"""Public shellctl package exports.
This package stays lazy on purpose. Hot-path runtime helpers live outside the
`shellctl` package, and importing this package root should
not pull the full client/server/public DTO surface unless a caller explicitly
asks for those exports.
"""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from shellctl.client import (
ShellctlClient,
ShellctlClientError,
)
from shellctl.shared import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_BASE_URL,
DEFAULT_BASE_URL_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
DeleteJobResponse,
HealthResponse,
InputJobRequest,
JobInfo,
JobResult,
JobStatusName,
JobStatusView,
ListJobsResponse,
RunJobRequest,
TerminalSize,
TerminateJobRequest,
WaitJobRequest,
generate_job_id,
read_output_window,
tail_output_window,
)
__all__ = [
"DEFAULT_AUTH_TOKEN_ENV",
"DEFAULT_BASE_URL",
"DEFAULT_BASE_URL_ENV",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS",
"DEFAULT_GC_INTERVAL_SECONDS",
"DEFAULT_IDLE_FLUSH_SECONDS",
"DEFAULT_LIST_LIMIT",
"DEFAULT_OUTPUT_LIMIT_BYTES",
"DEFAULT_TERMINAL_COLS",
"DEFAULT_TERMINAL_ROWS",
"DEFAULT_TERMINATE_GRACE_SECONDS",
"DEFAULT_TIMEOUT_SECONDS",
"DeleteJobResponse",
"HealthResponse",
"InputJobRequest",
"JobInfo",
"JobResult",
"JobStatusName",
"JobStatusView",
"ListJobsResponse",
"RunJobRequest",
"ShellctlClient",
"ShellctlClientError",
"TerminalSize",
"TerminateJobRequest",
"WaitJobRequest",
"generate_job_id",
"read_output_window",
"tail_output_window",
]
_EXPORTS = {
"ShellctlClient": "shellctl.client",
"ShellctlClientError": "shellctl.client",
"DEFAULT_AUTH_TOKEN_ENV": "shellctl.shared",
"DEFAULT_BASE_URL": "shellctl.shared",
"DEFAULT_BASE_URL_ENV": "shellctl.shared",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS": "shellctl.shared",
"DEFAULT_GC_INTERVAL_SECONDS": "shellctl.shared",
"DEFAULT_IDLE_FLUSH_SECONDS": "shellctl.shared",
"DEFAULT_LIST_LIMIT": "shellctl.shared",
"DEFAULT_OUTPUT_LIMIT_BYTES": "shellctl.shared",
"DEFAULT_TERMINAL_COLS": "shellctl.shared",
"DEFAULT_TERMINAL_ROWS": "shellctl.shared",
"DEFAULT_TERMINATE_GRACE_SECONDS": "shellctl.shared",
"DEFAULT_TIMEOUT_SECONDS": "shellctl.shared",
"DeleteJobResponse": "shellctl.shared",
"HealthResponse": "shellctl.shared",
"InputJobRequest": "shellctl.shared",
"JobInfo": "shellctl.shared",
"JobResult": "shellctl.shared",
"JobStatusName": "shellctl.shared",
"JobStatusView": "shellctl.shared",
"ListJobsResponse": "shellctl.shared",
"RunJobRequest": "shellctl.shared",
"TerminalSize": "shellctl.shared",
"TerminateJobRequest": "shellctl.shared",
"WaitJobRequest": "shellctl.shared",
"generate_job_id": "shellctl.shared",
"read_output_window": "shellctl.shared",
"tail_output_window": "shellctl.shared",
}
def __getattr__(name: str) -> Any:
if name not in _EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module = import_module(_EXPORTS[name])
value = getattr(module, name) # noqa: no-new-getattr lazy export proxy
globals()[name] = value
return value
def __dir__() -> list[str]:
return sorted(set(globals()) | set(__all__))
+587
View File
@@ -0,0 +1,587 @@
"""Typer CLI for network-backed shellctl commands.
Job-management commands in this module intentionally stay on the SDK side of
the boundary: they parse CLI options, call `ShellctlClient`, and render compact
JSON. That keeps `shellctl --help` and `shellctl run --help` free of FastAPI,
SQLAlchemy, tmux, and local runtime bootstrap imports.
Only `serve` lazily imports server-side modules when that subcommand is
actually invoked.
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import NoReturn
import anyio
import httpx2 as httpx
import typer
from pydantic import BaseModel, ValidationError
from shellctl.client import ShellctlClient, ShellctlClientError
from shellctl.shared.constants import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_BASE_URL,
DEFAULT_BASE_URL_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
MAX_LIST_LIMIT,
MAX_OUTPUT_LIMIT_BYTES,
)
from shellctl.shared.schemas import (
DeleteJobResponse,
HealthResponse,
JobInfo,
JobResult,
JobStatusName,
JobStatusView,
RunJobRequest,
TerminalSize,
)
cli = typer.Typer(
no_args_is_help=True,
pretty_exceptions_enable=False,
rich_markup_mode=None,
)
@cli.command("health")
def health_command(
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help="Accepted for CLI consistency but ignored because /healthz is public.",
),
) -> None:
"""Call the public health endpoint and report JSON."""
del auth_token
async def action(client: ShellctlClient) -> HealthResponse:
return await client.health()
_run_client_action(
base_url=base_url,
auth_token=None,
action=action,
emit=_emit_model,
)
@cli.command("run")
def run_command(
script: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
cwd: Path | None = typer.Option(None, "--cwd"),
env: list[str] | None = typer.Option(None, "--env"),
timeout: float = typer.Option(DEFAULT_TIMEOUT_SECONDS, "--timeout"),
output_limit: int = typer.Option(DEFAULT_OUTPUT_LIMIT_BYTES, "--output-limit"),
idle_flush_seconds: float = typer.Option(
DEFAULT_IDLE_FLUSH_SECONDS,
"--idle-flush-seconds",
),
cols: int | None = typer.Option(None, "--cols"),
rows: int | None = typer.Option(None, "--rows"),
) -> None:
"""Create a job through the running shellctl server."""
request = _build_model(
RunJobRequest,
script=script,
cwd=str(cwd) if cwd is not None else None,
env=_parse_env(env),
terminal=_terminal_size(cols=cols, rows=rows),
timeout=timeout,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
)
async def action(client: ShellctlClient) -> JobResult:
return await client.run(
request.script,
cwd=request.cwd,
env=request.env,
timeout=request.timeout,
terminal=request.terminal,
)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
action=action,
emit=_emit_model,
)
@cli.command("wait")
def wait_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
offset: int = typer.Option(..., "--offset"),
timeout: float = typer.Option(DEFAULT_TIMEOUT_SECONDS, "--timeout"),
output_limit: int = typer.Option(DEFAULT_OUTPUT_LIMIT_BYTES, "--output-limit"),
idle_flush_seconds: float = typer.Option(
DEFAULT_IDLE_FLUSH_SECONDS,
"--idle-flush-seconds",
),
) -> None:
"""Wait for incremental output, completion, truncation, or timeout."""
async def action(client: ShellctlClient) -> JobResult:
return await client.wait(job_id, offset=offset, timeout=timeout)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
action=action,
emit=_emit_model,
)
@cli.command("status")
def status_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
) -> None:
"""Materialize the current status view for one job."""
async def action(client: ShellctlClient) -> JobStatusView:
return await client.status(job_id)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_model,
)
@cli.command("list")
def list_command(
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
status: JobStatusName | None = typer.Option(None, "--status"),
limit: int = typer.Option(DEFAULT_LIST_LIMIT, "--limit", min=1, max=MAX_LIST_LIMIT),
) -> None:
"""List recent jobs, optionally filtered by lifecycle status."""
async def action(client: ShellctlClient) -> list[JobInfo]:
return await client.list_jobs(status=status, limit=limit)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_job_list,
)
@cli.command("input")
def input_command(
job_id: str = typer.Argument(...),
text: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
offset: int = typer.Option(..., "--offset"),
timeout: float = typer.Option(DEFAULT_TIMEOUT_SECONDS, "--timeout"),
output_limit: int = typer.Option(DEFAULT_OUTPUT_LIMIT_BYTES, "--output-limit"),
idle_flush_seconds: float = typer.Option(
DEFAULT_IDLE_FLUSH_SECONDS,
"--idle-flush-seconds",
),
) -> None:
"""Send text input to a running job and wait for the next result window."""
async def action(client: ShellctlClient) -> JobResult:
return await client.input(job_id, text, offset=offset, timeout=timeout)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
action=action,
emit=_emit_model,
)
@cli.command("tail")
def tail_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
output_limit: int = typer.Option(
DEFAULT_OUTPUT_LIMIT_BYTES,
"--output-limit",
min=1,
max=MAX_OUTPUT_LIMIT_BYTES,
),
) -> None:
"""Read a UTF-8-safe output tail for one job."""
async def action(client: ShellctlClient) -> JobResult:
return await client.tail(job_id)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
output_limit=output_limit,
action=action,
emit=_emit_model,
)
@cli.command("terminate")
def terminate_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
grace_seconds: float = typer.Option(
DEFAULT_TERMINATE_GRACE_SECONDS,
"--grace-seconds",
),
) -> None:
"""Terminate a job and return its materialized status."""
async def action(client: ShellctlClient) -> JobStatusView:
return await client.terminate(job_id, grace_seconds=grace_seconds)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_model,
)
@cli.command("delete")
def delete_command(
job_id: str = typer.Argument(...),
base_url: str = typer.Option(
DEFAULT_BASE_URL,
"--base-url",
envvar=DEFAULT_BASE_URL_ENV,
help="shellctl server base URL. You can also set SHELLCTL_BASE_URL.",
),
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty when the server does not require auth."
),
),
force: bool = typer.Option(False, "--force"),
grace_seconds: float = typer.Option(
DEFAULT_TERMINATE_GRACE_SECONDS,
"--grace-seconds",
),
) -> None:
"""Delete a job row and artifacts, optionally terminating first."""
async def action(client: ShellctlClient) -> DeleteJobResponse:
return await client.delete(
job_id,
force=force,
grace_seconds=grace_seconds,
)
_run_client_action(
base_url=base_url,
auth_token=auth_token,
action=action,
emit=_emit_model,
)
@cli.command("serve")
def serve_command(
listen: str = "127.0.0.1:8765",
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty to disable HTTP bearer auth."
),
),
state_dir: Path | None = None,
runtime_dir: Path | None = None,
gc_interval_seconds: float = typer.Option(
DEFAULT_GC_INTERVAL_SECONDS,
"--gc-interval-seconds",
),
gc_finished_job_retention_seconds: float = typer.Option(
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
"--gc-finished-job-retention-seconds",
),
) -> None:
"""Run the local shellctl FastAPI server via uvicorn."""
from shellctl.server.serve import (
serve_command as server_serve_command,
)
server_serve_command(
listen=listen,
auth_token=auth_token,
state_dir=state_dir,
runtime_dir=runtime_dir,
gc_interval_seconds=gc_interval_seconds,
gc_finished_job_retention_seconds=gc_finished_job_retention_seconds,
)
def main() -> None:
"""CLI entrypoint used by the console script and `python -m` invocations."""
cli()
def _parse_env(values: list[str] | None) -> dict[str, str] | None:
if not values:
return None
parsed: dict[str, str] = {}
for value in values:
if "=" not in value:
raise typer.BadParameter(
"env entries must use NAME=VALUE format",
param_hint="--env",
)
name, env_value = value.split("=", 1)
if not name:
raise typer.BadParameter(
"env names must be non-empty",
param_hint="--env",
)
parsed[name] = env_value
return parsed
def _terminal_size(*, cols: int | None, rows: int | None) -> TerminalSize | None:
if cols is None and rows is None:
return None
return _build_model(
TerminalSize,
cols=cols if cols is not None else DEFAULT_TERMINAL_COLS,
rows=rows if rows is not None else DEFAULT_TERMINAL_ROWS,
)
def _build_model[ModelT: BaseModel](model_type: type[ModelT], /, **data: object) -> ModelT:
try:
return model_type(**data)
except ValidationError as exc:
raise typer.BadParameter(_validation_error_message(exc)) from exc
async def _with_client[ResponseT](
base_url: str,
auth_token: str | None,
output_limit: int,
idle_flush_seconds: float,
action: Callable[[ShellctlClient], Awaitable[ResponseT]],
) -> ResponseT:
async with ShellctlClient(
base_url,
output_limit=output_limit,
idle_flush_seconds=idle_flush_seconds,
token=auth_token,
) as client:
return await action(client)
def _run_client_action[ResponseT](
*,
base_url: str,
auth_token: str | None,
action: Callable[[ShellctlClient], Awaitable[ResponseT]],
emit: Callable[[ResponseT], None],
output_limit: int = DEFAULT_OUTPUT_LIMIT_BYTES,
idle_flush_seconds: float = DEFAULT_IDLE_FLUSH_SECONDS,
) -> None:
try:
payload = anyio.run(
_with_client,
base_url,
auth_token,
output_limit,
idle_flush_seconds,
action,
)
except ShellctlClientError as exc:
_emit_error_and_exit(exc.code, exc.message)
except httpx.TimeoutException:
_emit_error_and_exit("request_timeout", "request timed out")
except httpx.TransportError as exc:
_emit_error_and_exit("connection_error", str(exc))
emit(payload)
def _emit_model(model: BaseModel) -> None:
typer.echo(model.model_dump_json(exclude_none=True), color=False)
def _emit_job_list(jobs: list[JobInfo]) -> None:
typer.echo(
json.dumps(
[item.model_dump(mode="json", exclude_none=True) for item in jobs],
separators=(",", ":"),
),
color=False,
)
def _emit_error_and_exit(code: str, message: str) -> NoReturn:
typer.echo(
json.dumps(
{"error": {"code": code, "message": message}},
separators=(",", ":"),
),
err=True,
color=False,
)
raise typer.Exit(code=1)
def _validation_error_message(exc: ValidationError) -> str:
detail = exc.errors(include_url=False)[0]
location = ".".join(str(part) for part in detail.get("loc", ()))
message = detail["msg"]
return f"{location}: {message}" if location else str(message)
__all__ = [
"cli",
"delete_command",
"health_command",
"input_command",
"list_command",
"main",
"run_command",
"serve_command",
"status_command",
"tail_command",
"terminate_command",
"wait_command",
]
@@ -0,0 +1,12 @@
"""Async HTTP client package for shellctl.
`shellctl.client` remains importable as before, but it is
now a package so future client helpers can live beside the main SDK class.
"""
from shellctl.client.sdk import (
ShellctlClient,
ShellctlClientError,
)
__all__ = ["ShellctlClient", "ShellctlClientError"]
+319
View File
@@ -0,0 +1,319 @@
"""Async HTTP client for the shellctl server API.
The SDK keeps transport-level knobs (`output_limit`, `idle_flush_seconds`, base
URL selection, and bearer-token handling) on the client instance so individual
method calls stay close to the network CLI's high-level workflow. Blocking shell
operations reuse the shared client, but they override the HTTP read timeout per
request so the transport does not fail before the server-side shell wait timeout
or terminate-grace budget does.
"""
from __future__ import annotations
import os
from typing import Any
import httpx2 as httpx
from shellctl.shared.constants import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_BASE_URL,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
)
from shellctl.shared.schemas import (
DeleteJobResponse,
HealthResponse,
JobInfo,
JobResult,
JobStatusView,
ListJobsResponse,
RunJobRequest,
TerminalSize,
)
class ShellctlClientError(RuntimeError):
"""Raised for API-declared failures and response decode/shape problems.
`ShellctlClient` raises this error when the server returns a structured
error payload, and also when an otherwise successful HTTP response contains
invalid JSON or a top-level payload shape that does not match the SDK
contract. Transport and timeout failures remain raw `httpx2` exceptions so
library callers can decide how to handle network-layer failures.
"""
def __init__(self, status_code: int, code: str, message: str) -> None:
super().__init__(f"{code} ({status_code}): {message}")
self.status_code = status_code
self.code = code
self.message = message
class ShellctlClient:
"""Thin async SDK for the shellctl HTTP API.
The client owns a reusable `httpx.AsyncClient` unless one is injected via the
`client` argument. Callers can therefore either keep one instance for a full
workflow or treat it as an async context manager. Injected clients keep their
original lifecycle; `close()` only closes clients that this SDK created.
"""
def __init__(
self,
base_url: str = DEFAULT_BASE_URL,
*,
output_limit: int = DEFAULT_OUTPUT_LIMIT_BYTES,
idle_flush_seconds: float = DEFAULT_IDLE_FLUSH_SECONDS,
token: str | None = None,
client: httpx.AsyncClient | None = None,
transport: httpx.AsyncBaseTransport | None = None,
request_timeout_grace_seconds: float = 10.0,
) -> None:
self.base_url = base_url.rstrip("/")
self.output_limit = output_limit
self.idle_flush_seconds = idle_flush_seconds
self.request_timeout_grace_seconds = request_timeout_grace_seconds
self.token = token if token is not None else os.environ.get(DEFAULT_AUTH_TOKEN_ENV)
self._owns_client = client is None
self._client = client or httpx.AsyncClient(
base_url=self.base_url,
follow_redirects=True,
timeout=httpx.Timeout(DEFAULT_TIMEOUT_SECONDS, connect=DEFAULT_TIMEOUT_SECONDS),
transport=transport,
)
async def __aenter__(self) -> ShellctlClient:
return self
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
await self.close()
async def close(self) -> None:
"""Close the underlying HTTP client if this SDK instance owns it."""
if self._owns_client:
await self._client.aclose()
def _wait_request_timeout(self, timeout: float) -> httpx.Timeout:
"""Return a request timeout for blocking shell calls.
The shellctl server enforces the payload timeout, while the SDK keeps a
small HTTP read-timeout grace so the transport can wait slightly longer
for that response without loosening connect/write/pool timeouts.
"""
return httpx.Timeout(
connect=DEFAULT_TIMEOUT_SECONDS,
read=timeout + self.request_timeout_grace_seconds,
write=DEFAULT_TIMEOUT_SECONDS,
pool=DEFAULT_TIMEOUT_SECONDS,
)
def _terminate_request_timeout(self, grace_seconds: float | None = None) -> httpx.Timeout:
"""Return a request timeout for terminate-style calls.
`terminate()` and forced `delete()` block until the server finishes the
terminate grace window, so their HTTP read timeout must cover that
business wait budget even when the request relies on the API default.
"""
effective_grace_seconds = DEFAULT_TERMINATE_GRACE_SECONDS if grace_seconds is None else grace_seconds
return self._wait_request_timeout(effective_grace_seconds)
async def health(self) -> HealthResponse:
"""Call the public health endpoint and decode it as `HealthResponse`."""
return HealthResponse.model_validate(await self.healthz())
async def healthz(self) -> dict[str, Any]:
"""Call the public health endpoint without requiring auth."""
response = await self._client.get("/healthz")
return self._decode_response(response)
async def run(
self,
script: str,
*,
cwd: str | None = None,
env: dict[str, str] | None = None,
timeout: float = DEFAULT_TIMEOUT_SECONDS,
terminal: TerminalSize | None = None,
) -> JobResult:
"""Create a new job and wait for initial output or completion.
`cwd` and `env` preset the script's working directory and environment
overlay on the server side.
"""
payload = RunJobRequest(
script=script,
cwd=cwd,
env=env,
terminal=terminal,
timeout=timeout,
output_limit=self.output_limit,
idle_flush_seconds=self.idle_flush_seconds,
)
response = await self._client.post(
"/v1/jobs/run",
json=payload.model_dump(mode="json", exclude_none=True),
headers=self._auth_headers(),
timeout=self._wait_request_timeout(timeout),
)
return JobResult.model_validate(self._decode_response(response))
async def wait(
self,
job_id: str,
*,
offset: int,
timeout: float = DEFAULT_TIMEOUT_SECONDS,
) -> JobResult:
"""Wait for incremental output, completion, truncation, or timeout."""
response = await self._client.post(
f"/v1/jobs/{job_id}/wait",
json={
"offset": offset,
"timeout": timeout,
"output_limit": self.output_limit,
"idle_flush_seconds": self.idle_flush_seconds,
},
headers=self._auth_headers(),
timeout=self._wait_request_timeout(timeout),
)
return JobResult.model_validate(self._decode_response(response))
async def status(self, job_id: str) -> JobStatusView:
"""Fetch the materialized status view for one job."""
response = await self._client.get(
f"/v1/jobs/{job_id}",
headers=self._auth_headers(),
)
return JobStatusView.model_validate(self._decode_response(response))
async def list_jobs(
self,
*,
status: str | None = None,
limit: int = DEFAULT_LIST_LIMIT,
) -> list[JobInfo]:
"""List recent jobs, optionally filtered by lifecycle status."""
params: dict[str, Any] = {"limit": limit}
if status is not None:
params["status"] = status
response = await self._client.get(
"/v1/jobs",
params=params,
headers=self._auth_headers(),
)
payload = ListJobsResponse.model_validate(self._decode_response(response))
return payload.jobs
async def input(
self,
job_id: str,
text: str,
*,
offset: int,
timeout: float = DEFAULT_TIMEOUT_SECONDS,
) -> JobResult:
"""Send text input to a running job and then wait like `wait()`."""
response = await self._client.post(
f"/v1/jobs/{job_id}/input",
json={
"text": text,
"offset": offset,
"timeout": timeout,
"output_limit": self.output_limit,
"idle_flush_seconds": self.idle_flush_seconds,
},
headers=self._auth_headers(),
timeout=self._wait_request_timeout(timeout),
)
return JobResult.model_validate(self._decode_response(response))
async def tail(self, job_id: str) -> JobResult:
"""Fetch an immediate UTF-8-safe tail snapshot for a job."""
response = await self._client.get(
f"/v1/jobs/{job_id}/log/tail",
params={"output_limit": self.output_limit},
headers=self._auth_headers(),
)
return JobResult.model_validate(self._decode_response(response))
async def terminate(
self,
job_id: str,
grace_seconds: float = DEFAULT_TERMINATE_GRACE_SECONDS,
) -> JobStatusView:
"""Terminate a job, waiting long enough for the grace window to finish."""
response = await self._client.post(
f"/v1/jobs/{job_id}/terminate",
json={"grace_seconds": grace_seconds},
headers=self._auth_headers(),
timeout=self._terminate_request_timeout(grace_seconds),
)
return JobStatusView.model_validate(self._decode_response(response))
async def delete(
self,
job_id: str,
*,
force: bool = False,
grace_seconds: float | None = None,
) -> DeleteJobResponse:
"""Delete job artifacts, optionally waiting for forced termination first."""
params: dict[str, Any] = {"force": str(force).lower()}
if grace_seconds is not None:
params["grace_seconds"] = grace_seconds
request_kwargs: dict[str, Any] = {
"params": params,
"headers": self._auth_headers(),
}
if force:
request_kwargs["timeout"] = self._terminate_request_timeout(grace_seconds)
response = await self._client.delete(
f"/v1/jobs/{job_id}",
**request_kwargs,
)
return DeleteJobResponse.model_validate(self._decode_response(response))
def _auth_headers(self) -> dict[str, str]:
if not self.token:
return {}
return {"Authorization": f"Bearer {self.token}"}
def _decode_response(self, response: httpx.Response) -> dict[str, Any]:
try:
payload = response.json()
except ValueError as exc: # pragma: no cover - network/proxy corruption
raise ShellctlClientError(response.status_code, "invalid_json", response.text) from exc
if response.is_error:
error = payload.get("error") if isinstance(payload, dict) else None
if isinstance(error, dict):
code = str(error.get("code", "request_failed"))
message = str(error.get("message", response.text))
else:
code = "request_failed"
message = response.text
raise ShellctlClientError(response.status_code, code, message)
if not isinstance(payload, dict):
raise ShellctlClientError(response.status_code, "invalid_payload", response.text)
return payload
__all__ = ["ShellctlClient", "ShellctlClientError"]
+1
View File
@@ -0,0 +1 @@
@@ -0,0 +1,58 @@
"""shellctl server package.
The server stack sits behind lazy exports so importing the network CLI does not
pull in FastAPI, SQLAlchemy, tmux, or the local service runtime unless a
server-side symbol is actually used.
"""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from shellctl.server.api import create_app
from shellctl.server.cli import (
cli,
main,
)
from shellctl.server.config import ShellctlConfig
from shellctl.server.db import JobRow
from shellctl.server.errors import ShellctlServerError
from shellctl.server.serve import serve_command
from shellctl.server.service import ShellctlService
__all__ = [
"JobRow",
"ShellctlConfig",
"ShellctlServerError",
"ShellctlService",
"cli",
"create_app",
"main",
"serve_command",
]
_EXPORTS = {
"JobRow": "shellctl.server.db",
"ShellctlConfig": "shellctl.server.config",
"ShellctlServerError": "shellctl.server.errors",
"ShellctlService": "shellctl.server.service",
"cli": "shellctl.server.cli",
"create_app": "shellctl.server.api",
"main": "shellctl.server.cli",
"serve_command": "shellctl.server.serve",
}
def __getattr__(name: str) -> Any:
if name not in _EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module = import_module(_EXPORTS[name])
value = getattr(module, name) # noqa: no-new-getattr lazy export proxy
globals()[name] = value
return value
def __dir__() -> list[str]:
return sorted(set(globals()) | set(__all__))
@@ -0,0 +1,6 @@
"""Support `python -m shellctl.server`."""
from shellctl.cli import main
if __name__ == "__main__":
main()
+198
View File
@@ -0,0 +1,198 @@
"""FastAPI wiring for shellctl server endpoints."""
from __future__ import annotations
from contextlib import asynccontextmanager
from typing import Annotated, cast
from fastapi import Depends, FastAPI, Header, Query, Request
from fastapi.responses import JSONResponse
from shellctl.server.config import ShellctlConfig
from shellctl.server.errors import ShellctlServerError
from shellctl.server.service import ShellctlService
from shellctl.shared.constants import (
DEFAULT_HEALTH_STATUS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINATE_GRACE_SECONDS,
MAX_LIST_LIMIT,
MAX_OUTPUT_LIMIT_BYTES,
)
from shellctl.shared.schemas import (
DeleteJobResponse,
ErrorDetail,
ErrorResponse,
HealthResponse,
InputJobRequest,
JobResult,
JobStatusName,
JobStatusView,
ListJobsResponse,
RunJobRequest,
TerminateJobRequest,
WaitJobRequest,
)
def create_app(
config: ShellctlConfig | None = None,
*,
service: ShellctlService | None = None,
) -> FastAPI:
"""Create the FastAPI application used by `shellctl serve`."""
resolved_config = config or ShellctlConfig()
resolved_service = service or ShellctlService(resolved_config)
@asynccontextmanager
async def lifespan(_app: FastAPI):
await resolved_service.initialize()
resolved_service.start_background_gc()
resolved_service.start_background_pipe_monitor()
try:
yield
finally:
await resolved_service.shutdown()
app = FastAPI(title="shellctl", version="0.1.0", lifespan=lifespan)
app.state.shellctl_service = resolved_service
@app.exception_handler(ShellctlServerError)
async def handle_shellctl_error(
_request: Request,
exc: ShellctlServerError,
) -> JSONResponse:
return JSONResponse(
status_code=exc.status_code,
content=ErrorResponse(error=ErrorDetail(code=exc.code, message=exc.message)).model_dump(mode="json"),
)
@app.exception_handler(RuntimeError)
async def handle_runtime_error(_request: Request, exc: RuntimeError) -> JSONResponse:
return JSONResponse(
status_code=500,
content=ErrorResponse(
error=ErrorDetail(
code="internal_error",
message=str(exc) or "internal server error",
)
).model_dump(mode="json"),
)
def get_service() -> ShellctlService:
return cast(ShellctlService, app.state.shellctl_service)
def verify_auth(
authorization: Annotated[str | None, Header()] = None,
) -> None:
token = resolved_config.auth_token
if token is None:
return
expected = f"Bearer {token}"
if authorization != expected:
raise ShellctlServerError(401, "unauthorized", "Missing or invalid bearer token")
@app.get("/healthz", response_model=HealthResponse)
async def healthz() -> HealthResponse:
return HealthResponse(status=DEFAULT_HEALTH_STATUS)
@app.post(
"/v1/jobs/run",
response_model=JobResult,
dependencies=[Depends(verify_auth)],
)
async def run_job(
payload: RunJobRequest,
svc: ShellctlService = Depends(get_service),
) -> JobResult:
return await svc.run_job(payload)
@app.post(
"/v1/jobs/{job_id}/wait",
response_model=JobResult,
dependencies=[Depends(verify_auth)],
)
async def wait_job(
job_id: str,
payload: WaitJobRequest,
svc: ShellctlService = Depends(get_service),
) -> JobResult:
return await svc.wait_job(job_id, payload)
@app.get(
"/v1/jobs/{job_id}/log/tail",
response_model=JobResult,
dependencies=[Depends(verify_auth)],
)
async def tail_job(
job_id: str,
output_limit: Annotated[int, Query(ge=1, le=MAX_OUTPUT_LIMIT_BYTES)] = DEFAULT_OUTPUT_LIMIT_BYTES,
svc: ShellctlService = Depends(get_service),
) -> JobResult:
return await svc.tail_job(job_id, output_limit=output_limit)
@app.get(
"/v1/jobs/{job_id}",
response_model=JobStatusView,
dependencies=[Depends(verify_auth)],
)
async def job_status(
job_id: str,
svc: ShellctlService = Depends(get_service),
) -> JobStatusView:
return await svc.get_job_status(job_id)
@app.get(
"/v1/jobs",
response_model=ListJobsResponse,
dependencies=[Depends(verify_auth)],
)
async def list_jobs(
status: Annotated[JobStatusName | None, Query()] = None,
limit: Annotated[int, Query(ge=1, le=MAX_LIST_LIMIT)] = DEFAULT_LIST_LIMIT,
svc: ShellctlService = Depends(get_service),
) -> ListJobsResponse:
return await svc.list_jobs(status=status, limit=limit)
@app.post(
"/v1/jobs/{job_id}/input",
response_model=JobResult,
dependencies=[Depends(verify_auth)],
)
async def input_job(
job_id: str,
payload: InputJobRequest,
svc: ShellctlService = Depends(get_service),
) -> JobResult:
return await svc.send_input(job_id, payload)
@app.post(
"/v1/jobs/{job_id}/terminate",
response_model=JobStatusView,
dependencies=[Depends(verify_auth)],
)
async def terminate_job(
job_id: str,
payload: TerminateJobRequest,
svc: ShellctlService = Depends(get_service),
) -> JobStatusView:
return await svc.terminate_job(job_id, payload)
@app.delete(
"/v1/jobs/{job_id}",
response_model=DeleteJobResponse,
dependencies=[Depends(verify_auth)],
)
async def delete_job(
job_id: str,
force: bool = False,
grace_seconds: float = DEFAULT_TERMINATE_GRACE_SECONDS,
svc: ShellctlService = Depends(get_service),
) -> DeleteJobResponse:
return await svc.delete_job(job_id, force=force, grace_seconds=grace_seconds)
return app
__all__ = ["create_app"]
@@ -0,0 +1,62 @@
"""Per-job artifact names used by shellctl server/runtime code.
Normal job completion is coordinated through small marker files inside each
`jobs/<job_id>/` directory so the tmux output-pipe finalizer can publish the
SQLite `exited(exit_code, ended_at)` state only after PTY output is fully
drained into `output.log`. The same artifact directory also stores the request's
environment overlay so the runner can merge arbitrary key/value pairs without
shell-escaping them into the generated script. Separate failure markers and a
dedicated `pipe-error.log` stderr capture keep startup diagnostics available
when the sanitizer never reaches its ready-file handshake.
"""
from __future__ import annotations
from pathlib import Path
RUNNER_EXIT_CODE_FILENAME = ".runner-exit-code"
RUNNER_ENDED_AT_FILENAME = ".runner-ended-at"
JOB_ENV_FILENAME = ".job-env.json"
PIPE_DRAINED_FILENAME = ".pipe-drained"
PIPE_FAILED_FILENAME = ".pipe-failed"
PIPE_ERROR_LOG_FILENAME = "pipe-error.log"
def runner_exit_code_path(job_dir: Path) -> Path:
return job_dir / RUNNER_EXIT_CODE_FILENAME
def runner_ended_at_path(job_dir: Path) -> Path:
return job_dir / RUNNER_ENDED_AT_FILENAME
def job_env_path(job_dir: Path) -> Path:
return job_dir / JOB_ENV_FILENAME
def pipe_drained_path(job_dir: Path) -> Path:
return job_dir / PIPE_DRAINED_FILENAME
def pipe_failed_path(job_dir: Path) -> Path:
return job_dir / PIPE_FAILED_FILENAME
def pipe_error_log_path(job_dir: Path) -> Path:
return job_dir / PIPE_ERROR_LOG_FILENAME
__all__ = [
"JOB_ENV_FILENAME",
"PIPE_DRAINED_FILENAME",
"PIPE_ERROR_LOG_FILENAME",
"PIPE_FAILED_FILENAME",
"RUNNER_ENDED_AT_FILENAME",
"RUNNER_EXIT_CODE_FILENAME",
"job_env_path",
"pipe_drained_path",
"pipe_error_log_path",
"pipe_failed_path",
"runner_ended_at_path",
"runner_exit_code_path",
]
+17
View File
@@ -0,0 +1,17 @@
"""Compatibility shim for historical `shellctl.server.cli` imports.
Job-management commands now live in `shellctl.cli`. This
module is intentionally a thin legacy re-export for callers that still import
CLI symbols from `shellctl.server.cli`.
"""
from __future__ import annotations
from shellctl.cli import cli, main
from shellctl.server.serve import serve_command
__all__ = [
"cli",
"main",
"serve_command",
]
+105
View File
@@ -0,0 +1,105 @@
"""Configuration objects for shellctl server/runtime modules."""
from __future__ import annotations
import os
from dataclasses import dataclass, field
from pathlib import Path
from typing import cast
from shellctl.shared.constants import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
MAX_LIST_LIMIT,
MAX_OUTPUT_LIMIT_BYTES,
MAX_WAIT_TIMEOUT_SECONDS,
)
from shellctl.shared.runtime import (
default_runtime_dir,
default_state_dir,
)
from shellctl_runtime.paths import DEFAULT_SQLITE_BUSY_TIMEOUT_MS
@dataclass(slots=True, frozen=True)
class ShellctlConfig:
"""Runtime configuration for the shellctl service and CLI.
The tmux subprocess hooks use dedicated console scripts instead of
`python -m shellctl...` entrypoints so every job does not pay
the shellctl client/server import cost just to sanitize PTY bytes or record
an exit row. `sqlite_busy_timeout_ms` still applies to the out-of-process
`runner-exit` callback, so non-default deployments keep one SQLite timeout
policy across the service and tmux finalizer.
Bearer auth is opt-in: if the explicit `auth_token` and the fallback
`SHELLCTL_AUTH_TOKEN` environment variable are both missing or empty,
`shellctl serve` accepts requests without checking an Authorization header.
"""
listen: str = "127.0.0.1:8765"
auth_token: str | None = None
state_dir: Path = field(default_factory=default_state_dir)
runtime_dir: Path | None = None
gc_interval_seconds: float = DEFAULT_GC_INTERVAL_SECONDS
gc_finished_job_retention_seconds: float = DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS
default_timeout_seconds: float = DEFAULT_TIMEOUT_SECONDS
max_wait_timeout_seconds: float = MAX_WAIT_TIMEOUT_SECONDS
idle_flush_seconds: float = DEFAULT_IDLE_FLUSH_SECONDS
default_cwd: Path = field(default_factory=Path.home)
default_terminal_cols: int = DEFAULT_TERMINAL_COLS
default_terminal_rows: int = DEFAULT_TERMINAL_ROWS
default_list_limit: int = DEFAULT_LIST_LIMIT
max_list_limit: int = MAX_LIST_LIMIT
default_output_limit_bytes: int = DEFAULT_OUTPUT_LIMIT_BYTES
max_output_limit_bytes: int = MAX_OUTPUT_LIMIT_BYTES
default_terminate_grace_seconds: float = DEFAULT_TERMINATE_GRACE_SECONDS
poll_interval_seconds: float = 0.05
pipe_monitor_interval_seconds: float = 1.0
pipe_ready_timeout_seconds: float = 10.0
sqlite_busy_timeout_ms: int = DEFAULT_SQLITE_BUSY_TIMEOUT_MS
sanitize_pty_command: tuple[str, ...] = ("shellctl-sanitize-pty",)
runner_exit_command: tuple[str, ...] = ("shellctl-runner-exit",)
def __post_init__(self) -> None:
if self.runtime_dir is None:
object.__setattr__(self, "runtime_dir", default_runtime_dir(self.state_dir))
token = self.auth_token
if token is None:
token = os.environ.get(DEFAULT_AUTH_TOKEN_ENV)
if not token:
token = None
object.__setattr__(self, "auth_token", token)
@property
def jobs_dir(self) -> Path:
return self.state_dir / "jobs"
@property
def db_path(self) -> Path:
return self.state_dir / "shellctl.db"
@property
def database_url(self) -> str:
return f"sqlite+aiosqlite:///{self.db_path}"
@property
def tmux_socket(self) -> Path:
runtime_dir = cast(Path, self.runtime_dir)
return runtime_dir / "tmux.sock"
@property
def runner_path(self) -> Path:
runtime_dir = cast(Path, self.runtime_dir)
return runtime_dir / "bin" / "shellctl-runner"
__all__ = ["ShellctlConfig"]
+52
View File
@@ -0,0 +1,52 @@
"""SQLite models and engine helpers for shellctl."""
from __future__ import annotations
from typing import Any, cast
from sqlalchemy import event
from sqlalchemy.ext.asyncio import AsyncEngine
from sqlmodel import Field, SQLModel
class JobRow(SQLModel, table=True):
"""SQLite source-of-truth row for shellctl jobs.
The table intentionally combines immutable metadata, mutable lifecycle state,
and exit facts so state transitions can be expressed as single conditional
`UPDATE` statements without synchronizing separate records.
"""
__tablename__ = cast(Any, "jobs")
job_id: str = Field(primary_key=True)
script_path: str
output_path: str
cwd: str
terminal_cols: int
terminal_rows: int
status: str = Field(index=True)
session_name: str
pane_target: str
exit_code: int | None = Field(default=None, nullable=True)
reason: str | None = Field(default=None, nullable=True)
message: str | None = Field(default=None, nullable=True)
created_at: str = Field(index=True)
started_at: str | None = Field(default=None, nullable=True)
ended_at: str | None = Field(default=None, nullable=True, index=True)
updated_at: str
def configure_sqlite_engine(engine: AsyncEngine, *, busy_timeout_ms: int) -> None:
"""Install SQLite pragmas required by the proposal's concurrency model."""
@event.listens_for(engine.sync_engine, "connect")
def _set_pragmas(dbapi_connection: Any, _connection_record: Any) -> None:
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute("PRAGMA foreign_keys=ON")
cursor.execute(f"PRAGMA busy_timeout={busy_timeout_ms}")
cursor.close()
__all__ = ["JobRow", "configure_sqlite_engine"]
+14
View File
@@ -0,0 +1,14 @@
"""Shared exception types for shellctl server-side modules."""
class ShellctlServerError(RuntimeError):
"""Structured server-side error that maps directly to API responses."""
def __init__(self, status_code: int, code: str, message: str) -> None:
super().__init__(message)
self.status_code = status_code
self.code = code
self.message = message
__all__ = ["ShellctlServerError"]
+103
View File
@@ -0,0 +1,103 @@
"""Server-side `shellctl serve` implementation.
This module stays separate from the top-level CLI so ordinary client commands
do not import the FastAPI or uvicorn stack. `serve_command()` performs those
imports lazily because only the long-running server path needs them.
"""
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING, Any
import typer
from shellctl.shared.constants import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
)
from shellctl.shared.runtime import (
default_state_dir,
)
if TYPE_CHECKING:
from shellctl.server.config import ShellctlConfig
def serve_command(
listen: str = "127.0.0.1:8765",
auth_token: str | None = typer.Option(
None,
"--auth-token",
envvar=DEFAULT_AUTH_TOKEN_ENV,
help=(
"Bearer token value. You can also set SHELLCTL_AUTH_TOKEN. "
"Leave it unset or empty to disable HTTP bearer auth."
),
),
state_dir: Path | None = None,
runtime_dir: Path | None = None,
gc_interval_seconds: float = DEFAULT_GC_INTERVAL_SECONDS,
gc_finished_job_retention_seconds: float = DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
) -> None:
"""Build `ShellctlConfig` from CLI inputs and run the local HTTP server.
Args:
listen: Host/port pair for the uvicorn listener in `host:port` form.
auth_token: Optional bearer token value. An explicit empty string
disables HTTP auth, an explicit non-empty token enables it, and an
omitted/`None` value may still resolve from `SHELLCTL_AUTH_TOKEN`
through the Typer env var binding or `ShellctlConfig` fallback.
state_dir: Persistent shellctl state directory; defaults to the shared
XDG-style state path when omitted.
runtime_dir: Optional runtime directory override for tmux/runtime
artifacts.
gc_interval_seconds: Background GC wake-up cadence for finished jobs.
gc_finished_job_retention_seconds: Retention window before finished jobs
are eligible for GC.
This entrypoint is the only CLI path that should pull in the FastAPI and
uvicorn stack. It parses the listener, constructs `ShellctlConfig`, and
then hands the configured app to uvicorn.
"""
from shellctl.server.config import ShellctlConfig
host, port = _parse_listen(listen)
config = ShellctlConfig(
listen=listen,
auth_token=auth_token,
state_dir=state_dir or default_state_dir(),
runtime_dir=runtime_dir,
gc_interval_seconds=gc_interval_seconds,
gc_finished_job_retention_seconds=gc_finished_job_retention_seconds,
)
_uvicorn_run(_create_app(config), host=host, port=port, log_level="info")
def _parse_listen(value: str) -> tuple[str, int]:
if ":" not in value:
raise typer.BadParameter("listen must use host:port format")
host, raw_port = value.rsplit(":", 1)
host = host.strip("[]")
try:
port = int(raw_port)
except ValueError as exc:
raise typer.BadParameter(f"invalid port: {raw_port}") from exc
return host, port
def _create_app(config: ShellctlConfig) -> Any:
from shellctl.server.api import create_app
return create_app(config)
def _uvicorn_run(app: Any, *, host: str, port: int, log_level: str) -> None:
import uvicorn
uvicorn.run(app, host=host, port=port, log_level=log_level)
__all__ = ["serve_command"]
File diff suppressed because it is too large Load Diff
+343
View File
@@ -0,0 +1,343 @@
"""tmux control layer for shellctl jobs.
The rest of shellctl treats this module as the only place that knows tmux CLI
command shapes. Tests can replace `TmuxControllerProtocol` with a fake to
exercise SQLite/output semantics without depending on a local tmux daemon.
"""
from __future__ import annotations
import os
import shlex
import subprocess
import tempfile
from pathlib import Path
from typing import Protocol, cast
import anyio
from shellctl.server.artifacts import (
pipe_drained_path,
pipe_error_log_path,
pipe_failed_path,
runner_ended_at_path,
runner_exit_code_path,
)
from shellctl.server.config import ShellctlConfig
from shellctl.server.errors import ShellctlServerError
from shellctl.shared.runtime import (
job_pane_target,
job_session_name,
)
from shellctl.shared.schemas import TerminalSize
class TmuxControllerProtocol(Protocol):
"""Protocol used by `ShellctlService` for tmux interactions."""
async def start_server(self) -> None: ...
async def list_sessions(self) -> set[str]: ...
async def session_exists(self, session_name: str) -> bool: ...
async def is_output_pipe_active(self, *, job_id: str) -> bool | None: ...
async def create_job_session(
self,
*,
job_id: str,
job_dir: Path,
cwd: Path,
terminal: TerminalSize,
) -> None: ...
async def enable_output_pipe(self, *, job_id: str, job_dir: Path, ready_file: Path) -> None: ...
async def send_input(self, *, job_id: str, text: str) -> None: ...
async def send_interrupt(self, *, job_id: str) -> None: ...
async def cleanup_session(self, *, job_id: str) -> None: ...
class TmuxController:
"""Best-effort wrapper around a dedicated tmux socket.
The controller always clears `TMUX` from the child environment and always
passes `-S <socket>` so shellctl sessions stay isolated from the user's
default tmux server.
"""
def __init__(self, config: ShellctlConfig) -> None:
self._config = config
async def start_server(self) -> None:
await self._run_tmux("start-server")
async def list_sessions(self) -> set[str]:
result = await self._run_tmux("list-sessions", "-F", "#{session_name}", check=False)
if result.returncode != 0:
stderr = result.stderr.decode("utf-8", errors="replace")
if _tmux_target_missing(stderr):
return set()
raise ShellctlServerError(500, "tmux_error", stderr.strip() or "tmux list-sessions failed")
output = result.stdout.decode("utf-8", errors="replace")
return {line.strip() for line in output.splitlines() if line.strip()}
async def session_exists(self, session_name: str) -> bool:
return session_name in await self.list_sessions()
async def is_output_pipe_active(self, *, job_id: str) -> bool | None:
result = await self._run_tmux(
"display-message",
"-p",
"-t",
job_pane_target(job_id),
"#{pane_pipe}",
check=False,
)
if result.returncode != 0:
stderr = result.stderr.decode("utf-8", errors="replace")
if _tmux_target_missing(stderr):
return None
raise ShellctlServerError(
500,
"tmux_error",
stderr.strip() or f"Failed to inspect output pipe for {job_id}",
)
return result.stdout.decode("utf-8", errors="replace").strip() == "1"
async def create_job_session(
self,
*,
job_id: str,
job_dir: Path,
cwd: Path,
terminal: TerminalSize,
) -> None:
runner_command = self._shell_join(
[
str(self._config.runner_path),
str(job_dir),
job_id,
str(cwd),
]
)
result = await self._run_tmux(
"-f",
"/dev/null",
"new-session",
"-d",
"-s",
job_session_name(job_id),
"-x",
str(terminal.cols),
"-y",
str(terminal.rows),
runner_command,
check=False,
)
if result.returncode != 0:
raise ShellctlServerError(
500,
"tmux_new_session_failed",
result.stderr.decode("utf-8", errors="replace").strip()
or f"Failed to create tmux session for {job_id}",
)
async def enable_output_pipe(self, *, job_id: str, job_dir: Path, ready_file: Path) -> None:
output_command = self._pipe_command_source(
job_id=job_id,
job_dir=job_dir,
ready_file=ready_file,
)
result = await self._run_tmux(
"pipe-pane",
"-o",
"-t",
job_pane_target(job_id),
output_command,
check=False,
)
if result.returncode != 0:
raise ShellctlServerError(
500,
"pipe_failed",
result.stderr.decode("utf-8", errors="replace").strip() or f"Failed to attach output pipe for {job_id}",
)
def _pipe_command_source(self, *, job_id: str, job_dir: Path, ready_file: Path) -> str:
"""Build the tmux `pipe-pane` command that drains and finalizes output.
For normal exits, the runner now records completion metadata into job
artifacts and the pipe finalizer commits `runner-exit` only after
the lightweight sanitizer reaches EOF and flushes `output.log`
successfully. Sanitizer stderr is captured into `pipe-error.log` so
startup timeouts can distinguish slow imports from subprocess crashes.
If the follow-up `runner-exit` write fails, the drain marker remains in
place and stderr is appended to the same log with an explicit status
line. The pipe still exits with the sanitizer status so a drained job is
not misclassified as `pipe_failed` before reconciliation can recover the
SQLite write from the drained artifacts.
"""
sanitize_command = self._shell_join(
(
*self._config.sanitize_pty_command,
"--ready-file",
str(ready_file),
)
)
runner_exit_command = self._shell_join(
(
*self._config.runner_exit_command,
"--state-dir",
str(self._config.state_dir),
"--job-id",
job_id,
"--sqlite-busy-timeout-ms",
str(self._config.sqlite_busy_timeout_ms),
)
)
output_path = shlex.quote(str(job_dir / "output.log"))
drained_path = shlex.quote(str(pipe_drained_path(job_dir)))
error_log_path = shlex.quote(str(pipe_error_log_path(job_dir)))
failed_path = shlex.quote(str(pipe_failed_path(job_dir)))
exit_code_path = shlex.quote(str(runner_exit_code_path(job_dir)))
ended_at_path = shlex.quote(str(runner_ended_at_path(job_dir)))
return " ; ".join(
[
f"{sanitize_command} >> {output_path} 2> {error_log_path}",
"sanitize_status=$?",
"runner_exit_status=0",
(
'if [ "$sanitize_status" -eq 0 ]; then '
f": > {drained_path}; "
f"if [ -s {exit_code_path} ] && [ -s {ended_at_path} ]; then "
f'{runner_exit_command} --exit-code "$(cat {exit_code_path})" '
f'--ended-at "$(cat {ended_at_path})" 2>> {error_log_path}; '
"runner_exit_status=$?; "
'if [ "$runner_exit_status" -ne 0 ]; then '
f"printf 'runner-exit failed with status %s\\n' \"$runner_exit_status\" >> {error_log_path}; "
"fi; fi; "
f"else : > {failed_path}; fi"
),
'if [ "$sanitize_status" -ne 0 ]; then exit "$sanitize_status"; fi',
'exit "$sanitize_status"',
]
)
async def send_input(self, *, job_id: str, text: str) -> None:
buffer_name = f"shellctl-in-{job_id}"
runtime_dir = cast(Path, self._config.runtime_dir)
runtime_dir.mkdir(parents=True, exist_ok=True)
fd, tmp_name = tempfile.mkstemp(prefix=f"shellctl-input-{job_id}-", dir=runtime_dir)
tmp_path = Path(tmp_name)
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
handle.write(text)
load_result = await self._run_tmux(
"load-buffer",
"-b",
buffer_name,
str(tmp_path),
check=False,
)
if load_result.returncode != 0:
stderr = load_result.stderr.decode("utf-8", errors="replace").strip()
if _tmux_target_missing(stderr):
raise ShellctlServerError(
409,
"tmux_target_missing",
stderr or f"The tmux pane for {job_id} is no longer available",
)
raise ShellctlServerError(
500,
"tmux_input_failed",
stderr or f"Failed to load input buffer for {job_id}",
)
paste_result = await self._run_tmux(
"paste-buffer",
"-t",
job_pane_target(job_id),
"-b",
buffer_name,
check=False,
)
if paste_result.returncode != 0:
stderr = paste_result.stderr.decode("utf-8", errors="replace").strip()
if _tmux_target_missing(stderr):
raise ShellctlServerError(
409,
"tmux_target_missing",
stderr or f"The tmux pane for {job_id} is no longer available",
)
raise ShellctlServerError(
500,
"tmux_input_failed",
stderr or f"Failed to paste input buffer for {job_id}",
)
finally:
await self._run_tmux("delete-buffer", "-b", buffer_name, check=False)
tmp_path.unlink(missing_ok=True)
async def send_interrupt(self, *, job_id: str) -> None:
await self._run_tmux(
"send-keys",
"-t",
job_pane_target(job_id),
"C-c",
check=False,
)
async def cleanup_session(self, *, job_id: str) -> None:
await self._run_tmux(
"kill-session",
"-t",
job_session_name(job_id),
check=False,
)
async def _run_tmux(self, *args: str, check: bool = True) -> subprocess.CompletedProcess[bytes]:
env = dict(os.environ)
env.pop("TMUX", None)
try:
result = await anyio.run_process(
["tmux", "-S", str(self._config.tmux_socket), *args],
env=env,
check=False,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
except FileNotFoundError as exc:
raise ShellctlServerError(500, "tmux_not_installed", "tmux executable was not found") from exc
if check and result.returncode != 0:
raise ShellctlServerError(
500,
"tmux_error",
result.stderr.decode("utf-8", errors="replace").strip() or "tmux command failed",
)
return result
@staticmethod
def _shell_join(parts: tuple[str, ...] | list[str]) -> str:
return " ".join(shlex.quote(part) for part in parts)
def _tmux_target_missing(stderr: str) -> bool:
normalized = stderr.lower()
return (
"can't find pane" in normalized
or "can't find session" in normalized
or "no server running" in normalized
or "failed to connect" in normalized
or "server exited unexpectedly" in normalized
)
__all__ = [
"TmuxController",
"TmuxControllerProtocol",
"_tmux_target_missing",
]
+182
View File
@@ -0,0 +1,182 @@
"""Shared shellctl transport/runtime helpers.
This package preserves the historical import surface while keeping the package
root lazy. Lightweight callers can import concrete submodules such as
`shared.runtime` without eagerly importing the pydantic schema layer, output
or helpers outside the shared compatibility surface.
"""
from __future__ import annotations
from importlib import import_module
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from shellctl.shared.constants import (
DEFAULT_AUTH_TOKEN_ENV,
DEFAULT_BASE_URL,
DEFAULT_BASE_URL_ENV,
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS,
DEFAULT_GC_INTERVAL_SECONDS,
DEFAULT_HEALTH_STATUS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_LIST_LIMIT,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
JOB_ID_ALPHABET,
JOB_ID_RANDOM_SUFFIX_LENGTH,
MAX_LIST_LIMIT,
MAX_OUTPUT_LIMIT_BYTES,
MAX_WAIT_TIMEOUT_SECONDS,
SESSION_NAME_PREFIX,
)
from shellctl.shared.output import (
OutputWindow,
read_output_window,
tail_output_window,
)
from shellctl.shared.runtime import (
default_runtime_dir,
default_state_dir,
format_timestamp,
generate_job_id,
is_terminal_status,
job_pane_target,
job_session_name,
parse_timestamp,
utc_now,
)
from shellctl.shared.schemas import (
TERMINAL_JOB_STATUSES,
DeleteJobResponse,
ErrorDetail,
ErrorResponse,
HealthResponse,
InputJobRequest,
JobInfo,
JobResult,
JobStatusName,
JobStatusView,
ListJobsResponse,
RunJobRequest,
ShellctlModel,
TerminalSize,
TerminateJobRequest,
WaitJobRequest,
)
__all__ = [
"DEFAULT_AUTH_TOKEN_ENV",
"DEFAULT_BASE_URL",
"DEFAULT_BASE_URL_ENV",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS",
"DEFAULT_GC_INTERVAL_SECONDS",
"DEFAULT_HEALTH_STATUS",
"DEFAULT_IDLE_FLUSH_SECONDS",
"DEFAULT_LIST_LIMIT",
"DEFAULT_OUTPUT_LIMIT_BYTES",
"DEFAULT_TERMINAL_COLS",
"DEFAULT_TERMINAL_ROWS",
"DEFAULT_TERMINATE_GRACE_SECONDS",
"DEFAULT_TIMEOUT_SECONDS",
"JOB_ID_ALPHABET",
"JOB_ID_RANDOM_SUFFIX_LENGTH",
"MAX_LIST_LIMIT",
"MAX_OUTPUT_LIMIT_BYTES",
"MAX_WAIT_TIMEOUT_SECONDS",
"SESSION_NAME_PREFIX",
"TERMINAL_JOB_STATUSES",
"DeleteJobResponse",
"ErrorDetail",
"ErrorResponse",
"HealthResponse",
"InputJobRequest",
"JobInfo",
"JobResult",
"JobStatusName",
"JobStatusView",
"ListJobsResponse",
"OutputWindow",
"RunJobRequest",
"ShellctlModel",
"TerminalSize",
"TerminateJobRequest",
"WaitJobRequest",
"default_runtime_dir",
"default_state_dir",
"format_timestamp",
"generate_job_id",
"is_terminal_status",
"job_pane_target",
"job_session_name",
"parse_timestamp",
"read_output_window",
"tail_output_window",
"utc_now",
]
_EXPORTS = {
"DEFAULT_AUTH_TOKEN_ENV": "shellctl.shared.constants",
"DEFAULT_BASE_URL": "shellctl.shared.constants",
"DEFAULT_BASE_URL_ENV": "shellctl.shared.constants",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS": "shellctl.shared.constants",
"DEFAULT_GC_INTERVAL_SECONDS": "shellctl.shared.constants",
"DEFAULT_HEALTH_STATUS": "shellctl.shared.constants",
"DEFAULT_IDLE_FLUSH_SECONDS": "shellctl.shared.constants",
"DEFAULT_LIST_LIMIT": "shellctl.shared.constants",
"DEFAULT_OUTPUT_LIMIT_BYTES": "shellctl.shared.constants",
"DEFAULT_TERMINAL_COLS": "shellctl.shared.constants",
"DEFAULT_TERMINAL_ROWS": "shellctl.shared.constants",
"DEFAULT_TERMINATE_GRACE_SECONDS": "shellctl.shared.constants",
"DEFAULT_TIMEOUT_SECONDS": "shellctl.shared.constants",
"JOB_ID_ALPHABET": "shellctl.shared.constants",
"JOB_ID_RANDOM_SUFFIX_LENGTH": "shellctl.shared.constants",
"MAX_LIST_LIMIT": "shellctl.shared.constants",
"MAX_OUTPUT_LIMIT_BYTES": "shellctl.shared.constants",
"MAX_WAIT_TIMEOUT_SECONDS": "shellctl.shared.constants",
"SESSION_NAME_PREFIX": "shellctl.shared.constants",
"OutputWindow": "shellctl.shared.output",
"read_output_window": "shellctl.shared.output",
"tail_output_window": "shellctl.shared.output",
"default_runtime_dir": "shellctl.shared.runtime",
"default_state_dir": "shellctl.shared.runtime",
"format_timestamp": "shellctl.shared.runtime",
"generate_job_id": "shellctl.shared.runtime",
"is_terminal_status": "shellctl.shared.runtime",
"job_pane_target": "shellctl.shared.runtime",
"job_session_name": "shellctl.shared.runtime",
"parse_timestamp": "shellctl.shared.runtime",
"utc_now": "shellctl.shared.runtime",
"TERMINAL_JOB_STATUSES": "shellctl.shared.schemas",
"DeleteJobResponse": "shellctl.shared.schemas",
"ErrorDetail": "shellctl.shared.schemas",
"ErrorResponse": "shellctl.shared.schemas",
"HealthResponse": "shellctl.shared.schemas",
"InputJobRequest": "shellctl.shared.schemas",
"JobInfo": "shellctl.shared.schemas",
"JobResult": "shellctl.shared.schemas",
"JobStatusName": "shellctl.shared.schemas",
"JobStatusView": "shellctl.shared.schemas",
"ListJobsResponse": "shellctl.shared.schemas",
"RunJobRequest": "shellctl.shared.schemas",
"ShellctlModel": "shellctl.shared.schemas",
"TerminalSize": "shellctl.shared.schemas",
"TerminateJobRequest": "shellctl.shared.schemas",
"WaitJobRequest": "shellctl.shared.schemas",
}
def __getattr__(name: str) -> Any:
if name not in _EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module = import_module(_EXPORTS[name])
value = getattr(module, name) # noqa: no-new-getattr lazy export proxy
globals()[name] = value
return value
def __dir__() -> list[str]:
return sorted(set(globals()) | set(__all__))
@@ -0,0 +1,49 @@
"""Shared shellctl constants.
This module intentionally contains only literal defaults and bounds so callers
can import stable configuration values without pulling in the heavier DTO,
sanitize, or server packages. The network CLI depends on this file for its
base URL and auth-token environment contract, so keep it import-light.
"""
DEFAULT_AUTH_TOKEN_ENV = "SHELLCTL_AUTH_TOKEN"
DEFAULT_BASE_URL_ENV = "SHELLCTL_BASE_URL"
DEFAULT_BASE_URL = "http://127.0.0.1:8765"
DEFAULT_OUTPUT_LIMIT_BYTES = 1024 * 8
MAX_OUTPUT_LIMIT_BYTES = 1024 * 1024
DEFAULT_TIMEOUT_SECONDS = 30.0
MAX_WAIT_TIMEOUT_SECONDS = 5.0 * 60.0
DEFAULT_IDLE_FLUSH_SECONDS = 0.5
DEFAULT_TERMINATE_GRACE_SECONDS = 5.0
DEFAULT_TERMINAL_COLS = 120
DEFAULT_TERMINAL_ROWS = 80
DEFAULT_LIST_LIMIT = 100
MAX_LIST_LIMIT = 1000
DEFAULT_GC_INTERVAL_SECONDS = 10.0 * 60.0
DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS = 24.0 * 60.0 * 60.0
JOB_ID_ALPHABET = "0123456789abcdefghjkmnpqrstvwxyz"
JOB_ID_RANDOM_SUFFIX_LENGTH = 3
SESSION_NAME_PREFIX = "shellctl-job-"
DEFAULT_HEALTH_STATUS = "ok"
__all__ = [
"DEFAULT_AUTH_TOKEN_ENV",
"DEFAULT_BASE_URL",
"DEFAULT_BASE_URL_ENV",
"DEFAULT_GC_FINISHED_JOB_RETENTION_SECONDS",
"DEFAULT_GC_INTERVAL_SECONDS",
"DEFAULT_HEALTH_STATUS",
"DEFAULT_IDLE_FLUSH_SECONDS",
"DEFAULT_LIST_LIMIT",
"DEFAULT_OUTPUT_LIMIT_BYTES",
"DEFAULT_TERMINAL_COLS",
"DEFAULT_TERMINAL_ROWS",
"DEFAULT_TERMINATE_GRACE_SECONDS",
"DEFAULT_TIMEOUT_SECONDS",
"JOB_ID_ALPHABET",
"JOB_ID_RANDOM_SUFFIX_LENGTH",
"MAX_LIST_LIMIT",
"MAX_OUTPUT_LIMIT_BYTES",
"MAX_WAIT_TIMEOUT_SECONDS",
"SESSION_NAME_PREFIX",
]
+120
View File
@@ -0,0 +1,120 @@
"""UTF-8-safe output slicing helpers for shellctl."""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
@dataclass(slots=True, frozen=True)
class OutputWindow:
"""UTF-8-safe slice of `output.log`."""
output: str
offset: int
truncated: bool
def read_output_window(path: Path, *, offset: int, limit: int) -> OutputWindow:
"""Read a forward UTF-8-safe slice from `output.log`."""
if offset < 0:
raise ValueError(f"offset must be >= 0; got {offset}")
if limit <= 0:
raise ValueError(f"limit must be > 0; got {limit}")
if not path.exists():
if offset == 0:
return OutputWindow(output="", offset=0, truncated=False)
raise ValueError(f"offset {offset} exceeds current file size 0")
size = path.stat().st_size
if offset > size:
raise ValueError(f"offset {offset} exceeds current file size {size}")
if offset == size:
return OutputWindow(output="", offset=offset, truncated=False)
with path.open("rb") as handle:
handle.seek(offset)
raw = handle.read(limit + 4)
start_shift = _advance_to_utf8_boundary(raw, 0)
data = raw[start_shift:]
budget = max(0, limit - start_shift)
consumed = _valid_utf8_prefix_length(data[:budget])
if consumed == 0 and data:
consumed = _first_complete_utf8_char_length(data)
output_bytes = data[:consumed]
new_offset = offset + start_shift + consumed
truncated = new_offset < size
return OutputWindow(
output=output_bytes.decode("utf-8", errors="strict"),
offset=new_offset,
truncated=truncated,
)
def tail_output_window(path: Path, *, limit: int) -> OutputWindow:
"""Read a UTF-8-safe tail snapshot from `output.log`."""
if limit <= 0:
raise ValueError(f"limit must be > 0; got {limit}")
if not path.exists():
return OutputWindow(output="", offset=0, truncated=False)
size = path.stat().st_size
if size == 0:
return OutputWindow(output="", offset=0, truncated=False)
start = max(0, size - limit)
padded_start = max(0, start - 4)
with path.open("rb") as handle:
handle.seek(padded_start)
raw = handle.read(size - padded_start)
relative_start = _advance_to_utf8_boundary(raw, start - padded_start)
payload = raw[relative_start:]
consumed = _valid_utf8_prefix_length(payload)
output_bytes = payload[:consumed]
return OutputWindow(
output=output_bytes.decode("utf-8", errors="strict"),
offset=padded_start + relative_start + consumed,
truncated=False,
)
def _advance_to_utf8_boundary(data: bytes, start: int) -> int:
while start < len(data) and _is_utf8_continuation_byte(data[start]):
start += 1
return start
def _valid_utf8_prefix_length(data: bytes) -> int:
end = len(data)
while end >= 0:
try:
data[:end].decode("utf-8", errors="strict")
return end
except UnicodeDecodeError as exc:
if exc.start < end - 4:
end = exc.start
else:
end -= 1
return 0
def _is_utf8_continuation_byte(value: int) -> bool:
return (value & 0b1100_0000) == 0b1000_0000
def _first_complete_utf8_char_length(data: bytes) -> int:
for end in range(1, len(data) + 1):
try:
decoded = data[:end].decode("utf-8", errors="strict")
except UnicodeDecodeError:
continue
if len(decoded) == 1:
return end
if decoded:
return len(decoded[0].encode("utf-8"))
return 0
__all__ = ["OutputWindow", "read_output_window", "tail_output_window"]
+95
View File
@@ -0,0 +1,95 @@
"""Stable runtime/naming helpers shared across shellctl modules."""
from __future__ import annotations
import os
import secrets
from datetime import UTC, datetime
from pathlib import Path
from shellctl.shared.constants import (
JOB_ID_ALPHABET,
JOB_ID_RANDOM_SUFFIX_LENGTH,
SESSION_NAME_PREFIX,
)
from shellctl.shared.schemas import (
TERMINAL_JOB_STATUSES,
JobStatusName,
)
def utc_now() -> datetime:
"""Return the current UTC time with timezone information."""
return datetime.now(UTC)
def format_timestamp(value: datetime | None = None) -> str:
"""Format a UTC timestamp in the stable artifact/API representation."""
moment = (value or utc_now()).astimezone(UTC).replace(microsecond=0)
return moment.isoformat().replace("+00:00", "Z")
def parse_timestamp(value: str) -> datetime:
"""Parse an artifact/API timestamp back into a timezone-aware datetime."""
return datetime.fromisoformat(value.replace("Z", "+00:00")).astimezone(UTC)
def is_terminal_status(status: JobStatusName) -> bool:
"""Check whether a lifecycle state is terminal."""
return status in TERMINAL_JOB_STATUSES
def generate_job_id(*, now: datetime | None = None) -> str:
"""Generate a short, human-readable job id."""
timestamp = (now or utc_now()).astimezone(UTC).strftime("%m%d%H%M")
suffix = "".join(secrets.choice(JOB_ID_ALPHABET) for _ in range(JOB_ID_RANDOM_SUFFIX_LENGTH))
return f"{timestamp}-{suffix}"
def job_session_name(job_id: str) -> str:
"""Return the dedicated tmux session name for a job."""
return f"{SESSION_NAME_PREFIX}{job_id}"
def job_pane_target(job_id: str) -> str:
"""Return the canonical tmux pane target for a single-pane job session."""
return f"{job_session_name(job_id)}:0.0"
def default_state_dir() -> Path:
"""Resolve the XDG-style default shellctl state directory."""
xdg_state_home = os.environ.get("XDG_STATE_HOME")
if xdg_state_home:
return Path(xdg_state_home) / "shellctl"
return Path.home() / ".local" / "state" / "shellctl"
def default_runtime_dir(state_dir: Path | None = None) -> Path:
"""Resolve the XDG-style default shellctl runtime directory."""
xdg_runtime_dir = os.environ.get("XDG_RUNTIME_DIR")
if xdg_runtime_dir:
return Path(xdg_runtime_dir) / "shellctl"
base_state_dir = state_dir or default_state_dir()
return base_state_dir / "run" / "shellctl"
__all__ = [
"default_runtime_dir",
"default_state_dir",
"format_timestamp",
"generate_job_id",
"is_terminal_status",
"job_pane_target",
"job_session_name",
"parse_timestamp",
"utc_now",
]
+213
View File
@@ -0,0 +1,213 @@
"""Shared pydantic transport models for shellctl."""
from __future__ import annotations
from enum import StrEnum
from pydantic import BaseModel, ConfigDict, Field, field_validator
from shellctl.shared.constants import (
DEFAULT_HEALTH_STATUS,
DEFAULT_IDLE_FLUSH_SECONDS,
DEFAULT_OUTPUT_LIMIT_BYTES,
DEFAULT_TERMINAL_COLS,
DEFAULT_TERMINAL_ROWS,
DEFAULT_TERMINATE_GRACE_SECONDS,
DEFAULT_TIMEOUT_SECONDS,
MAX_OUTPUT_LIMIT_BYTES,
MAX_WAIT_TIMEOUT_SECONDS,
)
class ShellctlModel(BaseModel):
"""Base pydantic model with strict extra-field handling.
shellctl uses these DTOs directly for HTTP request/response bodies, so
silently accepting unknown fields would make it harder to detect schema
drift.
"""
model_config = ConfigDict(extra="forbid")
class JobStatusName(StrEnum):
"""Lifecycle states materialized into SQLite rows and API responses."""
CREATED = "created"
STARTING = "starting"
RUNNING = "running"
EXITED = "exited"
TERMINATED = "terminated"
FAILED = "failed"
LOST = "lost"
TERMINAL_JOB_STATUSES = frozenset(
{
JobStatusName.EXITED,
JobStatusName.TERMINATED,
JobStatusName.FAILED,
JobStatusName.LOST,
}
)
class TerminalSize(ShellctlModel):
"""Requested initial PTY geometry for a job."""
cols: int = Field(default=DEFAULT_TERMINAL_COLS, ge=1, le=4096)
rows: int = Field(default=DEFAULT_TERMINAL_ROWS, ge=1, le=4096)
class JobResult(ShellctlModel):
"""Unified response shape for output-oriented job APIs."""
job_id: str
done: bool
status: JobStatusName
exit_code: int | None = None
output_path: str
output: str
offset: int = Field(ge=0)
truncated: bool
class JobStatusView(ShellctlModel):
"""Materialized lifecycle view returned by status-like APIs."""
job_id: str
status: JobStatusName
done: bool
exit_code: int | None = None
created_at: str
started_at: str | None = None
ended_at: str | None = None
offset: int = Field(ge=0)
class JobInfo(ShellctlModel):
"""Compact job listing record."""
job_id: str
status: JobStatusName
created_at: str
started_at: str | None = None
ended_at: str | None = None
class ListJobsResponse(ShellctlModel):
"""Response body for `GET /v1/jobs`."""
jobs: list[JobInfo]
class DeleteJobResponse(ShellctlModel):
"""Response body for successful delete operations."""
job_id: str
deleted: bool = True
class HealthResponse(ShellctlModel):
"""Public health check response."""
status: str = DEFAULT_HEALTH_STATUS
class ErrorDetail(ShellctlModel):
"""Machine-readable API error payload."""
code: str
message: str
class ErrorResponse(ShellctlModel):
"""Envelope used by server-side exception handlers."""
error: ErrorDetail
class RunJobRequest(ShellctlModel):
"""HTTP request body for `POST /v1/jobs/run`.
`env` augments the runner's inherited process environment instead of
replacing it, so callers can preset script-local variables without losing
ambient values such as `PATH`.
"""
script: str
cwd: str | None = None
env: dict[str, str] | None = None
terminal: TerminalSize | None = None
timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, gt=0, le=MAX_WAIT_TIMEOUT_SECONDS)
output_limit: int = Field(default=DEFAULT_OUTPUT_LIMIT_BYTES, ge=1, le=MAX_OUTPUT_LIMIT_BYTES)
idle_flush_seconds: float = Field(default=DEFAULT_IDLE_FLUSH_SECONDS, ge=0, le=30)
@field_validator("env")
@classmethod
def _validate_env(cls, env: dict[str, str] | None) -> dict[str, str] | None:
"""Reject env entries that cannot be represented in `execve`.
shellctl applies `env` as a process environment overlay, so validation
follows the low-level `NAME=value` constraints instead of shell variable
naming rules: names must be non-empty and cannot contain `=` or NUL,
while values cannot contain NUL.
"""
if env is None:
return None
for name, value in env.items():
if not name:
raise ValueError("env names must be non-empty")
if "=" in name:
raise ValueError(f"env name must not contain '=': {name!r}")
if "\x00" in name:
raise ValueError(f"env name must not contain NUL: {name!r}")
if "\x00" in value:
raise ValueError(f"env value must not contain NUL: {name!r}")
return env
class WaitJobRequest(ShellctlModel):
"""HTTP request body for `POST /v1/jobs/{job_id}/wait`."""
timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, ge=0, le=MAX_WAIT_TIMEOUT_SECONDS)
offset: int = Field(ge=0)
output_limit: int = Field(default=DEFAULT_OUTPUT_LIMIT_BYTES, ge=1, le=MAX_OUTPUT_LIMIT_BYTES)
idle_flush_seconds: float = Field(default=DEFAULT_IDLE_FLUSH_SECONDS, ge=0, le=30)
class InputJobRequest(ShellctlModel):
"""HTTP request body for `POST /v1/jobs/{job_id}/input`."""
text: str
timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, gt=0, le=MAX_WAIT_TIMEOUT_SECONDS)
offset: int = Field(ge=0)
output_limit: int = Field(default=DEFAULT_OUTPUT_LIMIT_BYTES, ge=1, le=MAX_OUTPUT_LIMIT_BYTES)
idle_flush_seconds: float = Field(default=DEFAULT_IDLE_FLUSH_SECONDS, ge=0, le=30)
class TerminateJobRequest(ShellctlModel):
"""HTTP request body for `POST /v1/jobs/{job_id}/terminate`."""
grace_seconds: float = Field(default=DEFAULT_TERMINATE_GRACE_SECONDS, ge=0, le=300)
__all__ = [
"TERMINAL_JOB_STATUSES",
"DeleteJobResponse",
"ErrorDetail",
"ErrorResponse",
"HealthResponse",
"InputJobRequest",
"JobInfo",
"JobResult",
"JobStatusName",
"JobStatusView",
"ListJobsResponse",
"RunJobRequest",
"ShellctlModel",
"TerminalSize",
"TerminateJobRequest",
"WaitJobRequest",
]
@@ -0,0 +1,7 @@
"""Stdlib-only shellctl subprocess entrypoints.
This package is reserved for tmux hot-path helpers that must start quickly for
every job. Keep `__init__` import-free so `shellctl-sanitize-pty` and
`shellctl-runner-exit` do not accidentally pull in the main shellctl client or
server stacks.
"""

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