Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b5e35fc2fc | ||
|
|
7bfbb2bbe8 | ||
|
|
d177998255 | ||
|
|
74ee665af6 | ||
|
|
d655153a3d | ||
|
|
925f97be20 | ||
|
|
9252d81826 | ||
|
|
9d5819a9c1 | ||
|
|
99e3b1a401 | ||
|
|
45957225cd | ||
|
|
2f99652203 | ||
|
|
dc1131b6df | ||
|
|
a5a7c762a3 | ||
|
|
063e390c5d | ||
|
|
3b3c25273a |
@@ -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.
|
||||
|
||||
@@ -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
@@ -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}, "
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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="",
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -276,6 +276,7 @@ class RequestRequestDownloadFile(BaseModel):
|
||||
"validation",
|
||||
]
|
||||
file: RequestDownloadFileMapping
|
||||
for_external: bool = True
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "dify-api"
|
||||
version = "1.15.0"
|
||||
version = "1.16.0-rc1"
|
||||
requires-python = "~=3.12.0"
|
||||
|
||||
dependencies = [
|
||||
|
||||
@@ -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 +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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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__))
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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()
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
Reference in New Issue
Block a user