Files
YG_FT/backend/app/schemas/data_process.py

374 lines
13 KiB
Python
Raw Normal View History

2026-07-30 14:38:56 +08:00
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 DataProcessWorkflowStep(StrEnum):
create = "create"
model = "model"
upload = "upload"
preview = "preview"
generate = "generate"
results = "results"
class DataProcessPreviewStatus(StrEnum):
idle = "idle"
queued = "queued"
running = "running"
completed = "completed"
failed = "failed"
cancelled = "cancelled"
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 DataProcessWorkflowStepUpdate(BaseModel):
"""仅保存创建向导位置,不修改配置或使下游产物失效。"""
model_config = ConfigDict(extra="forbid")
workflow_step: DataProcessWorkflowStep
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 ResultRegenerateRequest(BaseModel):
model_config = ConfigDict(extra="forbid")
expected_updated_at: str = Field(min_length=1, max_length=100)
class ResultBatchRegenerateItem(BaseModel):
model_config = ConfigDict(extra="forbid")
result_id: str = Field(min_length=1, max_length=100)
expected_updated_at: str = Field(min_length=1, max_length=100)
class ResultBatchRegenerateRequest(BaseModel):
model_config = ConfigDict(extra="forbid")
items: list[ResultBatchRegenerateItem] = Field(min_length=1, max_length=100)
@model_validator(mode="after")
def validate_unique_results(self) -> "ResultBatchRegenerateRequest":
result_ids = [item.result_id for item in self.items]
if len(result_ids) != len(set(result_ids)):
raise ValueError("result_id values must be unique")
return self
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