"""Durable, privacy-safe jobs for the accepted company-report processor.""" from __future__ import annotations import calendar import hashlib import json import os import re import stat import threading import time import uuid from collections import deque from dataclasses import asdict, dataclass from datetime import date, datetime, time as datetime_time, timezone from pathlib import Path, PurePosixPath from typing import Any, Callable, Deque, Dict, List, Mapping, Optional, Protocol, Tuple from zoneinfo import ZoneInfo from arr_web.contracts import PortalError, validate_month from arr_web.downloads import ArtifactDescriptor, MAX_DOWNLOAD_BYTES from company_reports.contracts import ( COMPANY_NAMES, ENGLISH_MONTH_NAMES, RESULT_SCHEMA_VERSION, ) from company_reports.service import CompanyReportService, RunRequest COMPANY_REPORT_TIME_ZONE = ZoneInfo("Asia/Bangkok") PERIOD_KEYS: Tuple[str, ...] = ("01-10", "11-20", "21-month-end") ACTIVE_STATES = frozenset({"queued", "running"}) TERMINAL_STATES = frozenset({"succeeded", "partial_failure", "failed"}) JOB_ID_RE = re.compile(r"^[0-9a-f]{32}$") SHA256_RE = re.compile(r"^[0-9a-f]{64}$") MAX_JOB_BYTES = 2 * 1024 * 1024 XLSX_MIME = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet" def period_to_as_of(report_month: str, period: str) -> date: validate_month(report_month) if period not in PERIOD_KEYS: raise PortalError( "COMPANY_REPORT_PERIOD_INVALID", "期间必须是 01-10、11-20 或 21-month-end", ) year, month = (int(part) for part in report_month.split("-")) if period == "01-10": day = 10 elif period == "11-20": day = 20 else: day = calendar.monthrange(year, month)[1] return date(year, month, day) def period_label(report_month: str, period: str) -> str: as_of = period_to_as_of(report_month, period) if period in {"01-10", "11-20"}: return period return f"21-{as_of.day:02d}" def period_complete_at(report_month: str, period: str) -> datetime: period_to_as_of(report_month, period) year, month = (int(part) for part in report_month.split("-")) if period == "01-10": released = date(year, month, 11) elif period == "11-20": released = date(year, month, 21) elif month == 12: released = date(year + 1, 1, 1) else: released = date(year, month + 1, 1) return datetime.combine(released, datetime_time.min, COMPANY_REPORT_TIME_ZONE) def period_release_at(report_month: str, period: str) -> datetime: """Backward-compatible alias for the period completion boundary.""" return period_complete_at(report_month, period) def _utc_now() -> str: return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") def _sha256_file(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as source: for chunk in iter(lambda: source.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() class CompanyReportExecutor(Protocol): def run( self, report_year: int, report_month: int, as_of_date: date, ) -> Mapping[str, Any]: ... @dataclass class ProgramCompanyReportExecutor: """Calls the accepted deterministic processor; no Agent or shell is involved.""" service: CompanyReportService def run( self, report_year: int, report_month: int, as_of_date: date, ) -> Mapping[str, Any]: return self.service.run( RunRequest(report_year, report_month, as_of_date, COMPANY_NAMES) ).to_dict() class CompanyReportCoordinator(Protocol): def create( self, report_month: str, period: str, *, source: Optional[Mapping[str, Any]] = None, ) -> Dict[str, Any]: ... def list_jobs( self, report_month: Optional[str] = None, limit: int = 50, offset: int = 0, ) -> tuple[List[Dict[str, Any]], int]: ... def list_month_counts(self) -> List[Dict[str, Any]]: ... def get_job(self, job_id: str) -> Dict[str, Any]: ... def resolve_download(self, job_id: str, company: str) -> ArtifactDescriptor: ... class UnavailableCompanyReportCoordinator: @staticmethod def _unavailable() -> PortalError: return PortalError( "COMPANY_REPORT_GENERATOR_UNAVAILABLE", "公司渠道明细生成服务暂不可用", 503, ) def create( self, report_month: str, period: str, *, source: Optional[Mapping[str, Any]] = None, ) -> Dict[str, Any]: _ = source raise self._unavailable() def list_jobs( self, report_month: Optional[str] = None, limit: int = 50, offset: int = 0, ) -> tuple[List[Dict[str, Any]], int]: raise self._unavailable() def list_month_counts(self) -> List[Dict[str, Any]]: return [] def get_job(self, job_id: str) -> Dict[str, Any]: raise self._unavailable() def resolve_download(self, job_id: str, company: str) -> ArtifactDescriptor: raise self._unavailable() class PersistentCompanyReportCoordinator: """One private queue, fixed five-company scope, and controlled artifacts.""" def __init__( self, project_root: Path, output_root: Path, executor: CompanyReportExecutor, *, jobs_root: Optional[Path] = None, now: Optional[Callable[[], datetime]] = None, ) -> None: self._project_root = project_root.resolve() self._output_root = output_root.resolve() try: self._output_root.relative_to(self._project_root) except ValueError: raise ValueError("company report output must be inside the project root") from None self._output_root.mkdir(parents=True, exist_ok=True, mode=0o700) self._jobs_root = (jobs_root or (self._output_root / ".web-jobs")).resolve() try: self._jobs_root.relative_to(self._output_root) except ValueError: raise ValueError("company report jobs must be inside the output root") from None self._jobs_root.mkdir(parents=True, exist_ok=True, mode=0o700) os.chmod(self._jobs_root, 0o700) self._executor = executor self._now = now or (lambda: datetime.now(COMPANY_REPORT_TIME_ZONE)) self._condition = threading.Condition(threading.RLock()) self._pending: Deque[str] = deque() self._active_job_id: Optional[str] = None self._stopping = False self._thread = threading.Thread( target=self._worker, name="arr-company-report-worker", daemon=True, ) with self._condition: self._recover_locked() self._thread.start() def close(self) -> None: with self._condition: self._stopping = True self._condition.notify_all() self._thread.join(timeout=10) @staticmethod def _source_metadata( source: Optional[Mapping[str, Any]], ) -> Optional[Dict[str, Any]]: if source is None: return None if not isinstance(source, Mapping): raise PortalError( "COMPANY_REPORT_SOURCE_INVALID", "当前 Excel 数据源信息无效", 409, ) source_batch_id = source.get("source_batch_id") source_type = source.get("source_type") filename = source.get("filename") activated_at = source.get("activated_at") disposition = source.get("disposition") if ( isinstance(source_batch_id, bool) or not isinstance(source_batch_id, int) or source_batch_id <= 0 or not isinstance(source_type, str) or source_type not in {"excel", "historical"} or ( filename is not None and ( not isinstance(filename, str) or not filename or len(filename) > 255 or PurePosixPath(filename).name != filename ) ) or ( activated_at is not None and (not isinstance(activated_at, str) or len(activated_at) > 80) ) or not isinstance(disposition, str) or not disposition or len(disposition) > 40 ): raise PortalError( "COMPANY_REPORT_SOURCE_INVALID", "当前 Excel 数据源信息无效", 409, ) counts: Dict[str, int] = {} for key in ( "source_rows", "worksheet_count", "distinct_group_codes", "room_quantity", ): value = source.get(key) if isinstance(value, bool) or not isinstance(value, int) or value < 0: raise PortalError( "COMPANY_REPORT_SOURCE_INVALID", "当前 Excel 数据源信息无效", 409, ) counts[key] = value return { "source_batch_id": source_batch_id, "source_type": source_type, "filename": filename, **counts, "activated_at": activated_at, "disposition": disposition, } def create( self, report_month: str, period: str, *, source: Optional[Mapping[str, Any]] = None, ) -> Dict[str, Any]: as_of_date = period_to_as_of(report_month, period) source_metadata = self._source_metadata(source) current = self._now() if current.tzinfo is None: current = current.replace(tzinfo=COMPANY_REPORT_TIME_ZONE) current = current.astimezone(COMPANY_REPORT_TIME_ZONE) report_year, report_month_number = ( int(part) for part in report_month.split("-") ) if (report_year, report_month_number) > (current.year, current.month): raise PortalError( "COMPANY_REPORT_MONTH_IN_FUTURE", "未来报表月份暂不可生成", 409, ) with self._condition: for existing in self._states_locked(): if ( existing.get("state") in ACTIVE_STATES and existing.get("report_month") == report_month and existing.get("period") == period ): raise PortalError( "COMPANY_REPORT_ALREADY_RUNNING", "同一月份和期间已有生成任务", 409, ) job_id = uuid.uuid4().hex state: Dict[str, Any] = { "schema_version": "1.0", "job_id": job_id, "state": "queued", "stage": "queued", "report_month": report_month, "period": period, "period_label": period_label(report_month, period), "as_of_date": as_of_date.isoformat(), "requested_companies": list(COMPANY_NAMES), "source": source_metadata, "created_at": _utc_now(), "started_at": None, "finished_at": None, "duration_seconds": None, "progress": {"percent": 0, "label": "等待生成"}, "message": "任务已进入生成队列。", "company_results": [], "artifact_files": {}, "failure_code": None, } self._write_locked(state) self._pending.append(job_id) self._condition.notify_all() return self._public_locked(state) def list_jobs( self, report_month: Optional[str] = None, limit: int = 50, offset: int = 0, ) -> tuple[List[Dict[str, Any]], int]: if report_month is not None: validate_month(report_month) if ( isinstance(limit, bool) or isinstance(offset, bool) or not 1 <= limit <= 200 or not 0 <= offset <= 10_000_000 ): raise PortalError("QUERY_INVALID", "查询参数无效") with self._condition: states = self._states_locked() if report_month is not None: states = [ state for state in states if state.get("report_month") == report_month ] total = len(states) page = states[offset : offset + limit] return [self._public_locked(state) for state in page], total def list_month_counts(self) -> List[Dict[str, Any]]: with self._condition: counts: Dict[str, int] = {} for state in self._states_locked(): month_key = state.get("report_month") if not isinstance(month_key, str): continue try: validate_month(month_key) except PortalError: continue counts[month_key] = counts.get(month_key, 0) + 1 return [ {"month_key": month_key, "company_count": counts[month_key]} for month_key in sorted(counts, reverse=True) ] def get_job(self, job_id: str) -> Dict[str, Any]: with self._condition: return self._public_locked(self._read_locked(job_id)) def resolve_download(self, job_id: str, company: str) -> ArtifactDescriptor: if company not in COMPANY_NAMES: raise PortalError("COMPANY_REPORT_COMPANY_INVALID", "公司名称无效") with self._condition: state = self._read_locked(job_id) artifacts = state.get("artifact_files") raw = artifacts.get(company) if isinstance(artifacts, Mapping) else None if not isinstance(raw, Mapping): raise PortalError( "COMPANY_REPORT_DOWNLOAD_NOT_FOUND", "该公司的正式表格尚不可下载", 404, ) try: descriptor = ArtifactDescriptor(**dict(raw)) descriptor.validate() except (TypeError, PortalError): raise PortalError( "COMPANY_REPORT_DOWNLOAD_INVALID", "文件身份无效", 500, ) from None return descriptor def _job_path(self, job_id: str) -> Path: if JOB_ID_RE.fullmatch(job_id) is None: raise PortalError("COMPANY_REPORT_JOB_NOT_FOUND", "任务不存在", 404) return self._jobs_root / f"{job_id}.json" def _read_locked(self, job_id: str) -> Dict[str, Any]: path = self._job_path(job_id) try: metadata = path.lstat() if path.is_symlink() or not stat.S_ISREG(metadata.st_mode) or metadata.st_size > MAX_JOB_BYTES: raise OSError("unsafe job state") value = json.loads(path.read_text(encoding="utf-8")) except (OSError, ValueError, json.JSONDecodeError): raise PortalError("COMPANY_REPORT_JOB_NOT_FOUND", "任务不存在", 404) from None if not isinstance(value, dict) or value.get("job_id") != job_id: raise PortalError("COMPANY_REPORT_JOB_NOT_FOUND", "任务不存在", 404) return value def _write_locked(self, state: Mapping[str, Any]) -> None: job_id = str(state.get("job_id", "")) target = self._job_path(job_id) temporary = self._jobs_root / f".{job_id}.{uuid.uuid4().hex}.tmp" descriptor = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) try: with os.fdopen(descriptor, "w", encoding="utf-8") as destination: descriptor = -1 json.dump(state, destination, ensure_ascii=False, sort_keys=True, indent=2) destination.write("\n") destination.flush() os.fsync(destination.fileno()) os.replace(temporary, target) os.chmod(target, 0o600) finally: if descriptor >= 0: os.close(descriptor) temporary.unlink(missing_ok=True) def _states_locked(self) -> List[Dict[str, Any]]: states: List[Dict[str, Any]] = [] for path in self._jobs_root.glob("*.json"): if JOB_ID_RE.fullmatch(path.stem) is None: continue try: states.append(self._read_locked(path.stem)) except PortalError: continue return sorted(states, key=lambda value: str(value.get("created_at", "")), reverse=True) def _update_locked(self, job_id: str, **changes: Any) -> Dict[str, Any]: state = self._read_locked(job_id) state.update(changes) self._write_locked(state) return state def _recover_locked(self) -> None: queued: List[str] = [] for state in reversed(self._states_locked()): job_id = str(state["job_id"]) if state.get("state") == "running": self._fail_locked( job_id, "COMPANY_REPORT_JOB_INTERRUPTED", "服务重启中断了任务,请重新生成。", ) elif state.get("state") == "queued": queued.append(job_id) self._pending.extend(queued) def _queue_position_locked(self, job_id: str) -> Optional[int]: if self._active_job_id == job_id: return 0 try: return list(self._pending).index(job_id) + 1 except ValueError: return None def _public_locked(self, state: Mapping[str, Any]) -> Dict[str, Any]: job_id = str(state.get("job_id", "")) artifacts = state.get("artifact_files") artifact_companies = set(artifacts) if isinstance(artifacts, Mapping) else set() companies: List[Dict[str, Any]] = [] downloads: Dict[str, str] = {} for raw in state.get("company_results", []): if not isinstance(raw, Mapping) or raw.get("company") not in COMPANY_NAMES: continue company = str(raw["company"]) item: Dict[str, Any] = { "company": company, "status": raw.get("status"), "row_count": raw.get("row_count", 0), "period_row_counts": dict(raw.get("period_row_counts") or {}), "warnings": list(raw.get("warnings") or []), "errors": list(raw.get("errors") or []), "version_no": raw.get("version_no"), } if raw.get("status") == "success" and company in artifact_companies: url = f"/api/company-reports/jobs/{job_id}/downloads/{company}" item["download_url"] = url downloads[company] = url companies.append(item) return { "job_id": job_id, "state": state.get("state"), "stage": state.get("stage"), "report_month": state.get("report_month"), "period": state.get("period"), "period_label": state.get("period_label"), "as_of_date": state.get("as_of_date"), "requested_companies": list(COMPANY_NAMES), "source": ( dict(state["source"]) if isinstance(state.get("source"), Mapping) else None ), "created_at": state.get("created_at"), "started_at": state.get("started_at"), "finished_at": state.get("finished_at"), "duration_seconds": state.get("duration_seconds"), "progress": dict(state.get("progress") or {}), "message": state.get("message"), "failure_code": state.get("failure_code"), "queue_position": self._queue_position_locked(job_id), "company_results": companies, "downloads": downloads, "status_url": f"/api/company-reports/jobs/{job_id}", } @staticmethod def _safe_problem(value: Any) -> Dict[str, Any]: if not isinstance(value, Mapping): raise ValueError("problem must be an object") code = value.get("code") stage = value.get("stage") period = value.get("period") record_ids = value.get("record_ids", []) if not all( isinstance(item, str) and 0 < len(item) <= 100 for item in (code, stage, period) ): raise ValueError("problem fields are invalid") if not isinstance(record_ids, list) or len(record_ids) > 500: raise ValueError("problem record ids are invalid") safe_ids: List[int] = [] for item in record_ids: if isinstance(item, bool) or not isinstance(item, int) or item < 0: raise ValueError("problem record id is invalid") safe_ids.append(item) return { "code": code, "stage": stage, "period": period, "record_ids": safe_ids, } def _artifact_descriptor( self, company: str, report_year: int, report_month: int, value: Any, ) -> ArtifactDescriptor: if not isinstance(value, Mapping): raise ValueError("artifact is missing") filename = value.get("filename") storage_key = value.get("storage_key") sha256 = value.get("sha256") expected_filename = ( f"{company}-{ENGLISH_MONTH_NAMES[report_month]}-{report_year:04d}.xlsx" ) if ( filename != expected_filename or not isinstance(storage_key, str) or not isinstance(sha256, str) or SHA256_RE.fullmatch(sha256) is None ): raise ValueError("artifact metadata is invalid") storage = PurePosixPath(storage_key) if storage.is_absolute() or not storage.parts or ".." in storage.parts or "." in storage.parts: raise ValueError("artifact path is invalid") candidate = self._project_root.joinpath(*storage.parts) try: metadata = candidate.lstat() if candidate.is_symlink() or not stat.S_ISREG(metadata.st_mode): raise OSError("artifact is not a regular file") resolved = candidate.resolve(strict=True) resolved.relative_to(self._output_root) except (OSError, ValueError): raise ValueError("artifact escaped the controlled output") from None if resolved.name != expected_filename or resolved.suffix.lower() != ".xlsx": raise ValueError("artifact filename is invalid") if not 0 < metadata.st_size <= MAX_DOWNLOAD_BYTES: raise ValueError("artifact size is invalid") if _sha256_file(resolved) != sha256: raise ValueError("artifact identity is invalid") descriptor = ArtifactDescriptor( file_kind="company_ten_day_xlsx", original_filename=expected_filename, storage_key=storage_key, sha256=sha256, byte_size=metadata.st_size, mime_type=XLSX_MIME, ) descriptor.validate() return descriptor def _normalize_result( self, state: Mapping[str, Any], payload: Mapping[str, Any], ) -> Tuple[str, List[Dict[str, Any]], Dict[str, Dict[str, Any]]]: report_year, report_month = ( int(part) for part in str(state["report_month"]).split("-") ) if ( payload.get("schema_version") != RESULT_SCHEMA_VERSION or payload.get("status") not in {"success", "partial_failure", "failed"} or payload.get("report_year") != report_year or payload.get("report_month") != report_month or payload.get("as_of_date") != state["as_of_date"] or payload.get("requested_companies") != list(COMPANY_NAMES) or not isinstance(payload.get("companies"), list) ): raise ValueError("batch envelope is invalid") raw_companies = payload["companies"] if len(raw_companies) != len(COMPANY_NAMES): raise ValueError("company result count is invalid") by_name: Dict[str, Mapping[str, Any]] = {} for item in raw_companies: if not isinstance(item, Mapping) or item.get("company") not in COMPANY_NAMES: raise ValueError("company identity is invalid") company = str(item["company"]) if company in by_name: raise ValueError("company identity is duplicated") by_name[company] = item if set(by_name) != set(COMPANY_NAMES): raise ValueError("company identities are incomplete") expected_periods = { "01-10", "11-20", f"21-{calendar.monthrange(report_year, report_month)[1]:02d}", } normalized: List[Dict[str, Any]] = [] artifacts: Dict[str, Dict[str, Any]] = {} for company in COMPANY_NAMES: item = by_name[company] company_status = item.get("status") row_count = item.get("row_count") counts = item.get("period_row_counts") warnings = item.get("warnings", []) errors = item.get("errors", []) if ( company_status not in {"success", "failed"} or isinstance(row_count, bool) or not isinstance(row_count, int) or row_count < 0 or not isinstance(counts, Mapping) or set(counts) != expected_periods or not isinstance(warnings, list) or not isinstance(errors, list) ): raise ValueError("company result fields are invalid") safe_counts: Dict[str, int] = {} for key, count in counts.items(): if isinstance(count, bool) or not isinstance(count, int) or count < 0: raise ValueError("period row count is invalid") safe_counts[str(key)] = count safe_warnings = [self._safe_problem(value) for value in warnings] safe_errors = [self._safe_problem(value) for value in errors] version_no = item.get("version_no") if company_status == "success": if safe_errors or isinstance(version_no, bool) or not isinstance(version_no, int) or version_no <= 0: raise ValueError("successful company result is invalid") descriptor = self._artifact_descriptor( company, report_year, report_month, item.get("artifact"), ) artifacts[company] = asdict(descriptor) elif not safe_errors: raise ValueError("failed company has no error") normalized.append( { "company": company, "status": company_status, "row_count": row_count, "period_row_counts": safe_counts, "warnings": safe_warnings, "errors": safe_errors, "version_no": version_no if isinstance(version_no, int) else None, } ) successful = sum(item["status"] == "success" for item in normalized) expected_status = ( "success" if successful == len(COMPANY_NAMES) else "failed" if successful == 0 else "partial_failure" ) if payload["status"] != expected_status: raise ValueError("batch status is inconsistent") job_state = "succeeded" if expected_status == "success" else expected_status return job_state, normalized, artifacts def _fail_locked(self, job_id: str, code: str, message: str) -> None: state = self._read_locked(job_id) company_results = [ { "company": company, "status": "failed", "row_count": 0, "period_row_counts": {}, "warnings": [], "errors": [ { "code": code, "stage": "task", "period": state["report_month"], "record_ids": [], } ], "version_no": None, } for company in COMPANY_NAMES ] self._update_locked( job_id, state="failed", stage="finished", finished_at=_utc_now(), progress={"percent": 100, "label": "生成未完成"}, message=message, failure_code=code, company_results=company_results, artifact_files={}, ) def _process(self, job_id: str) -> None: started = time.monotonic() with self._condition: state = self._update_locked( job_id, state="running", stage="loading_source", started_at=_utc_now(), progress={"percent": 15, "label": "正在读取已校验的 finance 数据"}, message="正在按公司生成渠道明细表。", ) state = self._update_locked( job_id, stage="generating_workbooks", progress={"percent": 45, "label": "正在生成 5 家公司 Excel"}, ) try: report_year, report_month = ( int(part) for part in str(state["report_month"]).split("-") ) payload = self._executor.run( report_year, report_month, date.fromisoformat(str(state["as_of_date"])), ) if not isinstance(payload, Mapping): raise ValueError("executor result is not an object") job_state, companies, artifacts = self._normalize_result(state, payload) message = { "succeeded": "5 家公司渠道明细表已生成。", "partial_failure": "部分公司已生成,失败公司可按错误代码复核后重跑。", "failed": None, }[job_state] with self._condition: self._update_locked( job_id, state=job_state, stage="finished", finished_at=_utc_now(), duration_seconds=round(time.monotonic() - started, 3), progress={ "percent": 100, "label": "生成完成" if job_state == "succeeded" else "生成已结束", }, message=message, company_results=companies, artifact_files=artifacts, failure_code=None, ) except Exception: with self._condition: self._fail_locked( job_id, "COMPANY_REPORT_EXECUTION_FAILED", "渠道明细生成未完成;未成功发布的公司不会提供下载。", ) self._update_locked( job_id, duration_seconds=round(time.monotonic() - started, 3), ) def _worker(self) -> None: while True: with self._condition: while not self._pending and not self._stopping: self._condition.wait() if self._stopping: return job_id = self._pending.popleft() self._active_job_id = job_id try: self._process(job_id) except Exception: with self._condition: try: self._fail_locked( job_id, "COMPANY_REPORT_EXECUTION_FAILED", "渠道明细生成未完成,请重试。", ) except Exception: pass finally: with self._condition: self._active_job_id = None self._condition.notify_all()