"""专业软件任务的校验、幂等持久化和 offer 状态机。""" from __future__ import annotations import json import re from datetime import datetime, timedelta, timezone from hashlib import sha256 from uuid import UUID, uuid4 from sqlalchemy import and_, desc, or_, select from sqlalchemy.exc import IntegrityError from core.software_nodes import SUPPORTED_CAPABILITIES from core.storage.engine import session_scope from core.storage.models import Artifact, SoftwareJob, SoftwareNode, Task OFFER_SECONDS = 60 ALLOWED_PLOT_TYPES = frozenset( {"line", "scatter", "line_scatter"} ) ALLOWED_INPUT_SUFFIXES = frozenset({".csv", ".xlsx", ".json"}) MAX_INPUT_BYTES = 100 * 1024 * 1024 MAX_INPUTS = 16 MAX_INPUT_TOTAL_BYTES = 512 * 1024 * 1024 MAX_OUTPUT_ARTIFACT_BYTES = 256 * 1024 * 1024 class SoftwareJobError(Exception): pass MAX_OUTPUT_TOTAL_BYTES = 512 * 1024 * 1024 OUTPUT_ARTIFACTS = { "project": ("project.opju", "application/x-origin-project", "opju"), "figure_png": ("figure.png", "image/png", "png"), "figure_svg": ("figure.svg", "image/svg+xml", "svg"), "figure_pdf": ("figure.pdf", "application/pdf", "pdf"), "plot_spec": ("plot-spec.json", "application/json", None), "provenance": ("provenance.json", "application/json", None), } SOFTWARE_JOB_METADATA_IDS = frozenset({"plot_spec", "provenance"}) ORIGIN_OUTPUT_IDENTITIES = { ("project", "opju"): "project", ("figure", "png"): "figure_png", ("figure", "svg"): "figure_svg", ("figure", "pdf"): "figure_pdf", } def software_job_output_path(output_id: str) -> str: """返回 Job 输出目录内路径;技术元数据固定进入隐藏 `.meta/`。""" item = OUTPUT_ARTIFACTS.get(output_id) if item is None: raise SoftwareJobError("unsupported output artifact identity") filename = item[0] return f".meta/{filename}" if output_id in SOFTWARE_JOB_METADATA_IDS else filename def _has_only(value: dict, fields: set[str]) -> bool: return set(value).issubset(fields) def _canonical_origin_plot_request(request: dict) -> tuple[dict, str]: if not isinstance(request, dict) or set(request) != { "schema_version", "inputs", "operation", "outputs" }: raise SoftwareJobError("invalid origin plot request fields") if request.get("schema_version") != 2: raise SoftwareJobError("unsupported origin plot schema version") inputs = request.get("inputs") operation = request.get("operation") outputs = request.get("outputs") if not isinstance(inputs, list) or not 1 <= len(inputs) <= MAX_INPUTS: raise SoftwareJobError("inputs must contain 1 to 16 artifact bindings") if not isinstance(operation, dict) or set(operation) != {"plot"}: raise SoftwareJobError("origin operation must contain exactly plot") plot = operation.get("plot") if not isinstance(plot, dict): raise SoftwareJobError("origin plot request sections must be objects") input_keys: list[str] = [] for input_spec in inputs: if not isinstance(input_spec, dict) or set(input_spec) not in ( {"key", "artifact_id"}, {"key", "artifact_id", "selector"} ): raise SoftwareJobError("invalid origin input binding fields") input_key = input_spec.get("key") if not isinstance(input_key, str) or not re.fullmatch(r"[a-z][a-z0-9_]{0,31}", input_key): raise SoftwareJobError("input key must match [a-z][a-z0-9_]{0,31}") try: UUID(str(input_spec.get("artifact_id") or "")) except ValueError as exc: raise SoftwareJobError("inputs[].artifact_id must be an artifact UUID") from exc selector = input_spec.get("selector") if selector is not None and ( not isinstance(selector, dict) or set(selector) != {"sheet"} or not isinstance(selector.get("sheet"), str) or not 1 <= len(selector["sheet"]) <= 128 ): raise SoftwareJobError("origin input selector must contain a valid sheet") input_keys.append(input_key) if len(input_keys) != len(set(input_keys)): raise SoftwareJobError("input keys must be unique") if not _has_only( plot, {"type", "series", "template", "title", "x_axis", "y_axis", "legend", "error_bars"}, ): raise SoftwareJobError("unsupported origin plot fields") if plot.get("type") not in ALLOWED_PLOT_TYPES: raise SoftwareJobError("unsupported origin plot type") if "title" in plot and ( not isinstance(plot["title"], str) or len(plot["title"]) > 500 ): raise SoftwareJobError("plot.title must be a string") if plot.get("template", "publication_double_column") != "publication_double_column": raise SoftwareJobError("unsupported origin plot template") series = plot.get("series") if not isinstance(series, list) or not 1 <= len(series) <= 16: raise SoftwareJobError("plot.series must contain 1 to 16 series") identities: list[tuple[str, str, str]] = [] used_input_keys: set[str] = set() series_labels: dict[tuple[str, str], str] = {} for item in series: if not isinstance(item, dict) or not _has_only(item, {"input", "x", "y", "label"}): raise SoftwareJobError("invalid plot series fields") if not {"input", "x", "y"}.issubset(item): raise SoftwareJobError("plot series requires input, x, and y") input_key = item.get("input") x_column = item.get("x") y_column = item.get("y") if input_key not in input_keys: raise SoftwareJobError("plot series references an unknown input") if any( not isinstance(value, str) or not 1 <= len(value) <= 128 for value in (x_column, y_column) ): raise SoftwareJobError("plot series x and y must be column names") if "label" in item and ( not isinstance(item["label"], str) or not 1 <= len(item["label"]) <= 200 ): raise SoftwareJobError("plot series label must be a string") label_key = (input_key, y_column) effective_label = item.get("label", y_column) if label_key in series_labels and series_labels[label_key] != effective_label: raise SoftwareJobError("series sharing an input Y column must use one label") series_labels[label_key] = effective_label used_input_keys.add(input_key) identities.append((input_key, x_column, y_column)) if len(identities) != len(set(identities)): raise SoftwareJobError("plot series must be unique") if used_input_keys != set(input_keys): raise SoftwareJobError("every input must be referenced by a plot series") for axis_name in ("x_axis", "y_axis"): axis = plot.get(axis_name) if axis is not None and ( not isinstance(axis, dict) or not _has_only(axis, {"title", "unit", "scale"}) or axis.get("scale", "linear") != "linear" or any( name in axis and not isinstance(axis[name], str) for name in ("title", "unit") ) ): raise SoftwareJobError(f"invalid {axis_name}") legend = plot.get("legend") if legend is not None and ( not isinstance(legend, dict) or not _has_only(legend, {"enabled", "position"}) or ("enabled" in legend and not isinstance(legend["enabled"], bool)) or legend.get("enabled", True) is not True or legend.get("position", "top_right") != "top_right" ): raise SoftwareJobError("invalid plot.legend") if plot.get("error_bars") is not None: raise SoftwareJobError("error bars are not supported") if not isinstance(outputs, list) or not 1 <= len(outputs) <= 16: raise SoftwareJobError("outputs must contain 1 to 16 declarations") output_keys: list[str] = [] output_identities: list[tuple[str, str]] = [] for output in outputs: if not isinstance(output, dict) or set(output) not in ( {"key", "type", "format"}, {"key", "type", "format", "options"} ): raise SoftwareJobError("invalid origin output declaration fields") output_type = output.get("type") output_format = output.get("format") expected_key = ORIGIN_OUTPUT_IDENTITIES.get((output_type, output_format)) if output.get("key") != expected_key: raise SoftwareJobError("origin output key, type, and format do not match") options = output.get("options") if output_format == "png": if options is not None and ( not isinstance(options, dict) or set(options) != {"dpi"} or not isinstance(options.get("dpi"), int) or isinstance(options.get("dpi"), bool) or not 72 <= options["dpi"] <= 1200 ): raise SoftwareJobError("PNG output options must contain a valid dpi") elif options is not None: raise SoftwareJobError("output options are only supported for PNG") output_keys.append(expected_key) output_identities.append((output_type, output_format)) if len(output_keys) != len(set(output_keys)) or len(output_identities) != len(set(output_identities)): raise SoftwareJobError("outputs must be unique") encoded = json.dumps(request, ensure_ascii=False, sort_keys=True, separators=(",", ":")) if len(encoded.encode("utf-8")) > 256 * 1024: raise SoftwareJobError("origin plot request is too large") normalized = json.loads(encoded) return normalized, sha256(encoded.encode("utf-8")).hexdigest() REQUEST_VALIDATORS = { "origin.plot@v2": _canonical_origin_plot_request, } def _canonical_request(capability: str, request: dict) -> tuple[dict, str]: validator = REQUEST_VALIDATORS.get(capability) if validator is None: raise SoftwareJobError("unsupported capability") return validator(request) def _job_dict(row: SoftwareJob) -> dict: return { "job_id": str(row.job_id), "task_id": str(row.task_id), "capability": row.capability, "request_digest": row.request_digest, "node_id": str(row.node_id) if row.node_id else None, "status": row.status, "stage": row.stage, "progress": row.progress, "metrics": row.metrics, "error": row.error, "artifact_manifest": row.artifact_manifest, "output_dir": f"origin/{row.job_id}", "created_at": row.created_at.isoformat() if row.created_at else None, "started_at": row.started_at.isoformat() if row.started_at else None, "terminal_at": row.terminal_at.isoformat() if row.terminal_at else None, } def list_jobs( user_id: UUID, *, task_id: UUID | None = None, active_only: bool = False, limit: int = 50, before: tuple[datetime, UUID] | None = None, ) -> list[dict]: """列出用户的软件任务;用于全局 Job 中心和 Agent 查询。""" # Web 分页会多取一条判断是否还有下一页;公开接口仍把页大小限制为 100。 limit = max(1, min(int(limit), 101)) with session_scope() as session: statement = ( select(SoftwareJob, Task.name, SoftwareNode.name) .join(Task, Task.task_id == SoftwareJob.task_id) .outerjoin(SoftwareNode, SoftwareNode.node_id == SoftwareJob.node_id) .where(SoftwareJob.user_id == user_id) ) if task_id is not None: statement = statement.where(SoftwareJob.task_id == task_id) if active_only: statement = statement.where( SoftwareJob.status.in_( { "queued", "offered", "dispatched", "running", "disconnected", "cancelling", } ) ) if before is not None: before_created_at, before_job_id = before statement = statement.where( or_( SoftwareJob.created_at < before_created_at, and_( SoftwareJob.created_at == before_created_at, SoftwareJob.job_id < before_job_id, ), ) ) rows = session.execute( statement.order_by( desc(SoftwareJob.created_at), desc(SoftwareJob.job_id) ).limit(limit) ).all() results: list[dict] = [] for job, task_name, node_name in rows: item = _job_dict(job) item.update( { "task_name": task_name, "node_name": node_name, "input": job.input_manifest, "request_summary": _request_summary(job), } ) results.append(item) return results def _request_summary(job: SoftwareJob) -> dict: plot = (job.request.get("operation") or {}).get("plot") or {} outputs = job.request.get("outputs") or [] return { "display_name": ( "Origin 科研绘图" if job.capability == "origin.plot@v2" else job.capability ), "title": str(plot.get("title") or ""), "formats": [item.get("format") for item in outputs if isinstance(item, dict)], } def request_job_cancel(user_id: UUID, job_id: UUID) -> tuple[dict, dict | None]: """持久化取消意图;queued 直接终止,已分派任务返回 Node 消息。""" now = datetime.now(timezone.utc) with session_scope() as session: job = session.execute( select(SoftwareJob).where( SoftwareJob.job_id == job_id, SoftwareJob.user_id == user_id, ).with_for_update() ).scalar_one_or_none() if job is None: raise SoftwareJobError("job not found") if job.status in {"succeeded", "failed", "cancelled"}: return _job_dict(job), None if job.status == "queued" or job.node_id is None or job.lease_id is None: job.status = "cancelled" job.stage = "terminal" job.error = {"code": "USER_CANCELLED", "detail": "Cancelled before dispatch."} job.terminal_at = now return _job_dict(job), None job.status = "cancelling" job.stage = "cancel_requested" payload = { "job_id": str(job.job_id), "lease_id": str(job.lease_id), "request_digest": job.request_digest, } return _job_dict(job), {"node_id": job.node_id, "payload": payload} def pending_node_cancellations(node_id: UUID) -> list[dict]: """节点重连或心跳时重放未确认的取消意图。""" with session_scope() as session: rows = session.execute( select(SoftwareJob).where( SoftwareJob.node_id == node_id, SoftwareJob.status == "cancelling", ) ).scalars() return [ { "job_id": str(job.job_id), "lease_id": str(job.lease_id), "request_digest": job.request_digest, } for job in rows if job.lease_id is not None ] def create_job( user_id: UUID, task_id: UUID, *, idempotency_key: str, capability: str, request: dict, ) -> tuple[dict, bool]: key = idempotency_key.strip() if not key or len(key) > 200: raise SoftwareJobError("idempotency_key must contain 1 to 200 characters") if capability not in SUPPORTED_CAPABILITIES: raise SoftwareJobError("unsupported capability") normalized, digest = _canonical_request(capability, request) with session_scope() as session: task = session.execute( select(Task.task_id).where(Task.task_id == task_id, Task.user_id == user_id) ).first() if task is None: raise SoftwareJobError("task not found") artifact_ids = [UUID(item["artifact_id"]) for item in normalized["inputs"]] artifacts = session.execute( select(Artifact).where( Artifact.artifact_id.in_(set(artifact_ids)), Artifact.user_id == user_id, Artifact.status == "active", ) ).scalars().all() artifacts_by_id = {artifact.artifact_id: artifact for artifact in artifacts} input_manifest: list[dict] = [] total_input_bytes = 0 for binding, artifact_id in zip(normalized["inputs"], artifact_ids, strict=True): artifact = artifacts_by_id.get(artifact_id) if artifact is None: raise SoftwareJobError("input artifact not found") suffix = ( "." + artifact.current_path.rsplit(".", 1)[-1].lower() if "." in artifact.current_path else "" ) if suffix not in ALLOWED_INPUT_SUFFIXES: raise SoftwareJobError("input artifact type is not supported") if ( artifact.size_bytes is None or artifact.size_bytes < 0 or artifact.size_bytes > MAX_INPUT_BYTES or not artifact.content_sha256 or len(artifact.content_sha256) != 64 ): raise SoftwareJobError("input artifact metadata is incomplete or too large") total_input_bytes += artifact.size_bytes item = { "key": binding["key"], "artifact_id": str(artifact.artifact_id), "filename": artifact.current_path.replace("\\", "/").rsplit("/", 1)[-1], "size_bytes": artifact.size_bytes, "sha256": artifact.content_sha256, } if binding.get("selector") is not None: item["selector"] = binding["selector"] input_manifest.append(item) if total_input_bytes > MAX_INPUT_TOTAL_BYTES: raise SoftwareJobError("job inputs exceed the total size limit") existing = session.execute( select(SoftwareJob).where( SoftwareJob.user_id == user_id, SoftwareJob.idempotency_key == key, ) ).scalar_one_or_none() if existing is not None: if ( existing.task_id != task_id or existing.capability != capability or existing.request_digest != digest ): raise SoftwareJobError("idempotency key was already used for a different request") return _job_dict(existing), False row = SoftwareJob( job_id=uuid4(), user_id=user_id, task_id=task_id, idempotency_key=key, capability=capability, request=normalized, request_digest=digest, input_manifest=input_manifest, status="queued", stage="", metrics={}, error={}, artifact_manifest=[], ) try: with session.begin_nested(): session.add(row) session.flush() return _job_dict(row), True except IntegrityError: existing = session.execute( select(SoftwareJob).where( SoftwareJob.user_id == user_id, SoftwareJob.idempotency_key == key, ) ).scalar_one() if ( existing.task_id != task_id or existing.capability != capability or existing.request_digest != digest ): raise SoftwareJobError( "idempotency key was already used for a different request" ) return _job_dict(existing), False def get_job(user_id: UUID, job_id: UUID) -> dict | None: with session_scope() as session: row = session.execute( select(SoftwareJob).where(SoftwareJob.job_id == job_id, SoftwareJob.user_id == user_id) ).scalar_one_or_none() return _job_dict(row) if row else None def offer_next_job(node_ids: set[UUID]) -> dict | None: """从当前进程实际在线的节点中选择一个,为最早 queued job 创建短租约。""" if not node_ids: return None now = datetime.now(timezone.utc) with session_scope() as session: expired = session.execute( select(SoftwareJob) .where( SoftwareJob.status == "offered", SoftwareJob.lease_expires_at <= now, ) .with_for_update(skip_locked=True) ).scalars() for item in expired: item.status = "queued" item.node_id = None item.lease_id = None item.lease_expires_at = None job = session.execute( select(SoftwareJob) .where(SoftwareJob.status == "queued") .order_by(SoftwareJob.created_at, SoftwareJob.job_id) .with_for_update(skip_locked=True) .limit(1) ).scalar_one_or_none() if job is None: return None busy_node_ids = set( session.execute( select(SoftwareJob.node_id).where( SoftwareJob.node_id.is_not(None), SoftwareJob.status.in_({"offered", "dispatched", "running"}), ) ).scalars() ) nodes = session.execute( select(SoftwareNode) .where(SoftwareNode.node_id.in_(node_ids), SoftwareNode.status == "online") .order_by(SoftwareNode.last_seen_at.desc()) ).scalars() node = next( ( item for item in nodes if item.node_id not in busy_node_ids and job.capability in item.capabilities and int((item.runtime or {}).get("available_slots") or 0) > 0 ), None, ) if node is None: return None lease_id = uuid4() expires_at = now + timedelta(seconds=OFFER_SECONDS) job.node_id = node.node_id job.lease_id = lease_id job.lease_expires_at = expires_at job.status = "offered" return { "node_id": node.node_id, "payload": { "job_id": str(job.job_id), "lease_id": str(lease_id), "lease_expires_at": expires_at.isoformat(), "capability": job.capability, "request_digest": job.request_digest, "request": job.request, "input_transfers": [ { **item, "download_path": ( f"/v1/software-jobs/{job.job_id}/inputs/{item['key']}" ), } for item in job.input_manifest ], }, } def get_job_input(node_id: UUID, job_id: UUID, input_key: str) -> dict | None: """返回任务绑定的 artifact 定位信息;调用方仍需在 user_root 内安全解析。""" with session_scope() as session: job = session.execute( select(SoftwareJob).where( SoftwareJob.job_id == job_id, SoftwareJob.node_id == node_id, SoftwareJob.status.in_({"offered", "dispatched", "running", "disconnected"}), ) ).scalar_one_or_none() if job is None: return None manifest = next( (item for item in job.input_manifest if item.get("key") == input_key), None, ) if manifest is None: return None artifact_id = UUID(manifest["artifact_id"]) artifact = session.execute( select(Artifact).where( Artifact.artifact_id == artifact_id, Artifact.user_id == job.user_id, Artifact.status == "active", ) ).scalar_one_or_none() if artifact is None: return None return { "user_id": job.user_id, "current_path": artifact.current_path, **manifest, } def get_job_output_context(node_id: UUID, job_id: UUID, lease_id: UUID, digest: str) -> dict | None: """返回 Node 输出上传上下文,不向 Node 暴露任何云端文件路径。""" with session_scope() as session: row = session.execute( select(SoftwareJob, Task.working_dir) .join(Task, Task.task_id == SoftwareJob.task_id) .where(SoftwareJob.job_id == job_id) ).one_or_none() if row is None: return None job, working_dir = row if ( job.node_id != node_id or job.lease_id != lease_id or job.request_digest != digest or job.status not in {"dispatched", "running", "disconnected", "succeeded"} ): return None return { "user_id": job.user_id, "task_id": job.task_id, "working_dir": working_dir, "request": job.request, "status": job.status, "artifact_manifest": job.artifact_manifest, } def validate_output_manifest(request: dict, manifest: object) -> list[dict]: if not isinstance(manifest, list): raise SoftwareJobError("job artifact manifest must be a list") expected_ids = {"plot_spec", "provenance"} expected_ids.update( item.get("key") for item in request.get("outputs", []) if isinstance(item, dict) and isinstance(item.get("key"), str) ) if len(manifest) != len(expected_ids): raise SoftwareJobError("job artifact manifest is incomplete") normalized: list[dict] = [] seen: set[str] = set() total = 0 for raw in manifest: if not isinstance(raw, dict) or set(raw) != { "artifact_id", "filename", "media_type", "size_bytes", "sha256" }: raise SoftwareJobError("job artifact manifest entry is invalid") local_id = raw.get("artifact_id") if local_id not in expected_ids or local_id in seen: raise SoftwareJobError("job artifact manifest identity is invalid") filename, media_type, _ = OUTPUT_ARTIFACTS[local_id] size = raw.get("size_bytes") digest = raw.get("sha256") if raw.get("filename") != filename or raw.get("media_type") != media_type: raise SoftwareJobError("job artifact manifest metadata does not match its identity") if not isinstance(size, int) or isinstance(size, bool) or not 1 <= size <= MAX_OUTPUT_ARTIFACT_BYTES: raise SoftwareJobError("job output artifact size is invalid") if not isinstance(digest, str) or not re.fullmatch(r"[0-9a-f]{64}", digest): raise SoftwareJobError("job output artifact digest is invalid") total += size seen.add(local_id) normalized.append(dict(raw)) if seen != expected_ids or total > MAX_OUTPUT_TOTAL_BYTES: raise SoftwareJobError("job artifact manifest is incomplete or too large") return normalized def replay_succeeded_outputs(context: dict, manifest: list[dict]) -> list[dict] | None: """确认成功 Job 的 Node 重放,并返回云端已持久化的发布清单。 Node 可能在服务端完成发布后、写本地 ``upload-complete.json`` 前断线。 此时数据库中的 published manifest 是事实源;不能再按当前版本的目录规则 重做发布,否则跨版本布局调整会让已经成功的旧 Job 永久无法确认。 """ if context.get("status") != "succeeded": return None published = context.get("artifact_manifest") if not isinstance(published, list): raise SoftwareJobError("successful job artifact manifest is invalid") by_source = { item.get("source_artifact_id"): item for item in published if isinstance(item, dict) and isinstance(item.get("source_artifact_id"), str) } if len(by_source) != len(published) or len(by_source) != len(manifest): raise SoftwareJobError("successful job artifact manifest does not match replay") identity_fields = ("filename", "media_type", "size_bytes", "sha256") for submitted in manifest: stored = by_source.get(submitted["artifact_id"]) if stored is None or any(stored.get(field) != submitted.get(field) for field in identity_fields): raise SoftwareJobError("successful job artifact manifest does not match replay") return [dict(item) for item in published] def succeeded_output_upload_matches( context: dict, output_id: str, *, size_bytes: int, digest: str, ) -> bool: """判断单个重复 PUT 是否已包含在成功 Job 的持久化清单中。""" if context.get("status") != "succeeded": return False published = context.get("artifact_manifest") if not isinstance(published, list): raise SoftwareJobError("successful job artifact manifest is invalid") matches = [ item for item in published if isinstance(item, dict) and item.get("source_artifact_id") == output_id ] if ( len(matches) != 1 or matches[0].get("size_bytes") != size_bytes or matches[0].get("sha256") != digest ): raise SoftwareJobError("successful job output does not match replay") return True def abandon_offer(node_id: UUID, payload: dict) -> None: """WebSocket 发送失败时只回滚仍属于该连接租约的 offer。""" try: job_id = UUID(str(payload.get("job_id", ""))) lease_id = UUID(str(payload.get("lease_id", ""))) except ValueError: return with session_scope() as session: job = session.execute( select(SoftwareJob).where(SoftwareJob.job_id == job_id).with_for_update() ).scalar_one_or_none() if ( job is not None and job.status == "offered" and job.node_id == node_id and job.lease_id == lease_id ): job.status = "queued" job.node_id = None job.lease_id = None job.lease_expires_at = None def respond_to_offer(node_id: UUID, *, accepted: bool, payload: dict) -> None: try: job_id = UUID(str(payload.get("job_id", ""))) lease_id = UUID(str(payload.get("lease_id", ""))) except ValueError as exc: raise SoftwareJobError("invalid job offer response identity") from exc now = datetime.now(timezone.utc) with session_scope() as session: job = session.execute( select(SoftwareJob).where(SoftwareJob.job_id == job_id).with_for_update() ).scalar_one_or_none() if job is None or job.node_id != node_id or job.lease_id != lease_id: raise SoftwareJobError("job offer is stale or does not belong to this node") if ( accepted and job.status in {"dispatched", "running", "cancelling", "succeeded", "failed", "cancelled"} and payload.get("request_digest") == job.request_digest ): return if job.status != "offered": raise SoftwareJobError("job offer is stale or does not belong to this node") if job.lease_expires_at is None or job.lease_expires_at <= now: job.status = "queued" job.node_id = None job.lease_id = None job.lease_expires_at = None raise SoftwareJobError("job offer has expired") if accepted: if payload.get("request_digest") != job.request_digest: raise SoftwareJobError("job request digest mismatch") job.status = "dispatched" job.stage = "accepted" job.error = {} else: job.status = "queued" job.node_id = None job.lease_id = None job.lease_expires_at = None job.error = {"code": "node_rejected", "detail": str(payload.get("reason") or "")[:500]} def update_job_state(node_id: UUID, payload: dict) -> None: job_id, lease_id, digest = _message_identity(payload) stage = str(payload.get("stage") or "") progress = payload.get("progress") metrics = payload.get("metrics") or {} if not stage or len(stage) > 100: raise SoftwareJobError("job stage is required") if not isinstance(progress, int) or isinstance(progress, bool) or not 0 <= progress <= 100: raise SoftwareJobError("job progress must be between 0 and 100") if not isinstance(metrics, dict) or len(json.dumps(metrics, ensure_ascii=False)) > 64 * 1024: raise SoftwareJobError("job metrics are invalid") now = datetime.now(timezone.utc) with session_scope() as session: job = session.execute( select(SoftwareJob).where(SoftwareJob.job_id == job_id).with_for_update() ).scalar_one_or_none() _assert_job_message(job, node_id, lease_id, digest) if job.status in {"succeeded", "failed", "cancelled", "cancelling"}: return if not _can_accept_state(job.status): raise SoftwareJobError("job state cannot advance from its current status") job.status = ( "dispatched" if stage in {"accepted", "downloading_inputs", "ready_to_run"} else "running" ) job.stage = stage job.progress = progress job.metrics = metrics if job.status == "running" and job.started_at is None: job.started_at = now def record_job_terminal(node_id: UUID, payload: dict) -> None: job_id, lease_id, digest = _message_identity(payload) terminal_status = payload.get("status") if terminal_status not in {"succeeded", "failed", "cancelled"}: raise SoftwareJobError("invalid job terminal status") error = payload.get("error") or {} manifest = payload.get("artifact_manifest") or [] if not isinstance(error, dict) or len(json.dumps(error, ensure_ascii=False)) > 64 * 1024: raise SoftwareJobError("job terminal error is invalid") if not isinstance(manifest, list) or len(json.dumps(manifest, ensure_ascii=False)) > 256 * 1024: raise SoftwareJobError("job artifact manifest is invalid") now = datetime.now(timezone.utc) with session_scope() as session: job = session.execute( select(SoftwareJob).where(SoftwareJob.job_id == job_id).with_for_update() ).scalar_one_or_none() _assert_job_message(job, node_id, lease_id, digest) if terminal_status == "succeeded": expected = validate_output_manifest( job.request, [ { "artifact_id": item.get("source_artifact_id"), "filename": item.get("filename"), "media_type": item.get("media_type"), "size_bytes": item.get("size_bytes"), "sha256": item.get("sha256"), } for item in manifest if isinstance(item, dict) ], ) if len(expected) != len(manifest) or any( not _published_output_is_valid(job.job_id, item) for item in manifest ): raise SoftwareJobError("successful job artifacts have not been published") if job.status in {"succeeded", "failed", "cancelled"}: if job.status != terminal_status: raise SoftwareJobError("job terminal status conflicts with existing terminal") return if job.status not in {"offered", "dispatched", "running", "disconnected", "cancelling"}: raise SoftwareJobError("job terminal cannot advance from its current status") job.status = terminal_status job.stage = "terminal" job.progress = 100 if terminal_status == "succeeded" else job.progress job.error = error job.artifact_manifest = manifest job.terminal_at = now def _is_uuid(value: str) -> bool: try: UUID(value) except ValueError: return False return True def _published_output_is_valid(job_id: UUID, item: object) -> bool: if not isinstance(item, dict): return False output_id = item.get("source_artifact_id") if not isinstance(output_id, str): return False if output_id not in OUTPUT_ARTIFACTS: return False expected_path = f"origin/{job_id}/{software_job_output_path(output_id)}" if item.get("path") != expected_path: return False artifact_id = item.get("artifact_id") if output_id in SOFTWARE_JOB_METADATA_IDS: return artifact_id is None return isinstance(artifact_id, str) and _is_uuid(artifact_id) def mark_node_jobs_disconnected(node_id: UUID) -> None: """连接丢失后保留 Node 归属和 lease,禁止任务被自动重派。""" with session_scope() as session: jobs = session.execute( select(SoftwareJob) .where( SoftwareJob.node_id == node_id, SoftwareJob.status.in_({"dispatched", "running"}), ) .with_for_update() ).scalars() for job in jobs: job.status = "disconnected" def _can_accept_state(status: str) -> bool: return status in {"offered", "dispatched", "running", "disconnected"} def _message_identity(payload: dict) -> tuple[UUID, UUID, str]: try: job_id = UUID(str(payload.get("job_id", ""))) lease_id = UUID(str(payload.get("lease_id", ""))) except ValueError as exc: raise SoftwareJobError("invalid job message identity") from exc digest = str(payload.get("request_digest") or "") if len(digest) != 64: raise SoftwareJobError("invalid job request digest") return job_id, lease_id, digest def _assert_job_message( job: SoftwareJob | None, node_id: UUID, lease_id: UUID, digest: str, ) -> None: if ( job is None or job.node_id != node_id or job.lease_id != lease_id or job.request_digest != digest ): raise SoftwareJobError("job message does not belong to this node or lease")