openitalaw / src /hf_streaming_loader.py
OpenItaLaw Builder
Add .streamlit config, /data cache, dataset freshness dashboard panel
6873c03
Raw
History Blame Contribute Delete
13 kB
"""
HF Dataset Streaming Loader for OpenItaLaw App
Efficiently streams FAISS indexes and metadata from HF dataset without
requiring full 3.3 GB downloads. Supports both direct HTTP streaming
and caching for repeated queries.
Usage:
loader = HFStreamingLoader(repo_id="diatribe00/ItalianLawEngine")
index = loader.load_faiss_index(tier="1") # Streams from HF
records = loader.load_metadata_records(limit=100) # Lazy-loads
"""
import json
import logging
import os
from pathlib import Path
from typing import Optional, List, Dict, Iterator, Any
import tempfile
from functools import lru_cache
import requests
from huggingface_hub import hf_hub_url, hf_hub_download
logger = logging.getLogger(__name__)
class HFStreamingLoader:
"""Stream FAISS indexes and metadata from HF dataset efficiently."""
FAISS_DIR = "faiss_index"
VIGENTE_DIR = "vigente_pipeline"
DEFAULT_TIER_INDEXES = [
"faiss_index/index.faiss", # Main index (all vectors)
]
def __init__(
self,
repo_id: str = "diatribe00/ItalianLawEngine",
token: Optional[str] = None,
cache_dir: Optional[Path] = None,
):
"""
Initialize HF streaming loader.
Args:
repo_id: HF dataset repository ID
token: HF API token (for private datasets)
cache_dir: Directory for caching downloaded files (default: ~/.cache/italaw)
"""
self.repo_id = repo_id
self.token = token
# Prefer /data (HF Spaces persistent disk) so files survive container
# restarts and are not re-downloaded every session.
if cache_dir:
self.cache_dir = cache_dir
elif Path("/data").exists() and os.access("/data", os.W_OK):
self.cache_dir = Path("/data") / "hf_cache" / "italaw"
else:
self.cache_dir = Path.home() / ".cache" / "italaw"
self.cache_dir.mkdir(parents=True, exist_ok=True)
logger.info(f"Initialized HF streaming loader: {repo_id}")
logger.info(f"Cache directory: {self.cache_dir}")
@lru_cache(maxsize=1)
def get_file_list(self) -> List[str]:
"""Get list of available files from dataset (cached)."""
try:
api_url = f"https://huggingface.co/api/datasets/{self.repo_id}/tree/main?recursive=true&expand=false"
headers = {}
if self.token:
headers["Authorization"] = f"Bearer {self.token}"
resp = requests.get(api_url, headers=headers, timeout=30)
resp.raise_for_status()
data = resp.json()
files = data if isinstance(data, list) else data.get("tree", [])
file_paths = [f.get("path") for f in files if isinstance(f, dict) and "path" in f]
logger.info(f"Found {len(file_paths)} files in dataset")
return file_paths
except Exception as e:
logger.error(f"Failed to get file list: {e}")
return []
def load_faiss_index(self, tier: str = "all") -> Optional[Any]:
"""
Load FAISS index from HF, optionally tier-specific.
Tiers: "1" (current), "2" (active), "3" (recent), "4" (archive),
"all" (main index with all vectors)
"""
import faiss
# For now, main index works with metadata filtering
# Future: tier-specific indexes (t1_index.faiss, etc.)
index_file = f"{self.FAISS_DIR}/index.faiss"
try:
logger.info(f"Loading FAISS index: {index_file}")
# Download with caching (skips if already cached)
local_path = hf_hub_download(
repo_id=self.repo_id,
filename=index_file,
repo_type="dataset",
token=self.token,
cache_dir=str(self.cache_dir),
local_dir_use_symlinks=False,
)
index = faiss.read_index(local_path)
logger.info(f"✓ FAISS index loaded: {index.ntotal} vectors, {index.d} dims")
return index
except Exception as e:
logger.error(f"Failed to load FAISS index: {e}")
return None
def load_metadata_bulk(self, limit: Optional[int] = None) -> List[Dict[str, Any]]:
"""
Load metadata records from JSONL file.
For initial load/search. If limit=None, loads all (may be large).
WARNING: For 839K records, loads all in memory (~1.5 GB).
For Streamlit on limited resources, consider stream_metadata() instead.
"""
try:
meta_file = f"{self.FAISS_DIR}/doc_metadata.jsonl"
logger.info(f"Loading metadata: {meta_file}")
local_path = hf_hub_download(
repo_id=self.repo_id,
filename=meta_file,
repo_type="dataset",
token=self.token,
cache_dir=str(self.cache_dir),
local_dir_use_symlinks=False,
)
records = []
with open(local_path, encoding="utf-8") as f:
for i, line in enumerate(f):
if limit and i >= limit:
logger.info(f"Stopped at limit: {limit} records")
break
try:
record = json.loads(line.strip())
records.append(record)
except json.JSONDecodeError:
logger.warning(f"Skipped invalid JSON at line {i+1}")
logger.info(f"✓ Loaded {len(records)} metadata records")
return records
except Exception as e:
logger.error(f"Failed to load metadata: {e}")
return []
def get_metadata_index(self) -> Dict[str, Dict[str, Any]]:
"""
Build a queryable index of metadata by ID (for large datasets).
Loads all into memory but allows efficient lookups.
"""
try:
meta_file = f"{self.FAISS_DIR}/doc_metadata.jsonl"
local_path = hf_hub_download(
repo_id=self.repo_id,
filename=meta_file,
repo_type="dataset",
token=self.token,
cache_dir=str(self.cache_dir),
local_dir_use_symlinks=False,
)
index = {}
count = 0
with open(local_path, encoding="utf-8") as f:
for line in f:
try:
record = json.loads(line.strip())
if 'id' in record:
index[record['id']] = record
count += 1
except json.JSONDecodeError:
pass
logger.info(f"✓ Built metadata index with {count} records")
return index
except Exception as e:
logger.error(f"Failed to build metadata index: {e}")
return {}
def stream_metadata(self, start: int = 0, batch_size: int = 100) -> Iterator[Dict[str, Any]]:
"""
Stream metadata records in batches (memory-efficient).
Useful for large result sets or pagination.
"""
try:
meta_file = f"{self.FAISS_DIR}/doc_metadata.jsonl"
local_path = hf_hub_download(
repo_id=self.repo_id,
filename=meta_file,
repo_type="dataset",
token=self.token,
cache_dir=str(self.cache_dir),
local_dir_use_symlinks=False,
)
with open(local_path, encoding="utf-8") as f:
for i, line in enumerate(f):
if i < start:
continue
try:
record = json.loads(line.strip())
yield record
except json.JSONDecodeError:
logger.warning(f"Skipped invalid JSON at line {i+1}")
except Exception as e:
logger.error(f"Failed to stream metadata: {e}")
def get_metadata_record(self, doc_id: str, all_records: Optional[List[Dict]] = None) -> Optional[Dict]:
"""
Get a single metadata record by ID.
For efficiency, provide all_records if already loaded.
"""
if all_records:
for rec in all_records:
if rec.get("id") == doc_id:
return rec
return None
# Fallback: load all (expensive)
all_records = self.load_metadata_bulk()
for rec in all_records:
if rec.get("id") == doc_id:
return rec
return None
def filter_by_vigente(self, records: List[Dict]) -> List[Dict]:
"""Filter metadata records to only vigente (in-force) laws."""
return [
r for r in records
if r.get("validity_status") == "in_corso" # Italian: "in corso" = in force
]
def filter_by_era(self, records: List[Dict], era: str) -> List[Dict]:
"""Filter by legal era (e.g., 'Fascismo', 'Repubblica', 'Contemporaneo')."""
return [r for r in records if r.get("legal_era") == era]
def get_cache_size_mb(self) -> float:
"""Get total size of cached files in MB."""
total = 0
for file in self.cache_dir.rglob("*"):
if file.is_file():
total += file.stat().st_size
return total / (1024 * 1024)
def get_dataset_info(self) -> Dict[str, Any]:
"""
Fetch dataset-level info from HF API — no file download needed.
Returns lastModified timestamp, commit SHA, and sibling file list
with sizes. Used to drive the dashboard freshness panel.
"""
try:
url = f"https://huggingface.co/api/datasets/{self.repo_id}"
headers = {"Authorization": f"Bearer {self.token}"} if self.token else {}
resp = requests.get(url, headers=headers, timeout=15)
resp.raise_for_status()
data = resp.json()
siblings = data.get("siblings", [])
file_info = {s["rfilename"]: s for s in siblings if "rfilename" in s}
return {
"last_modified": data.get("lastModified") or data.get("updatedAt"),
"sha": data.get("sha"),
"private": data.get("private", False),
"files": file_info,
}
except Exception as e:
logger.warning(f"Could not fetch dataset info: {e}")
return {}
def sample_metadata_stats(self, sample_size: int = 5000) -> Dict[str, Any]:
"""
Stream the first *sample_size* metadata records and return aggregate
stats without downloading all 1.5 GB. Used by the dashboard.
"""
source_counts: Dict[str, int] = {}
doc_type_counts: Dict[str, int] = {}
years: list = []
count = 0
try:
for rec in self.stream_metadata(batch_size=sample_size):
if count >= sample_size:
break
st = rec.get("build_source") or rec.get("source_type") or "unknown"
source_counts[st] = source_counts.get(st, 0) + 1
dt = rec.get("doc_type") or rec.get("type") or "unknown"
doc_type_counts[dt] = doc_type_counts.get(dt, 0) + 1
date_str = rec.get("date") or rec.get("data_pubblicazione") or ""
if date_str and len(date_str) >= 4:
try:
years.append(int(date_str[:4]))
except ValueError:
pass
count += 1
except Exception as e:
logger.warning(f"sample_metadata_stats error: {e}")
return {
"sampled": count,
"source_counts": source_counts,
"doc_type_counts": doc_type_counts,
"year_min": min(years) if years else None,
"year_max": max(years) if years else None,
}
if __name__ == "__main__":
# Example usage
logging.basicConfig(level=logging.INFO)
loader = HFStreamingLoader()
# Load index
index = loader.load_faiss_index()
if index:
print(f"Index: {index.ntotal} vectors × {index.d} dims")
# Load first 10 records
records = loader.load_metadata_bulk(limit=10)
print(f"Loaded {len(records)} records")
print(f"Sample: {records[0] if records else 'None'}")
# Filter vigente
vigente = loader.filter_by_vigente(records)
print(f"Vigente: {len(vigente)} out of {len(records)}")
# Cache size
cache_mb = loader.get_cache_size_mb()
print(f"Cache size: {cache_mb:.1f} MB")