"""High-level Python interface for reproducible dataset generation."""
from __future__ import annotations
import asyncio
import json
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import httpx
from openworld_radio_twin.batch import (
BatchDatasetRequest,
BatchPlan,
ProgressCallback,
build_batch_plan,
generate_batch_dataset,
)
from openworld_radio_twin.config import Settings, get_settings
from openworld_radio_twin.models import EngineName
from openworld_radio_twin.providers.map_data import MapDataProvider
from openworld_radio_twin.simulation.sionna_engine import availability as sionna_availability
DatasetRequestSource = BatchDatasetRequest | Mapping[str, Any] | str | Path
[docs]
@dataclass(frozen=True)
class DatasetGenerationResult:
"""Completed dataset location together with its deterministic expanded plan."""
output_directory: Path
plan: BatchPlan
@property
def manifest_path(self) -> Path:
return self.output_directory / "dataset_manifest.json"
[docs]
def read_manifest(self) -> dict[str, Any]:
"""Read the completed dataset manifest."""
value = json.loads(self.manifest_path.read_text(encoding="utf-8"))
if not isinstance(value, dict):
raise ValueError(f"Dataset manifest is not a JSON object: {self.manifest_path}")
return value
[docs]
def load_dataset_request(source: DatasetRequestSource) -> BatchDatasetRequest:
"""Validate a model, mapping, or JSON file as a batch dataset request."""
if isinstance(source, BatchDatasetRequest):
return source
if isinstance(source, (str, Path)):
path = Path(source).expanduser().resolve()
return BatchDatasetRequest.model_validate_json(path.read_text(encoding="utf-8"))
return BatchDatasetRequest.model_validate(source)
[docs]
def plan_dataset(source: DatasetRequestSource) -> BatchPlan:
"""Expand a dataset request without provider access or solver execution."""
return build_batch_plan(load_dataset_request(source))
def _ignore_progress(completed: int, total: int, message: str) -> None:
del completed, total, message
[docs]
async def generate_dataset(
source: DatasetRequestSource,
*,
output_root: str | Path | None = None,
settings: Settings | None = None,
progress: ProgressCallback | None = None,
resume: bool = False,
) -> DatasetGenerationResult:
"""Generate a dataset without starting the web service.
The function owns and closes its HTTP client and worker thread. Use this asynchronous
entry point from notebooks, services, or existing event loops. For ordinary scripts, use
:func:`generate_dataset_sync`. Set ``resume=True`` to skip complete cases and regenerate
missing cases in a compatible existing dataset.
"""
runtime_settings = settings or get_settings()
request = load_dataset_request(source)
plan = build_batch_plan(request)
root = Path(output_root or runtime_settings.dataset_root).expanduser().resolve()
target = root / plan.dataset_slug
if target.exists() and not resume:
raise FileExistsError(
f"Dataset '{plan.dataset_slug}' already exists at {target}; "
"choose a new dataset name/output root, or pass resume=True"
)
uses_sionna = any(
case.request.engine == EngineName.sionna.value
for scene in plan.scenes
for case in scene.cases
)
if uses_sionna:
available, message = sionna_availability()
if not available:
raise RuntimeError(message)
root.mkdir(parents=True, exist_ok=True)
async with httpx.AsyncClient(
follow_redirects=True,
timeout=runtime_settings.http_request_timeout_seconds,
) as client:
provider = MapDataProvider(runtime_settings, client)
with ThreadPoolExecutor(max_workers=1, thread_name_prefix="owrt-dataset") as worker:
output_directory = await generate_batch_dataset(
request,
plan,
root,
provider,
runtime_settings.max_buildings,
asyncio.Lock(),
worker,
progress or _ignore_progress,
resume=resume,
)
return DatasetGenerationResult(output_directory=output_directory, plan=plan)
[docs]
def generate_dataset_sync(
source: DatasetRequestSource,
*,
output_root: str | Path | None = None,
settings: Settings | None = None,
progress: ProgressCallback | None = None,
resume: bool = False,
) -> DatasetGenerationResult:
"""Generate, resume, or deterministically extend a dataset from a normal script."""
try:
asyncio.get_running_loop()
except RuntimeError:
return asyncio.run(
generate_dataset(
source,
output_root=output_root,
settings=settings,
progress=progress,
resume=resume,
)
)
raise RuntimeError(
"generate_dataset_sync() cannot run inside an active event loop; "
"use 'await generate_dataset(...)' instead"
)