squid-main-space / fisheries_hf.py
mengxaingshuo's picture
fix: safe multi-day export and unified fisheries fallback
8312ded
Raw
History Blame Contribute Delete
25.6 kB
"""Restricted Hugging Face fisheries file reader and CSV export helpers.
The module deliberately exposes no arbitrary URL, shell, or filesystem access.
Callers must first validate a path against the configured Dataset live tree and
pass the expected file size to :func:`download_dataset_file`.
"""
from __future__ import annotations
import csv
import hashlib
import io
import itertools
import json
import os
import re
import secrets
import shutil
import time
import zipfile
from collections import defaultdict
from pathlib import Path
from typing import Any, Iterable
from urllib.parse import quote
import httpx
HF_SQUID_DATASET_REPO = (
os.environ.get("HF_SQUID_DATASET_REPO")
or os.environ.get("HF_DATASET_REPO")
or "globalsquiddatabase/squid_dataset"
).strip()
HF_TUNA_DATASET_REPO = (
os.environ.get("HF_TUNA_DATASET_REPO")
or "globalsquiddatabase/Tuna-Fisheries-Dataset"
).strip()
HF_DATASET_REPOS = {
"squid": HF_SQUID_DATASET_REPO,
"tuna": HF_TUNA_DATASET_REPO,
}
# Backwards-compatible default used when old callers omit a repository.
HF_DATASET_REPO = HF_SQUID_DATASET_REPO
LOCAL_SQUID_DATA_ROOT = Path(os.environ.get("LOCAL_SQUID_DATA_ROOT", "").strip()).expanduser() if os.environ.get("LOCAL_SQUID_DATA_ROOT", "").strip() else None
LOCAL_TUNA_DATA_ROOT = Path(os.environ.get("LOCAL_TUNA_DATA_ROOT", "").strip()).expanduser() if os.environ.get("LOCAL_TUNA_DATA_ROOT", "").strip() else None
LOCAL_DATA_ROOTS = {"squid": LOCAL_SQUID_DATA_ROOT, "tuna": LOCAL_TUNA_DATA_ROOT}
HF_SQUID_DATASET_REVISION = (
os.environ.get("HF_SQUID_DATASET_REVISION")
or os.environ.get("HF_DATASET_REVISION")
or ""
).strip()
HF_TUNA_DATASET_REVISION = os.environ.get("HF_TUNA_DATASET_REVISION", "").strip()
HF_TOKEN = os.environ.get("HF_TOKEN", "").strip()
PUBLIC_BASE_URL = os.environ.get("PUBLIC_BASE_URL", "").strip().rstrip("/")
CACHE_ROOT = Path(
os.environ.get("HF_FISHERIES_CACHE_ROOT", "/tmp/squid_hf_fisheries_cache")
)
EXPORT_ROOT = Path(
os.environ.get("FISHERIES_EXPORT_ROOT", "/tmp/squid_fisheries_exports")
)
CACHE_ROOT.mkdir(parents=True, exist_ok=True)
EXPORT_ROOT.mkdir(parents=True, exist_ok=True)
MAX_CSV_BYTES = int(os.environ.get("HF_FISHERIES_MAX_CSV_BYTES", 100 * 1024 * 1024))
MAX_ZIP_BYTES = int(os.environ.get("HF_FISHERIES_MAX_ZIP_BYTES", 1536 * 1024 * 1024))
EXPORT_TTL_SECONDS = int(os.environ.get("FISHERIES_EXPORT_TTL_SECONDS", "86400"))
_REVISION_CACHE: dict[str, dict[str, Any]] = {}
YEAR_ALIASES = ("year", "yearc", "年份", "yr")
MONTH_ALIASES = ("month", "月份", "mon")
DATE_ALIASES = ("date", "time", "datetime", "日期", "时间", "year_month")
LON_ALIASES = ("lon", "longitude", "decimal_longitude", "经度", "x")
LAT_ALIASES = ("lat", "latitude", "decimal_latitude", "纬度", "y")
def _headers() -> dict[str, str]:
return {"Authorization": f"Bearer {HF_TOKEN}"} if HF_TOKEN else {}
def normalize_repository(repository: str | None = None) -> str:
"""Resolve a safe repository selector to one of the two configured repos."""
value = str(repository or "squid").strip()
lowered = value.lower()
if lowered in HF_DATASET_REPOS:
return HF_DATASET_REPOS[lowered]
for repo in HF_DATASET_REPOS.values():
if value == repo:
return repo
raise ValueError("repository must be squid, tuna, or an exact configured repository id")
def local_root_for_repository(repository: str | None = None) -> Path | None:
"""Return the configured local mirror root for a logical repository."""
value = str(repository or "squid").strip()
domain = value.lower() if value.lower() in LOCAL_DATA_ROOTS else None
if domain is None:
for key, repo in HF_DATASET_REPOS.items():
if value == repo:
domain = key
break
root = LOCAL_DATA_ROOTS.get(domain or "")
if root and root.is_dir():
return root.resolve()
return None
def local_dataset_tree(repository: str | None = None) -> list[dict[str, Any]]:
"""List supported files from a local school-server mirror.
The mirror is a read-through cache, not an alternative dataset. Callers
merge this list with the Hugging Face live tree and prefer these entries
only when the same path exists in both places.
"""
root = local_root_for_repository(repository)
if root is None:
return []
items: list[dict[str, Any]] = []
for path in root.rglob("*"):
if not path.is_file() or path.suffix.lower() not in {".csv", ".tsv", ".zip"}:
continue
items.append({
"path": path.relative_to(root).as_posix(),
"type": "file",
"size": path.stat().st_size,
"storage_origin": "school_local_mirror",
})
return sorted(items, key=lambda item: str(item["path"]))
def merge_local_with_remote_tree(
remote_items: Iterable[dict[str, Any]], repository: str | None = None
) -> list[dict[str, Any]]:
"""Combine the authoritative HF tree with an optional local read cache.
Every remote file remains visible. A local file with the *same* safe
relative path replaces the remote entry so content reads take the local
route first; files absent locally keep their Hugging Face origin and can
be downloaded on demand. This gives the school and HF deployments the
same inventory instead of hiding remote-only files on the school server.
"""
merged: dict[str, dict[str, Any]] = {}
for item in remote_items:
if not isinstance(item, dict):
continue
path = str(item.get("path") or "").strip().lstrip("/")
if not path:
continue
row = dict(item)
row["path"] = path
row.setdefault("storage_origin", "huggingface")
merged[path] = row
for item in local_dataset_tree(repository):
path = str(item.get("path") or "").strip()
if path:
merged[path] = dict(item)
return [merged[path] for path in sorted(merged)]
def _revision(repository: str) -> str:
pinned = (
HF_TUNA_DATASET_REVISION
if repository == HF_TUNA_DATASET_REPO
else HF_SQUID_DATASET_REVISION
)
if pinned:
return pinned
now = time.time()
cache = _REVISION_CACHE.get(repository) or {}
cached = str(cache.get("sha") or "")
if cached and now - float(cache.get("ts") or 0) < 300:
return cached
url = f"https://huggingface.co/api/datasets/{repository}"
with httpx.Client(timeout=30.0, follow_redirects=True) as client:
response = client.get(url, headers=_headers())
if response.status_code in {401, 403}:
raise RuntimeError("Hugging Face Dataset 无读取权限,请检查 HF_TOKEN。")
response.raise_for_status()
body = response.json()
sha = str(body.get("sha") or "").strip()
if not re.fullmatch(r"[0-9a-fA-F]{40}", sha):
raise RuntimeError("无法获得 Hugging Face Dataset 的固定提交 SHA。")
_REVISION_CACHE[repository] = {"ts": now, "sha": sha}
return sha
def _safe_cache_path(repository: str, path: str, revision: str) -> Path:
suffix = Path(path).suffix.lower()
digest = hashlib.sha256(f"{repository}:{revision}:{path}".encode()).hexdigest()
return CACHE_ROOT / f"{digest}{suffix}"
def download_dataset_file(
path: str,
expected_size: int,
repository: str | None = None,
) -> tuple[Path, str]:
"""Download one validated Dataset file to a revision-keyed local cache."""
clean = str(path or "").strip().lstrip("/")
if not clean or "\x00" in clean or any(part in {"", ".", ".."} for part in clean.split("/")):
raise ValueError("invalid dataset path")
suffix = Path(clean).suffix.lower()
limit = MAX_ZIP_BYTES if suffix == ".zip" else MAX_CSV_BYTES
if suffix not in {".csv", ".tsv", ".zip"}:
raise ValueError("当前内容查询仅支持 CSV、TSV 和包含 CSV/TSV 的 ZIP。")
if expected_size <= 0:
raise ValueError("live tree did not provide a positive file size")
if expected_size > limit:
raise ValueError(
f"文件大小 {expected_size} bytes 超过在线处理上限 {limit} bytes;"
"请先在学校服务器生成查询就绪的分区文件。"
)
repository = normalize_repository(repository)
local_root = local_root_for_repository(repository)
if local_root is not None:
candidate = (local_root / clean).resolve()
if local_root not in candidate.parents and candidate != local_root:
raise ValueError("invalid local dataset path")
if candidate.is_file():
actual_size = candidate.stat().st_size
if expected_size and actual_size != expected_size:
raise RuntimeError(f"本地文件大小不一致:expected={expected_size}, actual={actual_size}")
return candidate, "local"
# A partial local mirror must never make HF-only files unreachable.
# Continue to the revision-pinned remote download below.
revision = _revision(repository)
target = _safe_cache_path(repository, clean, revision)
if target.exists() and target.stat().st_size == expected_size:
return target, revision
partial = target.with_suffix(target.suffix + ".part")
partial.unlink(missing_ok=True)
encoded_path = quote(clean, safe="/")
url = (
f"https://huggingface.co/datasets/{repository}/resolve/"
f"{revision}/{encoded_path}"
)
total = 0
try:
with httpx.stream(
"GET",
url,
headers=_headers(),
follow_redirects=True,
timeout=httpx.Timeout(connect=20, read=300, write=30, pool=30),
) as response:
if response.status_code in {401, 403}:
raise RuntimeError("Hugging Face 文件无读取权限,请检查 HF_TOKEN。")
response.raise_for_status()
with partial.open("wb") as stream:
for chunk in response.iter_bytes(1024 * 1024):
if not chunk:
continue
total += len(chunk)
if total > limit:
raise ValueError("download exceeded configured size limit")
stream.write(chunk)
if total != expected_size:
raise RuntimeError(
f"文件下载不完整:expected={expected_size}, received={total}"
)
partial.replace(target)
except Exception:
partial.unlink(missing_ok=True)
raise
return target, revision
def _decode_text(path: Path) -> tuple[str, str]:
raw = path.read_bytes()
for encoding in ("utf-8-sig", "utf-8", "gb18030"):
try:
return raw.decode(encoding), encoding
except UnicodeDecodeError:
pass
return raw.decode("utf-8", errors="replace"), "utf-8-replace"
def _dialect(text: str, suffix: str = ".csv") -> str:
sample = text[:20000]
try:
return csv.Sniffer().sniff(sample, delimiters=",\t;|").delimiter
except Exception:
return "\t" if suffix == ".tsv" else ","
def _column(columns: Iterable[str], aliases: Iterable[str]) -> str | None:
exact = {str(col).strip().lower(): str(col) for col in columns}
for alias in aliases:
if alias.lower() in exact:
return exact[alias.lower()]
return None
def _number(value: Any) -> float | None:
text = str(value or "").strip().replace(",", "")
if not text:
return None
try:
return float(text)
except Exception:
return None
def _year_month(row: dict[str, str], columns: list[str]) -> tuple[int | None, int | None]:
year_col = _column(columns, YEAR_ALIASES)
month_col = _column(columns, MONTH_ALIASES)
date_col = _column(columns, DATE_ALIASES)
year = None
month = None
if year_col:
value = _number(row.get(year_col))
if value is not None and 1800 <= int(value) <= 2200:
year = int(value)
if month_col:
value = _number(row.get(month_col))
if value is not None and 1 <= int(value) <= 12:
month = int(value)
if date_col and (year is None or month is None):
text = str(row.get(date_col) or "")
match = re.search(r"(19\d{2}|20\d{2}|21\d{2})[-/]?(0?[1-9]|1[0-2])?", text)
if match:
year = year or int(match.group(1))
month = month or (int(match.group(2)) if match.group(2) else None)
return year, month
def _metric_columns(columns: list[str], requested: str | None = None) -> list[str]:
if requested:
wanted = [x.strip() for x in requested.split(",") if x.strip()]
missing = [x for x in wanted if x not in columns]
if missing:
raise ValueError(f"requested metric columns not found: {missing}")
return wanted
keys = (
"effort", "fishing_hour", "fishing hours", "apparent_fishing",
"catch", "harvest", "landing", "cpue", "value",
)
excluded = set(YEAR_ALIASES + MONTH_ALIASES + LON_ALIASES + LAT_ALIASES)
return [
col for col in columns
if str(col).strip().lower() not in excluded
and any(key in str(col).strip().lower() for key in keys)
]
def _matches_filters(
row: dict[str, str],
columns: list[str],
*,
year: int | None,
lon_min: float | None,
lon_max: float | None,
lat_min: float | None,
lat_max: float | None,
) -> tuple[bool, int | None, int | None]:
row_year, row_month = _year_month(row, columns)
if year is not None and row_year != int(year):
return False, row_year, row_month
lon_col = _column(columns, LON_ALIASES)
lat_col = _column(columns, LAT_ALIASES)
if any(v is not None for v in (lon_min, lon_max)):
if not lon_col:
raise ValueError("经度筛选已请求,但文件中未识别到经度字段。")
lon = _number(row.get(lon_col))
if lon is None or (lon_min is not None and lon < lon_min) or (lon_max is not None and lon > lon_max):
return False, row_year, row_month
if any(v is not None for v in (lat_min, lat_max)):
if not lat_col:
raise ValueError("纬度筛选已请求,但文件中未识别到纬度字段。")
lat = _number(row.get(lat_col))
if lat is None or (lat_min is not None and lat < lat_min) or (lat_max is not None and lat > lat_max):
return False, row_year, row_month
return True, row_year, row_month
def _create_export(filename: str, rows: Iterable[dict[str, Any]], columns: list[str]) -> dict[str, str]:
now = time.time()
for item in list(EXPORT_ROOT.iterdir())[:2000]:
if not item.is_dir():
continue
try:
meta = json.loads((item / "meta.json").read_text(encoding="utf-8"))
expired = float(meta.get("expires_ts") or 0) < now
except Exception:
expired = True
if expired:
shutil.rmtree(item, ignore_errors=True)
token = secrets.token_urlsafe(24)
safe_name = re.sub(r"[^A-Za-z0-9._-]+", "_", filename).strip("._") or "fisheries_export.csv"
folder = EXPORT_ROOT / token
folder.mkdir(parents=True, exist_ok=False)
target = folder / safe_name
with target.open("w", encoding="utf-8-sig", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=columns, extrasaction="ignore")
writer.writeheader()
writer.writerows(rows)
meta = {
"filename": safe_name,
"content_type": "text/csv; charset=utf-8",
"size_bytes": target.stat().st_size,
"created_ts": now,
"expires_ts": now + EXPORT_TTL_SECONDS,
}
(folder / "meta.json").write_text(json.dumps(meta, ensure_ascii=False, indent=2), encoding="utf-8")
path = f"/api/fisheries/download/{token}"
return {
"token": token,
"filename": safe_name,
"download_path": path,
"download_url": f"{PUBLIC_BASE_URL}{path}" if PUBLIC_BASE_URL else path,
}
def _process_reader(
reader: csv.DictReader,
*,
source_label: str,
year: int | None,
lon_min: float | None,
lon_max: float | None,
lat_min: float | None,
lat_max: float | None,
metric_columns: str | None,
max_rows: int,
state: dict[str, Any],
) -> None:
columns = [str(x or "").strip() for x in (reader.fieldnames or [])]
if not columns:
return
if not state["columns"]:
state["columns"] = columns
state["metrics"] = _metric_columns(columns, metric_columns)
elif columns != state["columns"]:
state["schema_variants"].append({"source": source_label, "columns": columns})
seen = state["seen"]
for raw in reader:
if state["scanned"] >= max_rows:
state["truncated"] = True
return
state["scanned"] += 1
row = {str(k or "").strip(): "" if v is None else str(v).strip() for k, v in raw.items()}
matched, row_year, row_month = _matches_filters(
row,
columns,
year=year,
lon_min=lon_min,
lon_max=lon_max,
lat_min=lat_min,
lat_max=lat_max,
)
if not matched:
continue
# Quality counts are reported for the requested result set. This
# keeps large year/bbox ZIP queries bounded instead of retaining a
# duplicate key for every global row that was scanned and rejected.
key = tuple(row.get(col, "") for col in columns)
if key in seen:
state["duplicates"] += 1
else:
seen.add(key)
for col in columns:
if not row.get(col, "").strip():
state["missing"][col] += 1
state["matched"] += 1
state["rows"].append(row)
if row_year is not None:
state["years"].add(row_year)
if row_month is not None:
state["months"][row_month] += 1
if row_year is not None:
state["annual_counts"][row_year] += 1
for col in state["metrics"]:
value = _number(row.get(col))
if value is not None:
state["annual_values"][(row_year, col)] += value
def analyze_and_export(
local_path: Path,
*,
dataset_path: str,
revision: str,
repository: str | None = None,
year: int | None = None,
lon_min: float | None = None,
lon_max: float | None = None,
lat_min: float | None = None,
lat_max: float | None = None,
metric_columns: str | None = None,
max_rows: int = 2_000_000,
) -> dict[str, Any]:
if lon_min is not None and not -180 <= float(lon_min) <= 180:
raise ValueError("lon_min 必须位于 -180 至 180。")
if lon_max is not None and not -180 <= float(lon_max) <= 180:
raise ValueError("lon_max 必须位于 -180 至 180。")
if lat_min is not None and not -90 <= float(lat_min) <= 90:
raise ValueError("lat_min 必须位于 -90 至 90。")
if lat_max is not None and not -90 <= float(lat_max) <= 90:
raise ValueError("lat_max 必须位于 -90 至 90。")
if lon_min is not None and lon_max is not None and float(lon_min) > float(lon_max):
raise ValueError("lon_min 不能大于 lon_max。")
if lat_min is not None and lat_max is not None and float(lat_min) > float(lat_max):
raise ValueError("lat_min 不能大于 lat_max。")
has_filter = year is not None or any(
value is not None for value in (lon_min, lon_max, lat_min, lat_max)
)
if (
local_path.suffix.lower() == ".zip"
and local_path.stat().st_size > 200 * 1024 * 1024
and not has_filter
):
raise ValueError(
"大型 ZIP 查询必须提供 year 或经纬度范围,避免无边界解压扫描;"
"请补充筛选条件后重试。"
)
max_rows = max(1, min(int(max_rows or 2_000_000), 5_000_000))
state: dict[str, Any] = {
"columns": [], "metrics": [], "schema_variants": [], "scanned": 0,
"matched": 0, "duplicates": 0, "missing": defaultdict(int),
"years": set(), "months": defaultdict(int), "annual_counts": defaultdict(int),
"annual_values": defaultdict(float), "rows": [], "seen": set(), "truncated": False,
}
suffix = local_path.suffix.lower()
members: list[str] = []
if suffix in {".csv", ".tsv"}:
text, encoding = _decode_text(local_path)
delimiter = _dialect(text, suffix)
_process_reader(
csv.DictReader(io.StringIO(text), delimiter=delimiter),
source_label=dataset_path,
year=year, lon_min=lon_min, lon_max=lon_max,
lat_min=lat_min, lat_max=lat_max,
metric_columns=metric_columns, max_rows=max_rows, state=state,
)
else:
encoding = "utf-8-replace"
delimiter = "auto-by-member-extension"
with zipfile.ZipFile(local_path) as archive:
candidates = [
info for info in archive.infolist()
if not info.is_dir() and Path(info.filename).suffix.lower() in {".csv", ".tsv"}
]
if year is not None:
preferred = [info for info in candidates if str(year) in info.filename]
if preferred:
candidates = preferred
for info in candidates:
members.append(info.filename)
with archive.open(info) as binary:
text_stream = io.TextIOWrapper(binary, encoding="utf-8-sig", errors="replace", newline="")
sample_lines = list(itertools.islice(text_stream, 50))
member_delimiter = _dialect(
"".join(sample_lines),
Path(info.filename).suffix.lower(),
)
_process_reader(
csv.DictReader(
itertools.chain(sample_lines, text_stream),
delimiter=member_delimiter,
),
source_label=info.filename,
year=year, lon_min=lon_min, lon_max=lon_max,
lat_min=lat_min, lat_max=lat_max,
metric_columns=metric_columns, max_rows=max_rows, state=state,
)
if state["truncated"]:
break
if not state["columns"]:
raise ValueError("未在文件中发现可读取的 CSV/TSV 表格。")
annual_summary = []
for yr in sorted(state["annual_counts"]):
item: dict[str, Any] = {"year": yr, "record_count": state["annual_counts"][yr]}
for col in state["metrics"]:
item[f"sum_{col}"] = state["annual_values"].get((yr, col), 0.0)
annual_summary.append(item)
# Return three distinct, truthfully labelled artifacts when the source has
# a usable year field. Keep the raw-detail export as the backwards-
# compatible primary download_url/token/filename.
stamp = int(time.time())
source_stem = Path(dataset_path).stem or "fisheries_query"
raw_export = _create_export(
f"{source_stem}_{stamp}_filtered_raw.csv",
state["rows"],
state["columns"],
)
deduplicated_rows = []
exported_keys = set()
for row in state["rows"]:
key = tuple(row.get(col, "") for col in state["columns"])
if key in exported_keys:
continue
exported_keys.add(key)
deduplicated_rows.append(row)
deduplicated_export = _create_export(
f"{source_stem}_{stamp}_deduplicated.csv",
deduplicated_rows,
state["columns"],
)
exports = [
{"kind": "filtered_raw", **raw_export},
{
"kind": "deduplicated",
"record_count": len(deduplicated_rows),
**deduplicated_export,
},
]
if annual_summary:
annual_columns = list(annual_summary[0])
annual_export = _create_export(
f"{source_stem}_{stamp}_annual_summary.csv",
annual_summary,
annual_columns,
)
exports.append({"kind": "annual_summary", **annual_export})
return {
"status": "ok",
"repository": normalize_repository(repository),
"revision": revision,
"dataset_path": dataset_path,
"encoding": encoding,
"delimiter": delimiter,
"columns": state["columns"],
"metric_columns": state["metrics"],
"scanned_row_count": state["scanned"],
"matched_row_count": state["matched"],
"scan_truncated": state["truncated"],
"time_range": {
"min_year": min(state["years"]) if state["years"] else None,
"max_year": max(state["years"]) if state["years"] else None,
},
"monthly_record_counts": [
{"month": month, "record_count": state["months"][month]}
for month in sorted(state["months"])
],
"missing_values_by_column": {
col: int(state["missing"].get(col, 0)) for col in state["columns"]
},
"exact_duplicate_count": state["duplicates"],
"deduplicated_record_count": len(deduplicated_rows),
"annual_summary": annual_summary[:100],
"exports": exports,
"download_urls": [item["download_url"] for item in exports],
"zip_members_processed": members[:200],
"schema_variants": state["schema_variants"][:20],
**raw_export,
}