"""Job / Experiment 的 SQLAlchemy 实现(Phase 4)。""" from __future__ import annotations from sqlalchemy import Select, func, select from sqlalchemy.orm import Session from app.domain.entities.research import ExperimentRecord, ExperimentSummary, JobRecord from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel def _job_to_model(job: JobRecord) -> JobModel: return JobModel(**job.model_dump()) class SqlAlchemyJobRepository: def __init__(self, session: Session) -> None: self._session = session def create(self, job: JobRecord) -> JobRecord: self._session.add(_job_to_model(job)) self._session.flush() return job def get(self, job_id: str) -> JobRecord | None: row = self._session.get(JobModel, job_id) return JobRecord.model_validate(row, from_attributes=True) if row else None def update(self, job: JobRecord) -> None: row = self._session.get(JobModel, job.id) if row is None: raise KeyError(f"job {job.id} 不存在") for k, v in job.model_dump().items(): setattr(row, k, v) def list_recent(self, kind: str | None = None, limit: int = 20) -> list[JobRecord]: stmt = select(JobModel).order_by(JobModel.created_at.desc()).limit(limit) if kind: stmt = stmt.where(JobModel.kind == kind) return [ JobRecord.model_validate(r, from_attributes=True) for r in self._session.scalars(stmt).all() ] def list_by_status(self, status: str, limit: int = 100) -> list[JobRecord]: rows = self._session.scalars( select(JobModel) .where(JobModel.status == status) .order_by(JobModel.created_at) .limit(limit) ).all() return [JobRecord.model_validate(r, from_attributes=True) for r in rows] class SqlAlchemyExperimentRepository: """Experiment 归档仓储。 `list_filtered` 刻意**不加载 `result_json`**(MEDIUMTEXT,完整存档后单条数 MB), 只 SELECT 元数据列 + SQL 侧算出的字符长度,避免列表接口把上百 MB 拉进内存。 """ def __init__(self, session: Session) -> None: self._session = session def save(self, experiment: ExperimentRecord) -> ExperimentRecord: self._session.add(ExperimentModel(**experiment.model_dump())) self._session.flush() return experiment def upsert(self, experiment: ExperimentRecord) -> ExperimentRecord: """按 id 插入或覆盖(`merge` 语义):重建归档时保持原 id 不变。""" self._session.merge(ExperimentModel(**experiment.model_dump())) self._session.flush() return experiment def get(self, experiment_id: str) -> ExperimentRecord | None: row = self._session.get(ExperimentModel, experiment_id) return ExperimentRecord.model_validate(row, from_attributes=True) if row else None def list_recent(self, limit: int = 50) -> list[ExperimentRecord]: rows = self._session.scalars( select(ExperimentModel).order_by(ExperimentModel.created_at.desc()).limit(limit) ).all() return [ExperimentRecord.model_validate(r, from_attributes=True) for r in rows] # -- 列表 / 过滤(SQL 层完成过滤+分页,参数绑定,不做字符串拼接)------------- @staticmethod def _escape_like(value: str) -> str: """转义 LIKE 元字符:用户输入里的 % / _ / \\ 不应被当成通配符。""" return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") def _filter_conditions(self, *, kind: str | None, q: str | None): """kind 精确匹配 + q 大小写不敏感模糊匹配(id / spec_json=因子名 / summary_text)。""" conditions = [] if kind: conditions.append(ExperimentModel.kind == kind) if q: pattern = f"%{self._escape_like(q.lower())}%" conditions.append( func.lower(ExperimentModel.id).like(pattern, escape="\\") | func.lower(ExperimentModel.spec_json).like(pattern, escape="\\") | func.lower(func.coalesce(ExperimentModel.summary_text, "")).like( pattern, escape="\\" ) ) return conditions def _summary_stmt(self, conditions) -> Select: return select( ExperimentModel.id, ExperimentModel.kind, ExperimentModel.spec_json, ExperimentModel.summary_text, ExperimentModel.code_version, ExperimentModel.data_version, ExperimentModel.job_id, ExperimentModel.created_at, self._char_length().label("result_bytes"), ).where(*conditions) def _char_length(self): """归档 JSON 的**字符数**(不取大字段本身)。 MySQL 的 `LENGTH()` 返回字节数,字符数要用 `CHAR_LENGTH()`;SQLite 只有 `length()`(对 TEXT 返回字符数)。按方言选择,避免中文归档在两个库上口径 不一致。 """ if self._session.get_bind().dialect.name == "sqlite": return func.length(ExperimentModel.result_json) return func.char_length(ExperimentModel.result_json) def list_filtered( self, *, kind: str | None = None, q: str | None = None, limit: int = 200, offset: int = 0, ) -> list[ExperimentSummary]: stmt = ( self._summary_stmt(self._filter_conditions(kind=kind, q=q)) # created_at 在 MySQL datetime(0) 下只有秒精度,同秒创建的归档必须有 # 稳定的次级排序键,否则分页会重复/漏项 .order_by(ExperimentModel.created_at.desc(), ExperimentModel.id.desc()) .limit(max(limit, 0)) .offset(max(offset, 0)) ) return [ ExperimentSummary( id=row.id, kind=row.kind, spec_json=row.spec_json, summary_text=row.summary_text, code_version=row.code_version, data_version=row.data_version, job_id=row.job_id, created_at=row.created_at, result_bytes=int(row.result_bytes or 0), ) for row in self._session.execute(stmt).all() ] def count_filtered(self, *, kind: str | None = None, q: str | None = None) -> int: stmt = ( select(func.count()) .select_from(ExperimentModel) .where(*self._filter_conditions(kind=kind, q=q)) ) return int(self._session.execute(stmt).scalar() or 0) def delete(self, experiment_id: str) -> bool: """删除归档本身;返回 False 表示不存在。不触碰 job 表。""" row = self._session.get(ExperimentModel, experiment_id) if row is None: return False self._session.delete(row) self._session.flush() return True