1009 lines
36 KiB
Python
1009 lines
36 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import hashlib
|
||
|
|
import ipaddress
|
||
|
|
import json
|
||
|
|
import os
|
||
|
|
import socket
|
||
|
|
from contextlib import contextmanager
|
||
|
|
from dataclasses import asdict
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import Any, Iterator, Literal
|
||
|
|
from urllib.parse import urlsplit
|
||
|
|
|
||
|
|
import psycopg
|
||
|
|
from fastapi import (
|
||
|
|
APIRouter,
|
||
|
|
BackgroundTasks,
|
||
|
|
Body,
|
||
|
|
Depends,
|
||
|
|
File,
|
||
|
|
HTTPException,
|
||
|
|
Query,
|
||
|
|
UploadFile,
|
||
|
|
)
|
||
|
|
from psycopg.rows import dict_row
|
||
|
|
|
||
|
|
from app.modules.data_process.algorithms import (
|
||
|
|
chunk_unstructured,
|
||
|
|
decode_utf8,
|
||
|
|
desensitize_pii,
|
||
|
|
estimate_token_count,
|
||
|
|
generate_standard_records,
|
||
|
|
parse_text_content,
|
||
|
|
score_quality,
|
||
|
|
)
|
||
|
|
from app.modules.data_process.generation import generate_model_records
|
||
|
|
from app.modules.data_process.store import (
|
||
|
|
ConflictError,
|
||
|
|
DataProcessStore,
|
||
|
|
DataProcessStoreError,
|
||
|
|
InvalidStateError,
|
||
|
|
NotFoundError,
|
||
|
|
get_data_process_store,
|
||
|
|
)
|
||
|
|
from app.schemas.data_process import (
|
||
|
|
DataProcessTaskCreate,
|
||
|
|
DataProcessTaskUpdate,
|
||
|
|
DataProcessStatus,
|
||
|
|
ExternalPullRequest,
|
||
|
|
ExternalSourceRequest,
|
||
|
|
GenerateRequest,
|
||
|
|
PreviewBuildRequest,
|
||
|
|
PreviewItemCreate,
|
||
|
|
PreviewItemUpdate,
|
||
|
|
ProcessType,
|
||
|
|
PublishRequest,
|
||
|
|
ResultUpdate,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
router = APIRouter(prefix="/data-process")
|
||
|
|
MAX_SOURCE_FILE_BYTES = 200 * 1024 * 1024
|
||
|
|
MAX_SOURCE_FILE_COUNT = 20
|
||
|
|
MAX_SOURCE_BATCH_BYTES = 500 * 1024 * 1024
|
||
|
|
MAX_EXTERNAL_PULL_BYTES = 50 * 1024 * 1024
|
||
|
|
|
||
|
|
|
||
|
|
def ok(data: Any = None, message: str = "ok") -> dict[str, Any]:
|
||
|
|
return {"code": 0, "message": message, "data": data}
|
||
|
|
|
||
|
|
|
||
|
|
def fail(status_code: int, message: str) -> HTTPException:
|
||
|
|
return HTTPException(
|
||
|
|
status_code=status_code,
|
||
|
|
detail={"code": status_code, "message": message, "data": None},
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@contextmanager
|
||
|
|
def api_errors() -> Iterator[None]:
|
||
|
|
try:
|
||
|
|
yield
|
||
|
|
except NotFoundError as exc:
|
||
|
|
raise fail(404, str(exc)) from exc
|
||
|
|
except ConflictError as exc:
|
||
|
|
raise fail(409, str(exc)) from exc
|
||
|
|
except InvalidStateError as exc:
|
||
|
|
raise fail(409, str(exc)) from exc
|
||
|
|
except (DataProcessStoreError, ValueError) as exc:
|
||
|
|
raise fail(400, str(exc)) from exc
|
||
|
|
except psycopg.errors.UndefinedTable as exc:
|
||
|
|
raise fail(503, "data process schema is not installed; run schema_cli --check") from exc
|
||
|
|
except psycopg.OperationalError as exc:
|
||
|
|
raise fail(503, "data process database is unavailable") from exc
|
||
|
|
|
||
|
|
|
||
|
|
def _safe_file_name(value: str | None, fallback: str) -> str:
|
||
|
|
name = Path((value or "").replace("\\", "/")).name.replace("\x00", "").strip()
|
||
|
|
return name if name not in {"", ".", ".."} else fallback
|
||
|
|
|
||
|
|
|
||
|
|
def _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 _preprocess_options(config: dict[str, Any]) -> set[str]:
|
||
|
|
values = _value(config, "preprocess_options", "preprocessOptions", [])
|
||
|
|
return {str(item) for item in values} if isinstance(values, list) else set()
|
||
|
|
|
||
|
|
|
||
|
|
def _preview_quality(content: str, config: dict[str, Any]) -> dict[str, Any]:
|
||
|
|
records = generate_standard_records(
|
||
|
|
[{"id": "quality-preview", "edited_content": content}],
|
||
|
|
split={"train": 100, "validation": 0, "test": 0},
|
||
|
|
)
|
||
|
|
record = records[0] if records else {"instruction": "", "input": "", "output": ""}
|
||
|
|
minimum = int(_value(config, "min_output_length", "minOutputLength", 20) or 20)
|
||
|
|
return asdict(
|
||
|
|
score_quality(
|
||
|
|
record,
|
||
|
|
min_output_length=max(1, minimum),
|
||
|
|
source_content=content,
|
||
|
|
)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _build_preview_items(
|
||
|
|
task: dict[str, Any], source_files: list[dict[str, Any]]
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
config = task.get("config") or {}
|
||
|
|
process_type = task["process_type"]
|
||
|
|
preprocess_options = _preprocess_options(config)
|
||
|
|
should_desensitize = "desensitize" in preprocess_options
|
||
|
|
should_clean_invalid = bool(
|
||
|
|
preprocess_options & {"clean_invalid", "clean_invalid_content"}
|
||
|
|
)
|
||
|
|
should_deduplicate = bool(
|
||
|
|
preprocess_options & {"deduplicate", "deduplicate_content"}
|
||
|
|
)
|
||
|
|
seen_content_hashes: set[str] = set()
|
||
|
|
items: list[dict[str, Any]] = []
|
||
|
|
|
||
|
|
def append_item(item: dict[str, Any]) -> None:
|
||
|
|
content = str(item.get("edited_content") or "").strip()
|
||
|
|
if should_clean_invalid and not content:
|
||
|
|
return
|
||
|
|
content_hash = hashlib.sha256(content.encode("utf-8")).hexdigest()
|
||
|
|
if should_deduplicate and content_hash in seen_content_hashes:
|
||
|
|
return
|
||
|
|
seen_content_hashes.add(content_hash)
|
||
|
|
if not content:
|
||
|
|
item["status"] = "invalid"
|
||
|
|
items.append(item)
|
||
|
|
|
||
|
|
for source in source_files:
|
||
|
|
parsed = parse_text_content(
|
||
|
|
source.get("content") or "",
|
||
|
|
filename=source.get("name"),
|
||
|
|
file_format=source.get("file_format"),
|
||
|
|
)
|
||
|
|
if process_type == "unstructured":
|
||
|
|
chunks = chunk_unstructured(
|
||
|
|
parsed.text,
|
||
|
|
method=_value(config, "chunk_method", "chunkMethod", "semantic"),
|
||
|
|
chunk_size=int(_value(config, "chunk_size", "chunkSize", 800)),
|
||
|
|
chunk_overlap=int(_value(config, "chunk_overlap", "chunkOverlap", 100)),
|
||
|
|
min_chunk_size=int(_value(config, "min_chunk_size", "minChunkSize", 100)),
|
||
|
|
custom_delimiter=str(
|
||
|
|
_value(config, "custom_delimiter", "customDelimiter", "") or ""
|
||
|
|
),
|
||
|
|
preserve_code_blocks=bool(
|
||
|
|
_value(config, "preserve_code_blocks", "preserveCodeBlocks", False)
|
||
|
|
),
|
||
|
|
preserve_tables=bool(
|
||
|
|
_value(config, "preserve_tables", "preserveTables", False)
|
||
|
|
),
|
||
|
|
preserve_lists=bool(
|
||
|
|
_value(config, "preserve_lists", "preserveLists", False)
|
||
|
|
),
|
||
|
|
)
|
||
|
|
for chunk in chunks:
|
||
|
|
content = chunk.content
|
||
|
|
pii_counts: dict[str, int] = {}
|
||
|
|
if should_desensitize:
|
||
|
|
content, pii_counts = desensitize_pii(content)
|
||
|
|
quality = _preview_quality(content, config)
|
||
|
|
quality["pii_replacements"] = pii_counts
|
||
|
|
append_item(
|
||
|
|
{
|
||
|
|
"source_file_id": source["id"],
|
||
|
|
"original_content": chunk.content,
|
||
|
|
"edited_content": content,
|
||
|
|
"source_start": chunk.start,
|
||
|
|
"source_end": chunk.end,
|
||
|
|
"source_start_line": chunk.start_line,
|
||
|
|
"source_end_line": chunk.end_line,
|
||
|
|
"token_count": chunk.token_count,
|
||
|
|
"status": "modified" if content != chunk.content else "original",
|
||
|
|
"quality_score": quality,
|
||
|
|
}
|
||
|
|
)
|
||
|
|
continue
|
||
|
|
|
||
|
|
record_contents = [
|
||
|
|
json.dumps(record, ensure_ascii=False, separators=(",", ":"))
|
||
|
|
for record in parsed.records
|
||
|
|
if not should_clean_invalid
|
||
|
|
or any(value not in (None, "", [], {}) for value in record.values())
|
||
|
|
]
|
||
|
|
if not record_contents and parsed.text:
|
||
|
|
record_contents = [parsed.text]
|
||
|
|
for content in record_contents:
|
||
|
|
original_content = content
|
||
|
|
pii_counts = {}
|
||
|
|
if should_desensitize:
|
||
|
|
content, pii_counts = desensitize_pii(content)
|
||
|
|
quality = _preview_quality(content, config)
|
||
|
|
quality["pii_replacements"] = pii_counts
|
||
|
|
append_item(
|
||
|
|
{
|
||
|
|
"source_file_id": source["id"],
|
||
|
|
"original_content": original_content,
|
||
|
|
"edited_content": content,
|
||
|
|
"source_start": None,
|
||
|
|
"source_end": None,
|
||
|
|
"source_start_line": None,
|
||
|
|
"source_end_line": None,
|
||
|
|
"token_count": estimate_token_count(content),
|
||
|
|
"status": "modified" if content != original_content else "original",
|
||
|
|
"quality_score": quality,
|
||
|
|
}
|
||
|
|
)
|
||
|
|
return items
|
||
|
|
|
||
|
|
|
||
|
|
def _all_preview_items(store: DataProcessStore, task_id: str) -> list[dict[str, Any]]:
|
||
|
|
"""分页读取全部预览项,避免固定上限静默截断任务。"""
|
||
|
|
|
||
|
|
items: list[dict[str, Any]] = []
|
||
|
|
page = 1
|
||
|
|
page_size = 5_000
|
||
|
|
while True:
|
||
|
|
result = store.list_preview_items(task_id, page=page, page_size=page_size)
|
||
|
|
batch = result["items"]
|
||
|
|
items.extend(batch)
|
||
|
|
if len(items) >= int(result["total"]) or not batch:
|
||
|
|
return items
|
||
|
|
page += 1
|
||
|
|
|
||
|
|
|
||
|
|
def _run_generation(
|
||
|
|
store: DataProcessStore, task_id: str, generation_run_id: str
|
||
|
|
) -> None:
|
||
|
|
try:
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
if not store.generation_is_running(task_id, generation_run_id):
|
||
|
|
return
|
||
|
|
all_preview_items = _all_preview_items(store, task_id)
|
||
|
|
preview_items = [
|
||
|
|
item
|
||
|
|
for item in all_preview_items
|
||
|
|
if item.get("status") != "invalid"
|
||
|
|
and str(item.get("edited_content") or item.get("original_content") or "").strip()
|
||
|
|
]
|
||
|
|
pre_filtered_count = len(all_preview_items) - len(preview_items)
|
||
|
|
config = task.get("config") or {}
|
||
|
|
model_id = _value(config, "generation_model_id", "generationModelId", None)
|
||
|
|
generation_model: dict[str, Any] | None = None
|
||
|
|
if model_id:
|
||
|
|
generation_model = store.get_generation_model(str(model_id))
|
||
|
|
task = store.save_generation_model_snapshot(
|
||
|
|
task_id,
|
||
|
|
generation_model,
|
||
|
|
generation_run_id=generation_run_id,
|
||
|
|
)
|
||
|
|
config = task.get("config") or config
|
||
|
|
split = _value(
|
||
|
|
config,
|
||
|
|
"dataset_split",
|
||
|
|
"datasetSplit",
|
||
|
|
{"train": 80, "validation": 10, "test": 10},
|
||
|
|
)
|
||
|
|
pairs = (
|
||
|
|
_value(config, "qa_pairs_per_chunk", "qaPairsPerChunk", 1)
|
||
|
|
if task["process_type"] == "unstructured"
|
||
|
|
else _value(config, "qa_pairs_per_row", "qaPairsPerRow", 1)
|
||
|
|
)
|
||
|
|
if generation_model:
|
||
|
|
runtime_config = {
|
||
|
|
**config,
|
||
|
|
"generation_prompt": _value(
|
||
|
|
config, "generation_prompt", "generationPrompt", ""
|
||
|
|
),
|
||
|
|
"max_tokens": _value(config, "max_tokens", "maxTokens", 1024),
|
||
|
|
"json_mode": _value(config, "json_mode", "jsonMode", False),
|
||
|
|
}
|
||
|
|
def report_progress(processed_count: int, total_count: int) -> None:
|
||
|
|
if not store.update_generation_progress(
|
||
|
|
task_id,
|
||
|
|
generation_run_id,
|
||
|
|
processed_count,
|
||
|
|
total_count,
|
||
|
|
):
|
||
|
|
raise InvalidStateError("generation run is no longer active")
|
||
|
|
|
||
|
|
generated = generate_model_records(
|
||
|
|
preview_items,
|
||
|
|
model=generation_model,
|
||
|
|
config=runtime_config,
|
||
|
|
task_id=task_id,
|
||
|
|
split=split,
|
||
|
|
qa_pairs_per_item=int(pairs or 1),
|
||
|
|
on_progress=report_progress,
|
||
|
|
)
|
||
|
|
else:
|
||
|
|
generated = generate_standard_records(
|
||
|
|
preview_items,
|
||
|
|
qa_pairs_per_item=int(pairs or 1),
|
||
|
|
semantic_enrichment=bool(
|
||
|
|
_value(config, "semantic_enrichment", "semanticEnrichment", False)
|
||
|
|
),
|
||
|
|
split=split,
|
||
|
|
split_seed=task_id,
|
||
|
|
)
|
||
|
|
if not store.update_generation_progress(
|
||
|
|
task_id,
|
||
|
|
generation_run_id,
|
||
|
|
len(preview_items),
|
||
|
|
len(preview_items),
|
||
|
|
):
|
||
|
|
return
|
||
|
|
|
||
|
|
known_fingerprints: set[str] = set()
|
||
|
|
accepted: list[dict[str, Any]] = []
|
||
|
|
filtered_count = pre_filtered_count
|
||
|
|
duplicate_count = 0
|
||
|
|
error_count = 0
|
||
|
|
quality_filter = bool(
|
||
|
|
_value(config, "quality_filter_enabled", "qualityFilterEnabled", False)
|
||
|
|
)
|
||
|
|
filter_low = bool(_value(config, "filter_low_quality", "filterLowQuality", True))
|
||
|
|
filter_short = bool(
|
||
|
|
_value(config, "filter_short_content", "filterShortContent", True)
|
||
|
|
)
|
||
|
|
deduplicate = bool(
|
||
|
|
_preprocess_options(config)
|
||
|
|
& {"deduplicate", "deduplicate_content"}
|
||
|
|
)
|
||
|
|
minimum = max(1, int(_value(config, "min_output_length", "minOutputLength", 20) or 20))
|
||
|
|
preview_sources = {
|
||
|
|
str(item["id"]): str(
|
||
|
|
item.get("edited_content") or item.get("original_content") or ""
|
||
|
|
)
|
||
|
|
for item in preview_items
|
||
|
|
}
|
||
|
|
|
||
|
|
for record in generated:
|
||
|
|
quality = score_quality(
|
||
|
|
record,
|
||
|
|
min_output_length=minimum,
|
||
|
|
source_content=preview_sources.get(str(record.get("preview_item_id") or ""), ""),
|
||
|
|
known_fingerprints=known_fingerprints,
|
||
|
|
)
|
||
|
|
if "duplicate_record" in quality.flags:
|
||
|
|
duplicate_count += 1
|
||
|
|
else:
|
||
|
|
# 即使首条随后因短文本/低质量被过滤,也要阻止同批后续重复结果。
|
||
|
|
known_fingerprints.add(quality.fingerprint)
|
||
|
|
should_filter = (
|
||
|
|
(deduplicate and "duplicate_record" in quality.flags)
|
||
|
|
or (
|
||
|
|
quality_filter
|
||
|
|
and filter_short
|
||
|
|
and "output_too_short" in quality.flags
|
||
|
|
)
|
||
|
|
or (quality_filter and filter_low and not quality.is_valid)
|
||
|
|
)
|
||
|
|
if not quality.is_valid:
|
||
|
|
error_count += 1
|
||
|
|
record["status"] = "invalid"
|
||
|
|
record["error"] = ", ".join(quality.flags) or "quality validation failed"
|
||
|
|
if should_filter:
|
||
|
|
filtered_count += 1
|
||
|
|
continue
|
||
|
|
record["quality_score"] = asdict(quality)
|
||
|
|
accepted.append(record)
|
||
|
|
|
||
|
|
# stop 请求可能在纯函数计算期间到达,最终写入前再次检查状态。
|
||
|
|
if store.generation_is_running(task_id, generation_run_id):
|
||
|
|
store.complete_generation(
|
||
|
|
task_id,
|
||
|
|
accepted,
|
||
|
|
generation_run_id=generation_run_id,
|
||
|
|
filtered_count=filtered_count,
|
||
|
|
duplicate_count=duplicate_count,
|
||
|
|
error_count=error_count,
|
||
|
|
)
|
||
|
|
except Exception as exc: # noqa: BLE001 - background failures must be persisted
|
||
|
|
try:
|
||
|
|
if store.generation_is_running(task_id, generation_run_id):
|
||
|
|
store.mark_failed(
|
||
|
|
task_id,
|
||
|
|
str(exc),
|
||
|
|
generation_run_id=generation_run_id,
|
||
|
|
)
|
||
|
|
except Exception:
|
||
|
|
return
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("")
|
||
|
|
def list_tasks(
|
||
|
|
page: int = Query(default=1, ge=1),
|
||
|
|
page_size: int = Query(default=20, ge=1, le=200),
|
||
|
|
keyword: str | None = Query(default=None),
|
||
|
|
status: DataProcessStatus | None = Query(default=None),
|
||
|
|
process_type: ProcessType | None = Query(default=None),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
return ok(
|
||
|
|
store.list_tasks(
|
||
|
|
page=page,
|
||
|
|
page_size=page_size,
|
||
|
|
keyword=keyword,
|
||
|
|
status=status,
|
||
|
|
process_type=process_type,
|
||
|
|
)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("")
|
||
|
|
def create_task(
|
||
|
|
payload: DataProcessTaskCreate,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.create_task(payload.model_dump(mode="json"))
|
||
|
|
return ok(task, "data process task created")
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("/{task_id}")
|
||
|
|
def task_detail(
|
||
|
|
task_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
task["source_files"] = store.list_source_files(task_id)
|
||
|
|
return ok(task)
|
||
|
|
|
||
|
|
|
||
|
|
@router.put("/{task_id}")
|
||
|
|
def update_task(
|
||
|
|
task_id: str,
|
||
|
|
payload: DataProcessTaskUpdate,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
return ok(
|
||
|
|
store.update_task(task_id, payload.model_dump(exclude_unset=True, mode="json")),
|
||
|
|
"data process task updated",
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@router.delete("/{task_id}")
|
||
|
|
def delete_task(
|
||
|
|
task_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
store.delete_task(task_id)
|
||
|
|
return ok({"deleted": task_id}, "data process task deleted")
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("/{task_id}/source-files")
|
||
|
|
def source_files(
|
||
|
|
task_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
return ok({"files": store.list_source_files(task_id)})
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/source-files")
|
||
|
|
async def upload_source_files(
|
||
|
|
task_id: str,
|
||
|
|
files: list[UploadFile] = File(...),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
if not files:
|
||
|
|
raise fail(400, "at least one source file is required")
|
||
|
|
if len(files) > MAX_SOURCE_FILE_COUNT:
|
||
|
|
raise fail(413, f"a source batch may contain at most {MAX_SOURCE_FILE_COUNT} files")
|
||
|
|
prepared: list[dict[str, Any]] = []
|
||
|
|
batch_size = 0
|
||
|
|
with api_errors():
|
||
|
|
store.get_task(task_id)
|
||
|
|
for upload in files:
|
||
|
|
raw = await upload.read(MAX_SOURCE_FILE_BYTES + 1)
|
||
|
|
if len(raw) > MAX_SOURCE_FILE_BYTES:
|
||
|
|
raise fail(413, f"source file exceeds {MAX_SOURCE_FILE_BYTES} bytes")
|
||
|
|
content = decode_utf8(raw)
|
||
|
|
name = _safe_file_name(upload.filename, "source.txt")
|
||
|
|
suffix = Path(name).suffix.lower()
|
||
|
|
if suffix not in {
|
||
|
|
".txt",
|
||
|
|
".md",
|
||
|
|
".markdown",
|
||
|
|
".csv",
|
||
|
|
".tsv",
|
||
|
|
".json",
|
||
|
|
".jsonl",
|
||
|
|
".ndjson",
|
||
|
|
}:
|
||
|
|
raise fail(415, f"unsupported source file format: {suffix or 'none'}")
|
||
|
|
parsed = parse_text_content(content, filename=name)
|
||
|
|
if not parsed.text:
|
||
|
|
raise fail(400, f"source file is empty: {name}")
|
||
|
|
normalized_raw = parsed.text.encode("utf-8")
|
||
|
|
batch_size += len(normalized_raw)
|
||
|
|
if batch_size > MAX_SOURCE_BATCH_BYTES:
|
||
|
|
raise fail(413, f"source batch exceeds {MAX_SOURCE_BATCH_BYTES} bytes")
|
||
|
|
record_count = len(parsed.records) or (1 if parsed.text else 0)
|
||
|
|
prepared.append(
|
||
|
|
{
|
||
|
|
"name": name,
|
||
|
|
"content": parsed.text,
|
||
|
|
"raw_size": len(normalized_raw),
|
||
|
|
"checksum_sha256": hashlib.sha256(normalized_raw).hexdigest(),
|
||
|
|
"file_format": parsed.format,
|
||
|
|
"record_count": record_count,
|
||
|
|
"metadata": {
|
||
|
|
"content_type": upload.content_type or "text/plain",
|
||
|
|
"original_size_bytes": len(raw),
|
||
|
|
"original_checksum_sha256": hashlib.sha256(raw).hexdigest(),
|
||
|
|
},
|
||
|
|
"created_by": None,
|
||
|
|
}
|
||
|
|
)
|
||
|
|
created = store.add_source_files(task_id, prepared)
|
||
|
|
return ok({"files": created}, "source files uploaded")
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("/{task_id}/source-files/{file_id}/content")
|
||
|
|
def source_file_content(
|
||
|
|
task_id: str,
|
||
|
|
file_id: str,
|
||
|
|
start_line: int | None = Query(default=None, ge=1),
|
||
|
|
line_count: int = Query(default=200, ge=1, le=10_000),
|
||
|
|
offset: int = Query(default=0, ge=0),
|
||
|
|
limit: int = Query(default=100_000, ge=1, le=1_000_000),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
if start_line is not None:
|
||
|
|
return ok(store.source_content_lines(task_id, file_id, start_line, line_count))
|
||
|
|
return ok(store.source_content_window(task_id, file_id, offset, limit))
|
||
|
|
|
||
|
|
|
||
|
|
@router.delete("/{task_id}/source-files/{file_id}")
|
||
|
|
def delete_source_file(
|
||
|
|
task_id: str,
|
||
|
|
file_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
store.delete_source_file(task_id, file_id)
|
||
|
|
return ok({"deleted": file_id}, "source file removed")
|
||
|
|
|
||
|
|
|
||
|
|
def _external_postgres_connection(payload: ExternalSourceRequest) -> psycopg.Connection[Any]:
|
||
|
|
kind = payload.type.strip().lower()
|
||
|
|
parsed_url = urlsplit(payload.url)
|
||
|
|
scheme = parsed_url.scheme.lower()
|
||
|
|
if kind not in {"postgres", "postgresql"} or scheme not in {"postgres", "postgresql"}:
|
||
|
|
raise fail(501, f"external data source type is not supported: {payload.type}")
|
||
|
|
if parsed_url.username or parsed_url.password:
|
||
|
|
raise fail(400, "database credentials must use the account and password fields")
|
||
|
|
if payload.auth_mode not in {"none", "basic"}:
|
||
|
|
raise fail(400, "PostgreSQL supports only none or basic authentication")
|
||
|
|
if payload.auth_mode == "basic" and not payload.username:
|
||
|
|
raise fail(400, "database username is required for basic authentication")
|
||
|
|
hostname = parsed_url.hostname
|
||
|
|
if not hostname:
|
||
|
|
raise fail(400, "external PostgreSQL URL must include a hostname")
|
||
|
|
allow_private = os.getenv("DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB", "").lower() in {
|
||
|
|
"1",
|
||
|
|
"true",
|
||
|
|
"yes",
|
||
|
|
}
|
||
|
|
if not allow_private:
|
||
|
|
try:
|
||
|
|
addresses = {
|
||
|
|
item[4][0]
|
||
|
|
for item in socket.getaddrinfo(
|
||
|
|
hostname,
|
||
|
|
parsed_url.port or 5432,
|
||
|
|
type=socket.SOCK_STREAM,
|
||
|
|
)
|
||
|
|
}
|
||
|
|
except socket.gaierror as exc:
|
||
|
|
raise fail(400, "external PostgreSQL hostname cannot be resolved") from exc
|
||
|
|
if any(
|
||
|
|
(address := ipaddress.ip_address(value)).is_private
|
||
|
|
or address.is_loopback
|
||
|
|
or address.is_link_local
|
||
|
|
or address.is_reserved
|
||
|
|
or address.is_unspecified
|
||
|
|
for value in addresses
|
||
|
|
):
|
||
|
|
raise fail(
|
||
|
|
403,
|
||
|
|
"private or local database addresses are disabled; "
|
||
|
|
"set DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB=true only in a trusted deployment",
|
||
|
|
)
|
||
|
|
kwargs: dict[str, Any] = {
|
||
|
|
"connect_timeout": 5,
|
||
|
|
"row_factory": dict_row,
|
||
|
|
"application_name": "yg-ft-data-process-readonly",
|
||
|
|
"options": "-c default_transaction_read_only=on -c statement_timeout=30000",
|
||
|
|
}
|
||
|
|
if payload.auth_mode == "basic" and payload.username:
|
||
|
|
kwargs["user"] = payload.username
|
||
|
|
if payload.auth_mode == "basic" and payload.password:
|
||
|
|
kwargs["password"] = payload.password
|
||
|
|
return psycopg.connect(payload.url, **kwargs)
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/external/test")
|
||
|
|
def test_external_source(
|
||
|
|
task_id: str,
|
||
|
|
payload: ExternalSourceRequest,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
store.get_task(task_id)
|
||
|
|
try:
|
||
|
|
with _external_postgres_connection(payload) as conn:
|
||
|
|
conn.execute("SELECT 1 AS ok").fetchone()
|
||
|
|
except psycopg.Error as exc:
|
||
|
|
raise fail(502, "external PostgreSQL connection test failed") from exc
|
||
|
|
return ok({"connected": True, "type": payload.type})
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/external/pull")
|
||
|
|
def pull_external_source(
|
||
|
|
task_id: str,
|
||
|
|
payload: ExternalPullRequest,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
query = (payload.query or "").strip()
|
||
|
|
if query.endswith(";"):
|
||
|
|
query = query[:-1].rstrip()
|
||
|
|
if ";" in query:
|
||
|
|
raise fail(400, "external pull accepts exactly one read-only query")
|
||
|
|
first_token = query.split(maxsplit=1)[0].lower() if query else ""
|
||
|
|
if first_token not in {"select", "with"}:
|
||
|
|
raise fail(400, "a read-only SELECT or WITH query is required for external pull")
|
||
|
|
with api_errors():
|
||
|
|
store.get_task(task_id)
|
||
|
|
try:
|
||
|
|
with _external_postgres_connection(payload) as conn:
|
||
|
|
conn.execute("SET TRANSACTION READ ONLY")
|
||
|
|
conn.execute("SET LOCAL statement_timeout = '30s'")
|
||
|
|
cursor = conn.execute(query)
|
||
|
|
rows: list[dict[str, Any]] = []
|
||
|
|
content_parts: list[str] = []
|
||
|
|
content_size = 0
|
||
|
|
while len(rows) < payload.limit:
|
||
|
|
batch = cursor.fetchmany(min(1_000, payload.limit - len(rows)))
|
||
|
|
if not batch:
|
||
|
|
break
|
||
|
|
for row in batch:
|
||
|
|
line = json.dumps(row, ensure_ascii=False, default=str) + "\n"
|
||
|
|
content_size += len(line.encode("utf-8"))
|
||
|
|
if content_size > MAX_EXTERNAL_PULL_BYTES:
|
||
|
|
raise fail(413, "external pull result exceeds the 50 MiB safety limit")
|
||
|
|
rows.append(row)
|
||
|
|
content_parts.append(line)
|
||
|
|
conn.rollback()
|
||
|
|
except psycopg.Error as exc:
|
||
|
|
raise fail(502, "external PostgreSQL query failed") from exc
|
||
|
|
if not rows:
|
||
|
|
raise fail(400, "external query returned no rows")
|
||
|
|
content = "".join(content_parts)
|
||
|
|
raw = content.encode("utf-8")
|
||
|
|
source = store.add_source_file(
|
||
|
|
task_id,
|
||
|
|
name=_safe_file_name(payload.file_name, "external-data.jsonl"),
|
||
|
|
content=content,
|
||
|
|
raw_size=len(raw),
|
||
|
|
checksum_sha256=hashlib.sha256(raw).hexdigest(),
|
||
|
|
file_format="jsonl",
|
||
|
|
record_count=len(rows),
|
||
|
|
metadata={
|
||
|
|
"external_type": payload.type,
|
||
|
|
"external_host": urlsplit(payload.url).hostname,
|
||
|
|
"external_limit": payload.limit,
|
||
|
|
},
|
||
|
|
)
|
||
|
|
return ok({"files": [source]}, "external source pulled")
|
||
|
|
|
||
|
|
|
||
|
|
def _prepare_preview_items(
|
||
|
|
task_id: str,
|
||
|
|
store: DataProcessStore,
|
||
|
|
source_file_ids: list[str] | None = None,
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
source_summaries = store.list_source_files(task_id)
|
||
|
|
if source_file_ids is not None:
|
||
|
|
requested = set(source_file_ids)
|
||
|
|
source_summaries = [item for item in source_summaries if item["id"] in requested]
|
||
|
|
found = {item["id"] for item in source_summaries}
|
||
|
|
missing = requested - found
|
||
|
|
if missing:
|
||
|
|
raise NotFoundError(f"source files not found: {', '.join(sorted(missing))}")
|
||
|
|
sources = [
|
||
|
|
store.get_source_file(task_id, item["id"], include_content=True)
|
||
|
|
for item in source_summaries
|
||
|
|
]
|
||
|
|
if not sources:
|
||
|
|
raise InvalidStateError("at least one source file is required")
|
||
|
|
items = _build_preview_items(task, sources)
|
||
|
|
if not items:
|
||
|
|
raise InvalidStateError("source files did not produce preview items")
|
||
|
|
return items
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/preview/build")
|
||
|
|
def build_preview(
|
||
|
|
task_id: str,
|
||
|
|
payload: PreviewBuildRequest = Body(default_factory=PreviewBuildRequest),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
items = _prepare_preview_items(task_id, store, payload.source_file_ids)
|
||
|
|
created = store.replace_preview_items(task_id, items)
|
||
|
|
return ok(
|
||
|
|
{"items": created, "total": len(created), "page": 1, "page_size": len(created)},
|
||
|
|
"preview built",
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("/{task_id}/preview")
|
||
|
|
def preview_items(
|
||
|
|
task_id: str,
|
||
|
|
source_file_id: str | None = Query(default=None),
|
||
|
|
page: int = Query(default=1, ge=1),
|
||
|
|
page_size: int = Query(default=200, ge=1, le=1000),
|
||
|
|
keyword: str | None = Query(default=None),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
return ok(
|
||
|
|
store.list_preview_items(
|
||
|
|
task_id,
|
||
|
|
source_file_id=source_file_id,
|
||
|
|
page=page,
|
||
|
|
page_size=page_size,
|
||
|
|
keyword=keyword,
|
||
|
|
)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/preview")
|
||
|
|
def create_preview_item(
|
||
|
|
task_id: str,
|
||
|
|
payload: PreviewItemCreate,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
item_payload = payload.model_dump(mode="json")
|
||
|
|
content = payload.edited_content
|
||
|
|
item_payload["token_count"] = estimate_token_count(content)
|
||
|
|
item_payload["quality_score"] = _preview_quality(
|
||
|
|
content,
|
||
|
|
task.get("config") or {},
|
||
|
|
)
|
||
|
|
if not content.strip():
|
||
|
|
item_payload["status"] = "invalid"
|
||
|
|
item = store.create_preview_item(task_id, item_payload)
|
||
|
|
return ok(item, "preview item created")
|
||
|
|
|
||
|
|
|
||
|
|
@router.put("/{task_id}/preview/{preview_id}")
|
||
|
|
def update_preview_item(
|
||
|
|
task_id: str,
|
||
|
|
preview_id: str,
|
||
|
|
payload: PreviewItemUpdate,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
update = payload.model_dump(exclude_unset=True, mode="json")
|
||
|
|
update["quality_score"] = _preview_quality(payload.edited_content, task.get("config") or {})
|
||
|
|
item = store.update_preview_item(
|
||
|
|
task_id,
|
||
|
|
preview_id,
|
||
|
|
update,
|
||
|
|
)
|
||
|
|
return ok(item, "preview item updated")
|
||
|
|
|
||
|
|
|
||
|
|
@router.delete("/{task_id}/preview/{preview_id}")
|
||
|
|
def delete_preview_item(
|
||
|
|
task_id: str,
|
||
|
|
preview_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
store.delete_preview_item(task_id, preview_id)
|
||
|
|
return ok({"deleted": preview_id}, "preview item deleted")
|
||
|
|
|
||
|
|
|
||
|
|
def _start_generation(
|
||
|
|
task_id: str,
|
||
|
|
payload: GenerateRequest,
|
||
|
|
background_tasks: BackgroundTasks,
|
||
|
|
store: DataProcessStore,
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.start_generation(task_id, replace_existing=payload.replace_existing)
|
||
|
|
background_tasks.add_task(
|
||
|
|
_run_generation,
|
||
|
|
store,
|
||
|
|
task_id,
|
||
|
|
str(task["generation_run_id"]),
|
||
|
|
)
|
||
|
|
return ok(store.progress(task_id), "data process generation started")
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/generate")
|
||
|
|
def generate(
|
||
|
|
task_id: str,
|
||
|
|
background_tasks: BackgroundTasks,
|
||
|
|
payload: GenerateRequest = Body(default_factory=GenerateRequest),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
return _start_generation(task_id, payload, background_tasks, store)
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/start")
|
||
|
|
def start(
|
||
|
|
task_id: str,
|
||
|
|
background_tasks: BackgroundTasks,
|
||
|
|
payload: GenerateRequest = Body(default_factory=GenerateRequest),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
items = _prepare_preview_items(task_id, store)
|
||
|
|
store.replace_preview_items(task_id, items)
|
||
|
|
return _start_generation(task_id, payload, background_tasks, store)
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/stop")
|
||
|
|
def stop(
|
||
|
|
task_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
store.stop_task(task_id)
|
||
|
|
return ok(store.progress(task_id), "data process task stopped")
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("/{task_id}/progress")
|
||
|
|
def progress(
|
||
|
|
task_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
return ok(store.progress(task_id))
|
||
|
|
|
||
|
|
|
||
|
|
@router.get("/{task_id}/results")
|
||
|
|
def results(
|
||
|
|
task_id: str,
|
||
|
|
page: int = Query(default=1, ge=1),
|
||
|
|
page_size: int = Query(default=100, ge=1, le=1000),
|
||
|
|
status: Literal["valid", "modified", "invalid"] | None = Query(default=None),
|
||
|
|
split: Literal["train", "validation", "test"] | None = Query(default=None),
|
||
|
|
keyword: str | None = Query(default=None),
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
return ok(
|
||
|
|
store.list_results(
|
||
|
|
task_id,
|
||
|
|
page=page,
|
||
|
|
page_size=page_size,
|
||
|
|
status=status,
|
||
|
|
split=split,
|
||
|
|
keyword=keyword,
|
||
|
|
)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@router.put("/{task_id}/results/{result_id}")
|
||
|
|
def update_result(
|
||
|
|
task_id: str,
|
||
|
|
result_id: str,
|
||
|
|
payload: ResultUpdate,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
current = store.get_result(task_id, result_id)
|
||
|
|
update = payload.model_dump(exclude_unset=True, mode="json")
|
||
|
|
merged = {**current, **update}
|
||
|
|
preview_id = current.get("preview_item_id")
|
||
|
|
source_content = ""
|
||
|
|
if preview_id:
|
||
|
|
preview = store.get_preview_item(task_id, str(preview_id))
|
||
|
|
source_content = str(
|
||
|
|
preview.get("edited_content") or preview.get("original_content") or ""
|
||
|
|
)
|
||
|
|
minimum = max(
|
||
|
|
1,
|
||
|
|
int(
|
||
|
|
_value(
|
||
|
|
task.get("config") or {},
|
||
|
|
"min_output_length",
|
||
|
|
"minOutputLength",
|
||
|
|
20,
|
||
|
|
)
|
||
|
|
or 20
|
||
|
|
),
|
||
|
|
)
|
||
|
|
quality = score_quality(
|
||
|
|
merged,
|
||
|
|
min_output_length=minimum,
|
||
|
|
source_content=source_content,
|
||
|
|
)
|
||
|
|
update["quality_score"] = asdict(quality)
|
||
|
|
result = store.update_result(
|
||
|
|
task_id,
|
||
|
|
result_id,
|
||
|
|
update,
|
||
|
|
)
|
||
|
|
return ok(result, "data process result updated")
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/results/{result_id}/restore")
|
||
|
|
def restore_result(
|
||
|
|
task_id: str,
|
||
|
|
result_id: str,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
task = store.get_task(task_id)
|
||
|
|
current = store.get_result(task_id, result_id)
|
||
|
|
restored = {
|
||
|
|
**current,
|
||
|
|
"instruction": current.get("original_instruction") or current.get("instruction") or "",
|
||
|
|
"input": current.get("original_input") or current.get("input") or "",
|
||
|
|
"output": current.get("original_output") or current.get("output") or "",
|
||
|
|
}
|
||
|
|
preview_id = current.get("preview_item_id")
|
||
|
|
source_content = ""
|
||
|
|
if preview_id:
|
||
|
|
preview = store.get_preview_item(task_id, str(preview_id))
|
||
|
|
source_content = str(
|
||
|
|
preview.get("edited_content") or preview.get("original_content") or ""
|
||
|
|
)
|
||
|
|
minimum = max(
|
||
|
|
1,
|
||
|
|
int(
|
||
|
|
_value(
|
||
|
|
task.get("config") or {},
|
||
|
|
"min_output_length",
|
||
|
|
"minOutputLength",
|
||
|
|
20,
|
||
|
|
)
|
||
|
|
or 20
|
||
|
|
),
|
||
|
|
)
|
||
|
|
quality = score_quality(
|
||
|
|
restored,
|
||
|
|
min_output_length=minimum,
|
||
|
|
source_content=source_content,
|
||
|
|
)
|
||
|
|
restored = store.update_result(
|
||
|
|
task_id,
|
||
|
|
result_id,
|
||
|
|
{
|
||
|
|
"instruction": restored["instruction"],
|
||
|
|
"input": restored["input"],
|
||
|
|
"output": restored["output"],
|
||
|
|
"quality_score": asdict(quality),
|
||
|
|
"expected_updated_at": current.get("updated_at"),
|
||
|
|
},
|
||
|
|
)
|
||
|
|
return ok(restored, "data process result restored")
|
||
|
|
|
||
|
|
|
||
|
|
@router.post("/{task_id}/publish")
|
||
|
|
def publish(
|
||
|
|
task_id: str,
|
||
|
|
payload: PublishRequest,
|
||
|
|
store: DataProcessStore = Depends(get_data_process_store),
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
with api_errors():
|
||
|
|
result = store.publish(task_id, payload.model_dump(mode="json"))
|
||
|
|
message = "dataset published" if result["created"] else "dataset already published"
|
||
|
|
return ok(result, message)
|