from __future__ import annotations from enum import StrEnum from typing import Any, Literal from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from app.modules.data_process.constants import MAX_QA_PAIRS_PER_ITEM def _config_value(config: dict[str, Any], snake_name: str, camel_name: str, default: Any) -> Any: if snake_name in config: return config[snake_name] return config.get(camel_name, default) def _validate_process_config(config: dict[str, Any]) -> None: chunk_method = _config_value(config, "chunk_method", "chunkMethod", "layout_hybrid") if not isinstance(chunk_method, str) or chunk_method not in { "layout_hybrid", "semantic", "fixed", }: raise ValueError("chunk_method must be one of: layout_hybrid, semantic, fixed") semantic_percentile = _config_value( config, "semantic_breakpoint_percentile", "semanticBreakpointPercentile", 95, ) if ( isinstance(semantic_percentile, bool) or not isinstance(semantic_percentile, int) or not 1 <= semantic_percentile <= 99 ): raise ValueError("semantic_breakpoint_percentile must be an integer in [1, 99]") split = _config_value(config, "dataset_split", "datasetSplit", None) if split is not None: if not isinstance(split, dict) or set(split) != {"train", "validation", "test"}: raise ValueError("dataset_split must contain train, validation and test") values = list(split.values()) if any(isinstance(value, bool) or not isinstance(value, int) for value in values): raise ValueError("dataset_split values must be integers") if any(value < 0 or value > 100 for value in values) or sum(values) != 100: raise ValueError("dataset_split values must be in [0, 100] and total 100") chunk_fields = { "chunk_size", "chunkSize", "chunk_overlap", "chunkOverlap", "min_chunk_size", "minChunkSize", } if chunk_fields.intersection(config): chunk_size = _config_value(config, "chunk_size", "chunkSize", 800) overlap = _config_value(config, "chunk_overlap", "chunkOverlap", 100) minimum = _config_value(config, "min_chunk_size", "minChunkSize", 100) if any( isinstance(value, bool) or not isinstance(value, int) for value in (chunk_size, overlap, minimum) ): raise ValueError("chunk_size, chunk_overlap and min_chunk_size must be integers") if not 16 <= chunk_size <= 32_768: raise ValueError("chunk_size must be in [16, 32768]") if overlap < 0 or overlap >= chunk_size: raise ValueError("chunk_overlap must be in [0, chunk_size)") if minimum <= 0 or minimum > chunk_size or overlap + minimum > chunk_size: raise ValueError("min_chunk_size and chunk_overlap exceed chunk_size") temperature = _config_value(config, "temperature", "temperature", None) if temperature is not None: if isinstance(temperature, bool) or not isinstance(temperature, (int, float)): raise ValueError("temperature must be a number") if not 0 <= float(temperature) <= 2: raise ValueError("temperature must be in [0, 2]") max_tokens = _config_value(config, "max_tokens", "maxTokens", None) if max_tokens is not None: if isinstance(max_tokens, bool) or not isinstance(max_tokens, int): raise ValueError("max_tokens must be an integer") if not 1 <= max_tokens <= 32_768: raise ValueError("max_tokens must be in [1, 32768]") for snake_name, camel_name in ( ("qa_pairs_per_row", "qaPairsPerRow"), ("qa_pairs_per_chunk", "qaPairsPerChunk"), ): pairs = _config_value(config, snake_name, camel_name, None) if pairs is None: continue if ( isinstance(pairs, bool) or not isinstance(pairs, int) or not 1 <= pairs <= MAX_QA_PAIRS_PER_ITEM ): raise ValueError( f"{snake_name} must be an integer in [1, {MAX_QA_PAIRS_PER_ITEM}]" ) class DataProcessStatus(StrEnum): pending = "pending" running = "running" completed = "completed" failed = "failed" stopped = "stopped" class ProcessType(StrEnum): structured = "structured" unstructured = "unstructured" external = "external" class DataProcessTaskCreate(BaseModel): model_config = ConfigDict(extra="forbid") name: str = Field(min_length=1, max_length=150) description: str = "" process_type: ProcessType source_dataset_id: str | None = None config: dict[str, Any] = Field(default_factory=dict) @field_validator("name") @classmethod def normalize_name(cls, value: str) -> str: value = value.strip() if not value: raise ValueError("task name cannot be empty") return value @model_validator(mode="after") def validate_config(self) -> "DataProcessTaskCreate": _validate_process_config(self.config) return self class DataProcessTaskUpdate(BaseModel): model_config = ConfigDict(extra="forbid") name: str | None = Field(default=None, min_length=1, max_length=150) description: str | None = None process_type: ProcessType | None = None source_dataset_id: str | None = None config: dict[str, Any] | None = None @field_validator("name") @classmethod def normalize_name(cls, value: str | None) -> str | None: if value is None: return None value = value.strip() if not value: raise ValueError("task name cannot be empty") return value @model_validator(mode="after") def validate_config(self) -> "DataProcessTaskUpdate": if self.config is not None: _validate_process_config(self.config) return self class DataProcessRegenerateRequest(BaseModel): """以一份完整配置准备任务重新生成。 ``expected_updated_at`` 用于防止详情页的旧快照覆盖其他人刚刚 保存的配置。重新生成不允许改变处理类型,避免旧源文件在新解析 规则下被静默误用。 """ model_config = ConfigDict(extra="forbid") name: str = Field(min_length=1, max_length=150) description: str process_type: ProcessType config: dict[str, Any] expected_updated_at: str = Field(min_length=1) @field_validator("name") @classmethod def normalize_name(cls, value: str) -> str: value = value.strip() if not value: raise ValueError("task name cannot be empty") return value @model_validator(mode="after") def validate_config(self) -> "DataProcessRegenerateRequest": _validate_process_config(self.config) return self class PreviewBuildRequest(BaseModel): model_config = ConfigDict(extra="forbid") replace_existing: Literal[True] = True source_file_ids: list[str] | None = None source_file_id: str | None = None @model_validator(mode="after") def validate_source_file_selection(self) -> "PreviewBuildRequest": if self.source_file_ids is not None and self.source_file_id is not None: raise ValueError("source_file_id and source_file_ids cannot be used together") values = self.source_file_ids if values is None and self.source_file_id is not None: values = [self.source_file_id] if values is None: return self normalized = list(dict.fromkeys(str(value).strip() for value in values)) if not normalized or any(not value for value in normalized): raise ValueError("at least one non-empty source file id is required") self.source_file_ids = normalized self.source_file_id = None return self class PreviewItemCreate(BaseModel): model_config = ConfigDict(extra="forbid") source_file_id: str | None = None original_content: str = "" edited_content: str = "" source_start: int | None = Field(default=None, ge=0) source_end: int | None = Field(default=None, ge=0) source_start_line: int | None = Field(default=None, ge=1) source_end_line: int | None = Field(default=None, ge=1) @model_validator(mode="after") def validate_ranges(self) -> "PreviewItemCreate": if self.source_start is not None and self.source_end is not None: if self.source_end < self.source_start: raise ValueError("source_end must be greater than or equal to source_start") if self.source_start_line is not None and self.source_end_line is not None: if self.source_end_line < self.source_start_line: raise ValueError( "source_end_line must be greater than or equal to source_start_line" ) return self class PreviewItemUpdate(BaseModel): model_config = ConfigDict(extra="forbid") edited_content: str expected_updated_at: str | None = None class GenerateRequest(BaseModel): model_config = ConfigDict(extra="forbid") replace_existing: Literal[True] = True class ExternalSourceRequest(BaseModel): model_config = ConfigDict(extra="forbid") type: str = Field(min_length=1, max_length=30) url: str = Field(min_length=1, max_length=2048) auth_mode: Literal["none", "basic"] = "none" username: str | None = Field(default=None, max_length=150) password: str | None = Field(default=None, max_length=500) limit: int = Field(default=1000, ge=1, le=100_000) class ExternalPullRequest(ExternalSourceRequest): query: str | None = Field(default=None, max_length=20_000) file_name: str = Field(default="external-data.jsonl", min_length=1, max_length=255) @field_validator("file_name") @classmethod def validate_file_name(cls, value: str) -> str: name = value.strip() if not name.lower().endswith((".jsonl", ".ndjson")): raise ValueError("external pull file_name must end with .jsonl or .ndjson") return name class ResultUpdate(BaseModel): model_config = ConfigDict(extra="forbid") instruction: str | None = None input: str | None = None output: str | None = None expected_updated_at: str | None = None class DatasetSplit(BaseModel): model_config = ConfigDict(extra="forbid") train: int = Field(default=80, ge=0, le=100) validation: int = Field(default=10, ge=0, le=100) test: int = Field(default=10, ge=0, le=100) @model_validator(mode="after") def validate_total(self) -> "DatasetSplit": if self.train + self.validation + self.test != 100: raise ValueError("dataset split must total 100") return self class PublishRequest(BaseModel): model_config = ConfigDict(extra="forbid") dataset_name: str = Field(min_length=1, max_length=150) dataset_type: Literal["train", "test", "eval", "val", "other"] = "train" storage_type: Literal["local"] = "local" split: DatasetSplit = Field(default_factory=DatasetSplit) format: Literal["alpaca_jsonl", "jsonl"] = "alpaca_jsonl" description: str = "" @field_validator("dataset_name") @classmethod def normalize_dataset_name(cls, value: str) -> str: value = value.strip() if not value: raise ValueError("dataset name cannot be empty") return value