From 882f8abcaa347a9c471edc22a6f68b8b6df5372e Mon Sep 17 00:00:00 2001 From: Yi Date: Fri, 10 Apr 2026 23:35:57 +0200 Subject: [PATCH] Refactor remote transport module and add local sync workflow --- README.md | 43 ++-- remote_api/__init__.py | 1 + remote_api/client.py | 442 ++++++++++++++++++++++++++++++++++++++++ remote_api/server.py | 401 +++++++++++++++++++++++++++++++++++++ sim_db.py | 265 ++++++++++++++++++++++++ sim_db_client.py | 445 +---------------------------------------- sim_db_server.py | 398 +----------------------------------- test_sim_db.py | 50 +++++ test_sim_db_rest.py | 10 +- 9 files changed, 1200 insertions(+), 855 deletions(-) create mode 100644 remote_api/__init__.py create mode 100644 remote_api/client.py create mode 100644 remote_api/server.py diff --git a/README.md b/README.md index 972b1a6..7a1bca2 100644 --- a/README.md +++ b/README.md @@ -1,17 +1,15 @@ # mini_sim_db -Tiny simulation run tracker. +Tiny local-first simulation run tracker. -Now SQLite-backed (`sqlite3` stdlib) with the same CLI + REST workflow. +Core storage is SQLite (`sim_db.py`). Remote host/client transport now lives in `remote_api/` as an optional module. -- If you pass a `.csv` path (old usage), it is treated as compatibility input and mapped to a sibling `.sqlite3` DB file. -- If that CSV exists and the SQLite file does not, rows are auto-imported on first open. +- `.csv` paths are still accepted for compatibility and map to a sibling `.sqlite3` DB. +- If legacy CSV exists and SQLite does not, rows are auto-imported on first open. ## Quick start (local CLI) ```bash -# default argument stays compatible: ~/sim_db.csv -# actual DB file is ~/sim_db.sqlite3 python sim_db.py init python sim_db.py add \ @@ -21,29 +19,36 @@ python sim_db.py add \ --status start python sim_db.py done --case case_001 - -# easy inspection table python sim_db.py list --table -python sim_db.py list --table --status done --limit 20 +``` + +## Local-first sync workflow (JSON artifact) + +```bash +# inspect unsynced local updates +python sim_db.py sync-status --table -# optional explicit legacy import -python sim_db.py import-csv --csv ./legacy.csv +# export pending rows to a portable artifact +python sim_db.py sync-export --out ./sync-out.json + +# import artifact from another machine +python sim_db.py sync-import --in ./sync-out.json ``` -## Quick start (REST) +Sync format: JSON `mini_sim_db_sync_v1` with full row snapshots. Merge policy: per `job_id`, newer `updated_at` wins; if local is newer, import reports a conflict for manual review. + +## Optional REST transport ```bash export SIM_DB_API_TOKEN='replace-me' -python sim_db_server.py --host 127.0.0.1 --port 8765 --db ~/sim_db.csv - -python sim_db_client.py --url http://127.0.0.1:8765 --token "$SIM_DB_API_TOKEN" init -python sim_db_client.py --url http://127.0.0.1:8765 --token "$SIM_DB_API_TOKEN" create \ - --case c100 --inp c100.inp --bin solver --status start -python sim_db_client.py --url http://127.0.0.1:8765 --token "$SIM_DB_API_TOKEN" summary --status start --limit 20 +python remote_api/server.py --host 127.0.0.1 --port 8765 --db ~/sim_db.csv +python remote_api/client.py --url http://127.0.0.1:8765 --token "$SIM_DB_API_TOKEN" health ``` +Compatibility entrypoints are kept (`sim_db_server.py`, `sim_db_client.py`). + ## Tests ```bash -python -m unittest -v +python3 -m unittest -v ``` diff --git a/remote_api/__init__.py b/remote_api/__init__.py new file mode 100644 index 0000000..d30eb13 --- /dev/null +++ b/remote_api/__init__.py @@ -0,0 +1 @@ +"""Optional remote transport module for mini_sim_db.""" diff --git a/remote_api/client.py b/remote_api/client.py new file mode 100644 index 0000000..85f6b60 --- /dev/null +++ b/remote_api/client.py @@ -0,0 +1,442 @@ +"""Tiny REST client for mini_sim_db server with local durable dual-write.""" + +from __future__ import annotations + +import argparse +import json +import os +import socket +import sys +from typing import Any +from urllib import error, parse, request + + +class RemoteRequestError(RuntimeError): + """Base error for remote request failures.""" + + +class RemoteTransportError(RemoteRequestError): + """Network/transport failure reaching remote server.""" + + +class RemoteResponseError(RemoteRequestError): + """Remote server responded with an HTTP/application error.""" + +from sim_db import _read_sim_db, add_sim_item, del_cases, init_sim_db, mark_done, resolve_case_ref, upd_cases + + +def _case_ref(*, case: str | None, job_id: str | None) -> str: + if bool(case) == bool(job_id): + raise ValueError("use exactly one of case or job_id") + return str(case or job_id) + + +class SimDbClient: + def __init__( + self, + base_url: str, + token: str, + timeout: float = 10.0, + local_db_path: str | None = None, + enable_local_write: bool = True, + ) -> None: + self.base_url = base_url.rstrip("/") + self.token = token + self.timeout = timeout + self.local_db_path = os.path.expanduser(local_db_path) if local_db_path else None + self.enable_local_write = enable_local_write + + def health(self) -> dict[str, Any]: + return self._request("GET", "/health") + + def init(self, db_path: str | None = None) -> dict[str, Any]: + payload: dict[str, Any] = {} + if db_path: + payload["db_path"] = db_path + return self._request("POST", "/init", payload) + + def create( + self, + *, + case: str, + bin_name: str, + status: str, + inp: str | None = None, + input_files: list[str] | None = None, + note: str | None = None, + work_dir: str | None = None, + extra_params: str | None = None, + db_path: str | None = None, + run_host: str | None = None, + ) -> dict[str, Any]: + payload: dict[str, Any] = { + "case": case, + "bin_name": bin_name, + "status": status, + "run_host": run_host or socket.gethostname(), + } + if inp is not None: + payload["inp"] = inp + if input_files is not None: + payload["input_files"] = input_files + if note is not None: + payload["note"] = note + if work_dir is not None: + payload["work_dir"] = work_dir + if extra_params is not None: + payload["extra_params"] = extra_params + if db_path is not None: + payload["db_path"] = db_path + return self._dual_write("create", payload) + + def add(self, **kwargs: Any) -> dict[str, Any]: + return self.create(**kwargs) + + def read(self, *, case: str | None = None, job_id: str | None = None, db_path: str | None = None) -> dict[str, Any]: + case_ref = _case_ref(case=case, job_id=job_id) + path = f"/cases/{parse.quote(case_ref, safe='')}" + if db_path: + path += "?" + parse.urlencode({"db_path": db_path}) + return self._request("GET", path) + + def done( + self, + *, + case: str | None = None, + job_id: str | None = None, + db_path: str | None = None, + run_host: str | None = None, + ) -> dict[str, Any]: + return self.update( + case=case, + job_id=job_id, + fields={"status": "done"}, + db_path=db_path, + run_host=run_host, + ) + + def update( + self, + *, + case: str | None = None, + job_id: str | None = None, + fields: dict[str, Any], + db_path: str | None = None, + run_host: str | None = None, + ) -> dict[str, Any]: + payload: dict[str, Any] = { + "case": case, + "job_id": job_id, + "fields": fields, + "run_host": run_host or socket.gethostname(), + } + if db_path is not None: + payload["db_path"] = db_path + return self._dual_write("update", payload) + + def delete( + self, + *, + case: str | None = None, + job_id: str | None = None, + db_path: str | None = None, + run_host: str | None = None, + ) -> dict[str, Any]: + payload: dict[str, Any] = { + "case": case, + "job_id": job_id, + "run_host": run_host or socket.gethostname(), + } + if db_path is not None: + payload["db_path"] = db_path + return self._dual_write("delete", payload) + + def list(self, db_path: str | None = None) -> dict[str, Any]: + path = "/cases" + if db_path: + path += "?" + parse.urlencode({"db_path": db_path}) + return self._request("GET", path) + + def summary( + self, + *, + db_path: str | None = None, + status: str | None = None, + run_host: str | None = None, + limit: int | None = None, + sort_by: str = "updated_at", + order: str = "desc", + ) -> dict[str, Any]: + query: dict[str, Any] = {"sort_by": sort_by, "order": order} + if db_path: + query["db_path"] = db_path + if status: + query["status"] = status + if run_host: + query["run_host"] = run_host + if limit is not None: + query["limit"] = limit + return self._request("GET", "/cases/summary?" + parse.urlencode(query)) + + def _dual_write(self, op: str, payload: dict[str, Any]) -> dict[str, Any]: + local_ok = None + local_error = None + if self.enable_local_write and self.local_db_path: + try: + self._apply_local(op, payload) + local_ok = True + except Exception as exc: + local_ok = False + local_error = str(exc) + + try: + remote = self._request_for_op(op, payload) + out = {"ok": True, "remote_ok": True, "remote": remote} + if local_ok is not None: + out["local_ok"] = local_ok + if local_error: + out["local_error"] = local_error + return out + except RemoteTransportError as exc: + if local_ok: + return { + "ok": True, + "remote_ok": False, + "remote_error": str(exc), + "local_ok": True, + "fallback": "local-only", + } + raise + + def _apply_local(self, op: str, payload: dict[str, Any]) -> None: + assert self.local_db_path is not None + init_sim_db(self.local_db_path) + run_host = payload.get("run_host") + + if op == "create": + if payload.get("job_id") not in (None, ""): + raise ValueError("field 'job_id' is auto-generated and cannot be set on create") + + case = payload["case"] + add_sim_item( + case=case, + inp=payload.get("inp"), + input_files=payload.get("input_files"), + bin_name=payload["bin_name"], + status=payload["status"], + db_path=self.local_db_path, + note=payload.get("note"), + work_dir=payload.get("work_dir"), + extra_params=payload.get("extra_params"), + ) + if run_host: + upd_cases(self.local_db_path, {case: {"run_host": str(run_host)}}) + return + + if op == "update": + _, rows = _read_sim_db(self.local_db_path) + case = resolve_case_ref(rows, _case_ref(case=payload.get("case"), job_id=payload.get("job_id"))) + fields = dict(payload.get("fields") or {}) + if run_host: + fields["run_host"] = str(run_host) + if fields.get("status") == "done": + mark_done(case=case, db_path=self.local_db_path) + fields.pop("status", None) + fields.pop("state_changed_at", None) + if fields: + upd_cases(self.local_db_path, {case: fields}) + return + + if op == "delete": + _, rows = _read_sim_db(self.local_db_path) + case = resolve_case_ref(rows, _case_ref(case=payload.get("case"), job_id=payload.get("job_id"))) + del_cases(self.local_db_path, [case]) + return + + raise ValueError(f"unsupported op: {op}") + + def _request_for_op(self, op: str, payload: dict[str, Any]) -> dict[str, Any]: + if op == "create": + return self._request("POST", "/cases", payload) + if op == "update": + case_ref = parse.quote(_case_ref(case=payload.get("case"), job_id=payload.get("job_id")), safe='') + remote_payload = {k: v for k, v in payload.items() if k not in {"case", "job_id"}} + return self._request("PATCH", f"/cases/{case_ref}", remote_payload) + if op == "delete": + case_ref = parse.quote(_case_ref(case=payload.get("case"), job_id=payload.get("job_id")), safe='') + db_path = payload.get("db_path") + path = f"/cases/{case_ref}" + if db_path: + path += "?" + parse.urlencode({"db_path": db_path}) + return self._request("DELETE", path) + raise ValueError(f"unsupported op: {op}") + + def _request(self, method: str, path: str, payload: dict[str, Any] | None = None) -> dict[str, Any]: + data = None + headers = {"Authorization": f"Bearer {self.token}"} + if payload is not None: + data = json.dumps(payload).encode("utf-8") + headers["Content-Type"] = "application/json" + + req = request.Request(self.base_url + path, method=method, headers=headers, data=data) + try: + with request.urlopen(req, timeout=self.timeout) as resp: + body = resp.read().decode("utf-8") + return json.loads(body) if body else {} + except error.HTTPError as exc: + body = exc.read().decode("utf-8") + msg = body or str(exc) + raise RemoteResponseError(f"HTTP {exc.code}: {msg}") from exc + except (error.URLError, TimeoutError) as exc: + raise RemoteTransportError(f"request failed: {exc}") from exc + + +def _parse_fields(pairs: list[str] | None) -> dict[str, str]: + fields: dict[str, str] = {} + for pair in pairs or []: + if "=" not in pair: + raise ValueError(f"invalid --field '{pair}', expected key=value") + key, value = pair.split("=", 1) + key = key.strip() + if not key: + raise ValueError(f"invalid --field '{pair}', empty key") + fields[key] = value + return fields + + +def _add_create_like_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument("--case", required=True) + parser.add_argument("--bin", dest="bin_name", required=True) + parser.add_argument("--status", required=True) + parser.add_argument("--inp", default=None) + parser.add_argument("--input-file", action="append", default=None) + parser.add_argument("--note", default=None) + parser.add_argument("--work-dir", default=None) + parser.add_argument("--extra-params", default=None, help="Raw extra runtime params string (for example JSON)") + parser.add_argument("--db", default=None) + + +def _build_parser() -> argparse.ArgumentParser: + p = argparse.ArgumentParser(description="mini_sim_db REST client") + p.add_argument("--url", default="http://127.0.0.1:8765", help="Server base URL") + p.add_argument("--token", default=None, help="Bearer token (or SIM_DB_API_TOKEN)") + p.add_argument( + "--local-db", + default=os.path.expanduser("~/.sim_db_client_local.sqlite3"), + help="Local durable mirror DB for dual-write fallback (default: ~/.sim_db_client_local.sqlite3)", + ) + p.add_argument("--no-local-write", action="store_true", help="Disable local dual-write fallback") + sub = p.add_subparsers(dest="cmd", required=True) + + sub.add_parser("health") + + p_init = sub.add_parser("init") + p_init.add_argument("--db", default=None, help="Optional remote db_path override") + + p_create = sub.add_parser("create") + _add_create_like_args(p_create) + + p_add = sub.add_parser("add") + _add_create_like_args(p_add) + + p_read = sub.add_parser("read") + read_target = p_read.add_mutually_exclusive_group(required=True) + read_target.add_argument("--case") + read_target.add_argument("--job-id", dest="job_id") + p_read.add_argument("--db", default=None) + + p_update = sub.add_parser("update") + update_target = p_update.add_mutually_exclusive_group(required=True) + update_target.add_argument("--case") + update_target.add_argument("--job-id", dest="job_id") + p_update.add_argument("--field", action="append", default=None, help="key=value (repeatable)") + p_update.add_argument("--db", default=None) + + p_done = sub.add_parser("done") + done_target = p_done.add_mutually_exclusive_group(required=True) + done_target.add_argument("--case") + done_target.add_argument("--job-id", dest="job_id") + p_done.add_argument("--db", default=None) + + p_delete = sub.add_parser("delete") + delete_target = p_delete.add_mutually_exclusive_group(required=True) + delete_target.add_argument("--case") + delete_target.add_argument("--job-id", dest="job_id") + p_delete.add_argument("--db", default=None) + + p_list = sub.add_parser("list") + p_list.add_argument("--db", default=None) + + p_summary = sub.add_parser("summary") + p_summary.add_argument("--db", default=None) + p_summary.add_argument("--status", default=None) + p_summary.add_argument("--run-host", default=None) + p_summary.add_argument("--limit", type=int, default=None) + p_summary.add_argument("--sort-by", default="updated_at") + p_summary.add_argument("--order", choices=["asc", "desc"], default="desc") + + return p + + +def main(argv: list[str] | None = None) -> int: + args = _build_parser().parse_args(argv) + token = args.token or os.getenv("SIM_DB_API_TOKEN") + if not token: + print("Missing token. Set --token or SIM_DB_API_TOKEN", file=sys.stderr) + return 2 + + client = SimDbClient( + base_url=args.url, + token=token, + local_db_path=None if args.no_local_write else args.local_db, + enable_local_write=not args.no_local_write, + ) + + try: + if args.cmd == "health": + result = client.health() + elif args.cmd == "init": + result = client.init(db_path=args.db) + elif args.cmd in {"create", "add"}: + result = client.create( + case=args.case, + inp=args.inp, + input_files=args.input_file, + bin_name=args.bin_name, + status=args.status, + note=args.note, + work_dir=args.work_dir, + extra_params=args.extra_params, + db_path=args.db, + ) + elif args.cmd == "read": + result = client.read(case=args.case, job_id=args.job_id, db_path=args.db) + elif args.cmd == "update": + result = client.update(case=args.case, job_id=args.job_id, fields=_parse_fields(args.field), db_path=args.db) + elif args.cmd == "done": + result = client.done(case=args.case, job_id=args.job_id, db_path=args.db) + elif args.cmd == "delete": + result = client.delete(case=args.case, job_id=args.job_id, db_path=args.db) + elif args.cmd == "list": + result = client.list(db_path=args.db) + elif args.cmd == "summary": + result = client.summary( + db_path=args.db, + status=args.status, + run_host=args.run_host, + limit=args.limit, + sort_by=args.sort_by, + order=args.order, + ) + else: + return 1 + except (RuntimeError, ValueError) as exc: + print(str(exc), file=sys.stderr) + return 2 + + print(json.dumps(result, indent=2, sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/remote_api/server.py b/remote_api/server.py new file mode 100644 index 0000000..d660c32 --- /dev/null +++ b/remote_api/server.py @@ -0,0 +1,401 @@ +"""HTTP host for mini_sim_db. + +Stdlib-only JSON API around sim_db.py for centralized CRUD updates. +""" + +from __future__ import annotations + +import argparse +import json +import os +import threading +from datetime import datetime +from pathlib import Path +from http import HTTPStatus +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any +from urllib.parse import parse_qs, unquote, urlsplit + +from sim_db import ( + ALLOWED_STATUS, + DEFAULT_DB_PATH, + add_sim_item, + del_cases, + init_sim_db, + list_items, + list_view, + upd_cases, +) + + +def _now_iso() -> str: + return datetime.now().isoformat(timespec="milliseconds") + + +class SecurityPolicy: + """Validates API token and requested DB paths.""" + + def __init__( + self, + token: str, + default_db_path: str, + allowed_db_path: str | None = None, + allowed_base_dir: str | None = None, + ) -> None: + self.token = token + self.default_db_path = _norm(default_db_path) + self.allowed_db_path = _norm(allowed_db_path) if allowed_db_path else None + self.allowed_base_dir = _norm(allowed_base_dir) if allowed_base_dir else None + + def is_authorized(self, auth_header: str | None) -> bool: + if not auth_header or not auth_header.startswith("Bearer "): + return False + return auth_header[len("Bearer ") :] == self.token + + def resolve_db_path(self, requested_db_path: str | None) -> str: + if not requested_db_path: + return self.default_db_path + + wanted = _norm(requested_db_path) + + if self.allowed_db_path is None and self.allowed_base_dir is None: + if wanted == self.default_db_path: + return wanted + raise ValueError("db_path is not allowed") + + if self.allowed_db_path and wanted == self.allowed_db_path: + return wanted + + if self.allowed_base_dir and _is_within_base(wanted, self.allowed_base_dir): + return wanted + + raise ValueError("db_path is outside allowed scope") + + +class SimDbApiServer(ThreadingHTTPServer): + def __init__(self, server_address: tuple[str, int], policy: SecurityPolicy) -> None: + self.policy = policy + self.mutation_lock = threading.Lock() + super().__init__(server_address, SimDbRequestHandler) + + +class SimDbRequestHandler(BaseHTTPRequestHandler): + server: SimDbApiServer + + def do_GET(self) -> None: # noqa: N802 + if self.path == "/health": + self._json(HTTPStatus.OK, {"ok": True}) + return + + if self.path.startswith("/cases"): + self._require_auth() + if not self._authorized: + return + + query = self._query_dict() + try: + db_path = self.server.policy.resolve_db_path(query.get("db_path")) + if urlsplit(self.path).path == "/cases/summary": + status = query.get("status") + run_host = query.get("run_host") + limit = int(query["limit"]) if query.get("limit") else None + sort_by = query.get("sort_by", "updated_at") + desc = query.get("order", "desc").lower() != "asc" + with self.server.mutation_lock: + rows = list_view(db_path=db_path, status=status, run_host=run_host, sort_by=sort_by, desc=desc, limit=limit) + self._json(HTTPStatus.OK, {"db_path": db_path, "count": len(rows), "items": rows}) + return + + case_ref = self._case_from_path() + with self.server.mutation_lock: + data = list_items(db_path) + if case_ref: + case = self._resolve_case_ref_from_table(data, case_ref) + self._json(HTTPStatus.OK, {"db_path": db_path, "case": case, "item": data[case]}) + return + self._json(HTTPStatus.OK, {"db_path": db_path, "cases": data}) + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) + return + + self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) + + def do_POST(self) -> None: # noqa: N802 + if self.path == "/init": + self._require_auth() + if not self._authorized: + return + payload = self._read_json_body() + if payload is None: + return + try: + db_path = self.server.policy.resolve_db_path(payload.get("db_path")) + with self.server.mutation_lock: + init_sim_db(db_path) + self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path}) + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) + return + + if self.path in {"/add", "/done", "/update", "/delete"}: + self._legacy_mutating_routes() + return + + if self.path == "/cases": + self._require_auth() + if not self._authorized: + return + payload = self._read_json_body() + if payload is None: + return + self._create_case(payload) + return + + self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) + + def do_PATCH(self) -> None: # noqa: N802 + if not self.path.startswith("/cases/"): + self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) + return + self._require_auth() + if not self._authorized: + return + + payload = self._read_json_body() + if payload is None: + return + + case = self._case_from_path() + if not case: + self._json(HTTPStatus.BAD_REQUEST, {"error": "missing case in path"}) + return + + self._update_case(case, payload) + + def do_DELETE(self) -> None: # noqa: N802 + if not self.path.startswith("/cases/"): + self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) + return + self._require_auth() + if not self._authorized: + return + + case_ref = self._case_from_path() + if not case_ref: + self._json(HTTPStatus.BAD_REQUEST, {"error": "missing case in path"}) + return + + query = self._query_dict() + req_db = query.get("db_path") + try: + db_path = self.server.policy.resolve_db_path(req_db) + with self.server.mutation_lock: + table = list_items(db_path) + case = self._resolve_case_ref_from_table(table, case_ref) + del_cases(db_path, [case]) + self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": case}) + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) + + def _legacy_mutating_routes(self) -> None: + self._require_auth() + if not self._authorized: + return + payload = self._read_json_body() + if payload is None: + return + + if self.path == "/add": + self._create_case(payload) + return + if self.path == "/done": + case_ref = payload.get("case") or payload.get("job_id") + if not case_ref: + self._json(HTTPStatus.BAD_REQUEST, {"error": "missing field: case or job_id"}) + return + self._update_case(case_ref, {"fields": {"status": "done"}, **payload}) + return + if self.path == "/update": + case_ref = payload.get("case") or payload.get("job_id") + if not case_ref: + self._json(HTTPStatus.BAD_REQUEST, {"error": "missing field: case or job_id"}) + return + self._update_case(case_ref, payload) + return + if self.path == "/delete": + case_ref = payload.get("case") or payload.get("job_id") + if not case_ref: + self._json(HTTPStatus.BAD_REQUEST, {"error": "missing field: case or job_id"}) + return + try: + db_path = self.server.policy.resolve_db_path(payload.get("db_path")) + with self.server.mutation_lock: + table = list_items(db_path) + case = self._resolve_case_ref_from_table(table, case_ref) + del_cases(db_path, [case]) + self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": case}) + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) + + def _create_case(self, payload: dict[str, Any]) -> None: + try: + if payload.get("job_id") not in (None, ""): + raise ValueError("field 'job_id' is auto-generated and cannot be set on create") + + db_path = self.server.policy.resolve_db_path(payload.get("db_path")) + with self.server.mutation_lock: + add_sim_item( + case=payload["case"], + inp=payload.get("inp"), + input_files=payload.get("input_files"), + bin_name=payload["bin_name"], + status=payload["status"], + db_path=db_path, + note=payload.get("note"), + work_dir=payload.get("work_dir"), + extra_params=payload.get("extra_params"), + ) + run_host = payload.get("run_host") + if run_host: + upd_cases(db_path, {payload["case"]: {"run_host": str(run_host)}}) + self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": payload["case"]}) + except KeyError as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": f"missing field: {exc.args[0]}"}) + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) + + @staticmethod + def _resolve_case_ref_from_table(data: dict[str, dict[str, Any]], case_ref: str) -> str: + if case_ref in data: + return case_ref + + matches = [case for case, item in data.items() if item.get("job_id") == case_ref] + if not matches: + raise ValueError(f"case/job_id not found: {case_ref}") + if len(matches) > 1: + joined = ", ".join(sorted(matches)) + raise ValueError(f"job_id matches multiple cases ({joined}), use case explicitly") + return matches[0] + + def _update_case(self, case_ref: str, payload: dict[str, Any]) -> None: + try: + db_path = self.server.policy.resolve_db_path(payload.get("db_path")) + fields = payload.get("fields") + if fields is None: + fields = {k: v for k, v in payload.items() if k not in {"case", "job_id", "db_path", "run_host"}} + if not isinstance(fields, dict) or not fields: + raise ValueError("fields must be a non-empty object") + if "case" in fields: + raise ValueError("field 'case' is immutable") + + status = fields.get("status") + if status is not None: + if status not in ALLOWED_STATUS: + allowed = ", ".join(sorted(ALLOWED_STATUS)) + raise ValueError(f"Invalid status '{status}'. Allowed: {allowed}") + fields["state_changed_at"] = _now_iso() + + fields["updated_at"] = _now_iso() + + run_host = payload.get("run_host") + if run_host: + fields["run_host"] = str(run_host) + + with self.server.mutation_lock: + table = list_items(db_path) + case = self._resolve_case_ref_from_table(table, case_ref) + upd_cases(db_path, {case: fields}) + self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": case, "updated": sorted(fields.keys())}) + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) + + def log_message(self, fmt: str, *args: Any) -> None: + return + + def _read_json_body(self) -> dict[str, Any] | None: + try: + raw_len = int(self.headers.get("Content-Length", "0")) + raw = self.rfile.read(raw_len) if raw_len > 0 else b"{}" + payload = json.loads(raw.decode("utf-8")) if raw else {} + if not isinstance(payload, dict): + raise ValueError("JSON body must be an object") + return payload + except Exception as exc: + self._json(HTTPStatus.BAD_REQUEST, {"error": f"invalid JSON: {exc}"}) + return None + + def _query_dict(self) -> dict[str, str]: + query = parse_qs(urlsplit(self.path).query, keep_blank_values=True) + out: dict[str, str] = {} + for k, vals in query.items(): + if vals: + out[k] = vals[0] + return out + + def _case_from_path(self) -> str | None: + path = urlsplit(self.path).path + parts = [p for p in path.split("/") if p] + if len(parts) >= 2 and parts[0] == "cases": + return unquote(parts[1]) + return None + + def _require_auth(self) -> None: + self._authorized = self.server.policy.is_authorized(self.headers.get("Authorization")) + if not self._authorized: + self._json(HTTPStatus.UNAUTHORIZED, {"error": "unauthorized"}) + + def _json(self, status: HTTPStatus, obj: dict[str, Any]) -> None: + data = json.dumps(obj, ensure_ascii=False).encode("utf-8") + self.send_response(int(status)) + self.send_header("Content-Type", "application/json; charset=utf-8") + self.send_header("Content-Length", str(len(data))) + self.end_headers() + self.wfile.write(data) + + +def _norm(path: str) -> str: + return str(Path(path).expanduser().resolve()) + + +def _is_within_base(path: str, base_dir: str) -> bool: + target = Path(path).expanduser().resolve() + base = Path(base_dir).expanduser().resolve() + try: + target.relative_to(base) + return True + except ValueError: + return False + + +def _build_parser() -> argparse.ArgumentParser: + p = argparse.ArgumentParser(description="mini_sim_db HTTP server") + p.add_argument("--host", default="127.0.0.1", help="Bind host (default: 127.0.0.1)") + p.add_argument("--port", type=int, default=8765, help="Bind port (default: 8765)") + p.add_argument("--db", default=DEFAULT_DB_PATH, help="Default DB path (default: ~/sim_db.csv)") + p.add_argument("--allowed-db-path", default=None, help="Optional exact writable DB path") + p.add_argument("--allowed-base-dir", default=None, help="Optional writable base directory") + p.add_argument("--token", default=None, help="Bearer token (or use SIM_DB_API_TOKEN env)") + return p + + +def main(argv: list[str] | None = None) -> int: + args = _build_parser().parse_args(argv) + token = args.token or os.getenv("SIM_DB_API_TOKEN") + if not token: + raise SystemExit("Missing token. Set --token or SIM_DB_API_TOKEN") + + policy = SecurityPolicy( + token=token, + default_db_path=args.db, + allowed_db_path=args.allowed_db_path, + allowed_base_dir=args.allowed_base_dir, + ) + server = SimDbApiServer((args.host, args.port), policy) + print(f"Serving mini_sim_db API at http://{args.host}:{args.port}") + print(f"Default DB path: {policy.default_db_path}") + server.serve_forever() + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/sim_db.py b/sim_db.py index c4fb855..f5c54c4 100644 --- a/sim_db.py +++ b/sim_db.py @@ -16,6 +16,7 @@ import json import os import re +import socket import sqlite3 import sys from datetime import datetime @@ -191,6 +192,16 @@ def _ensure_schema(conn: sqlite3.Connection) -> None: ) """ ) + conn.execute( + """ + CREATE TABLE IF NOT EXISTS sim_sync_state ( + job_id TEXT PRIMARY KEY, + last_synced_updated_at TEXT NOT NULL DEFAULT '', + last_exported_at TEXT NOT NULL DEFAULT '', + last_imported_at TEXT NOT NULL DEFAULT '' + ) + """ + ) conn.commit() @@ -265,6 +276,92 @@ def _table_from_conn(conn: sqlite3.Connection) -> dict[str, dict[str, str]]: return out +def _base_and_extra_fields(detail: Mapping[str, Any]) -> tuple[dict[str, str], dict[str, str]]: + base = {k: str(v) for k, v in detail.items() if k in set(CLI_FIELDS)} + extras = {k: str(v) for k, v in detail.items() if k not in set(CLI_FIELDS) and k != 'case'} + return base, extras + + +def _insert_full_case(conn: sqlite3.Connection, case: str, detail: Mapping[str, Any]) -> None: + base, extras = _base_and_extra_fields(detail) + conn.execute( + """ + INSERT INTO sim_cases("case", work_dir, bin, inp, input_files, job_id, extra_params, status, + note, notes, state_changed_at, created_at, updated_at, run_host) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + case, + base.get('work_dir', ''), + base.get('bin', ''), + base.get('inp', ''), + base.get('input_files', ''), + base.get('job_id', ''), + base.get('extra_params', ''), + base.get('status', ''), + base.get('note', base.get('notes', '')), + base.get('notes', base.get('note', '')), + base.get('state_changed_at', ''), + base.get('created_at', ''), + base.get('updated_at', ''), + base.get('run_host', ''), + ), + ) + for key, value in extras.items(): + conn.execute( + 'INSERT OR REPLACE INTO sim_case_extra("case", field, value) VALUES (?, ?, ?)', + (case, key, value), + ) + + +def _replace_full_case(conn: sqlite3.Connection, case: str, detail: Mapping[str, Any]) -> None: + conn.execute('DELETE FROM sim_cases WHERE "case" = ?', (case,)) + _insert_full_case(conn, case, detail) + + +def _upsert_sync_state( + conn: sqlite3.Connection, + job_id: str, + *, + synced_updated_at: str | None = None, + exported_at: str | None = None, + imported_at: str | None = None, +) -> None: + row = conn.execute('SELECT * FROM sim_sync_state WHERE job_id = ?', (job_id,)).fetchone() + current = dict(row) if row else { + 'last_synced_updated_at': '', + 'last_exported_at': '', + 'last_imported_at': '', + } + if synced_updated_at is not None: + current['last_synced_updated_at'] = str(synced_updated_at) + if exported_at is not None: + current['last_exported_at'] = str(exported_at) + if imported_at is not None: + current['last_imported_at'] = str(imported_at) + + conn.execute( + """ + INSERT OR REPLACE INTO sim_sync_state(job_id, last_synced_updated_at, last_exported_at, last_imported_at) + VALUES (?, ?, ?, ?) + """, + (job_id, current['last_synced_updated_at'], current['last_exported_at'], current['last_imported_at']), + ) + + +def _pending_sync_rows(conn: sqlite3.Connection) -> list[dict[str, str]]: + rows = conn.execute( + """ + SELECT c.* + FROM sim_cases c + LEFT JOIN sim_sync_state s ON c.job_id = s.job_id + WHERE s.last_synced_updated_at IS NULL OR s.last_synced_updated_at < c.updated_at + ORDER BY c.updated_at, c."case" + """ + ).fetchall() + return [{'case': str(row['case']), **_row_to_detail(conn, row)} for row in rows] + + def _set_fields(conn: sqlite3.Connection, case: str, fields: Mapping[str, Any]) -> None: base_fields = {k: str(v) for k, v in fields.items() if k in set(CLI_FIELDS)} extras = {k: str(v) for k, v in fields.items() if k not in set(CLI_FIELDS) and k != 'case'} @@ -677,6 +774,19 @@ def _build_cli() -> argparse.ArgumentParser: p_import.add_argument('--csv', required=True, help='Path to legacy CSV file') p_import.add_argument('--db', default=DEFAULT_DB_PATH, help='Target DB path (CSV path auto-maps to SQLite)') + p_sync_status = sub.add_parser('sync-status', help='Show local sync status and pending records') + p_sync_status.add_argument('--db', default=DEFAULT_DB_PATH, help='Path to DB (CSV path auto-maps to SQLite)') + p_sync_status.add_argument('--table', action='store_true', help='Show pending rows in compact table view') + + p_sync_export = sub.add_parser('sync-export', help='Export pending updates into a JSON sync artifact') + p_sync_export.add_argument('--out', required=True, help='Output JSON file path') + p_sync_export.add_argument('--db', default=DEFAULT_DB_PATH, help='Path to DB (CSV path auto-maps to SQLite)') + p_sync_export.add_argument('--all', action='store_true', help='Export all rows, not only pending ones') + + p_sync_import = sub.add_parser('sync-import', help='Import updates from a JSON sync artifact') + p_sync_import.add_argument('--in', dest='in_path', required=True, help='Input JSON file path') + p_sync_import.add_argument('--db', default=DEFAULT_DB_PATH, help='Path to DB (CSV path auto-maps to SQLite)') + return parser @@ -693,6 +803,132 @@ def import_csv(csv_path: str, db_path: str = DEFAULT_DB_PATH) -> int: return max(0, added) +def sync_status(db_path: str = DEFAULT_DB_PATH) -> dict[str, Any]: + conn, _ = _connect_db(db_path) + try: + total = int(conn.execute('SELECT COUNT(*) FROM sim_cases').fetchone()[0]) + pending = _pending_sync_rows(conn) + synced = max(0, total - len(pending)) + last_export = conn.execute('SELECT MAX(last_exported_at) FROM sim_sync_state').fetchone()[0] or '' + last_import = conn.execute('SELECT MAX(last_imported_at) FROM sim_sync_state').fetchone()[0] or '' + return { + 'total_cases': total, + 'pending_cases': len(pending), + 'synced_cases': synced, + 'last_exported_at': str(last_export), + 'last_imported_at': str(last_import), + 'pending': pending, + } + finally: + conn.close() + + +def sync_export(db_path: str, out_path: str, include_all: bool = False) -> dict[str, Any]: + conn, sqlite_path = _connect_db(db_path) + exported_at = _now_iso() + source_host = socket.gethostname() + try: + if include_all: + rows = [{'case': c, **d} for c, d in _table_from_conn(conn).items()] + else: + rows = _pending_sync_rows(conn) + + artifact = { + 'format': 'mini_sim_db_sync_v1', + 'exported_at': exported_at, + 'source_host': source_host, + 'source_db': sqlite_path, + 'count': len(rows), + 'items': rows, + } + out_file = Path(out_path).expanduser() + out_file.parent.mkdir(parents=True, exist_ok=True) + out_file.write_text(json.dumps(artifact, indent=2, ensure_ascii=False, sort_keys=True) + '\n', encoding='utf-8') + + for row in rows: + job_id = str(row.get('job_id', '')) + if not job_id: + continue + _upsert_sync_state( + conn, + job_id, + synced_updated_at=str(row.get('updated_at', '')), + exported_at=exported_at, + ) + conn.commit() + return {'ok': True, 'path': str(out_file), 'exported': len(rows), 'exported_at': exported_at} + finally: + conn.close() + + +def sync_import(db_path: str, in_path: str) -> dict[str, Any]: + in_file = Path(in_path).expanduser() + payload = json.loads(in_file.read_text(encoding='utf-8')) + if payload.get('format') != 'mini_sim_db_sync_v1': + raise ValueError('unsupported sync artifact format') + if not isinstance(payload.get('items'), list): + raise ValueError('sync artifact must contain items list') + + imported_at = _now_iso() + conn, _ = _connect_db(db_path) + created = 0 + updated = 0 + skipped = 0 + conflicts: list[dict[str, str]] = [] + try: + for raw in payload['items']: + if not isinstance(raw, dict): + continue + row = {k: str(v) for k, v in raw.items() if k != 'case'} + case = str(raw.get('case', '')).strip() + job_id = row.get('job_id', '').strip() + if not case or not job_id: + conflicts.append({'reason': 'missing_case_or_job_id', 'case': case, 'job_id': job_id}) + continue + + local_by_job = conn.execute('SELECT "case", updated_at FROM sim_cases WHERE job_id = ?', (job_id,)).fetchone() + if local_by_job is None: + case_taken = conn.execute('SELECT "case" FROM sim_cases WHERE "case" = ?', (case,)).fetchone() + if case_taken is not None: + conflicts.append({'reason': 'case_name_taken_by_other_job', 'case': case, 'job_id': job_id}) + continue + _insert_full_case(conn, case, row) + _upsert_sync_state(conn, job_id, synced_updated_at=row.get('updated_at', ''), imported_at=imported_at) + created += 1 + continue + + local_case = str(local_by_job['case']) + local_updated = str(local_by_job['updated_at'] or '') + remote_updated = row.get('updated_at', '') + if remote_updated > local_updated: + _replace_full_case(conn, local_case, row) + _upsert_sync_state(conn, job_id, synced_updated_at=remote_updated, imported_at=imported_at) + updated += 1 + elif remote_updated == local_updated: + _upsert_sync_state(conn, job_id, synced_updated_at=remote_updated, imported_at=imported_at) + skipped += 1 + else: + conflicts.append({ + 'reason': 'local_newer', + 'case': local_case, + 'job_id': job_id, + 'local_updated_at': local_updated, + 'incoming_updated_at': remote_updated, + }) + + conn.commit() + return { + 'ok': True, + 'imported_file': str(in_file), + 'created': created, + 'updated': updated, + 'skipped': skipped, + 'conflicts': conflicts, + } + finally: + conn.close() + + def main(argv: list[str] | None = None) -> int: parser = _build_cli() args = parser.parse_args(argv) @@ -737,6 +973,35 @@ def main(argv: list[str] | None = None) -> int: print(f'{case}: {row}') elif args.command == 'import-csv': import_csv(args.csv, args.db) + elif args.command == 'sync-status': + status = sync_status(args.db) + print( + f"total={status['total_cases']} pending={status['pending_cases']} synced={status['synced_cases']} " + f"last_export={status['last_exported_at'] or '-'} last_import={status['last_imported_at'] or '-'}" + ) + pending_rows = status['pending'] + if pending_rows: + if args.table: + print(_format_table(pending_rows)) + else: + for row in pending_rows: + case = row.pop('case') + print(f"{case}: {row}") + else: + print('(no pending rows)') + elif args.command == 'sync-export': + out = sync_export(args.db, args.out, include_all=args.all) + print(f"Exported {out['exported']} rows to {out['path']} at {out['exported_at']}") + elif args.command == 'sync-import': + out = sync_import(args.db, args.in_path) + print( + f"Imported {out['created']} new, {out['updated']} updated, {out['skipped']} unchanged " + f"from {out['imported_file']}" + ) + if out['conflicts']: + print('Conflicts:') + for conflict in out['conflicts']: + print(json.dumps(conflict, ensure_ascii=False, sort_keys=True)) else: parser.print_help() return 1 diff --git a/sim_db_client.py b/sim_db_client.py index 85f6b60..5a858d5 100644 --- a/sim_db_client.py +++ b/sim_db_client.py @@ -1,441 +1,14 @@ -"""Tiny REST client for mini_sim_db server with local durable dual-write.""" +"""Backward-compatible entrypoint for remote_api client.""" -from __future__ import annotations +from remote_api.client import RemoteRequestError, RemoteResponseError, RemoteTransportError, SimDbClient, main -import argparse -import json -import os -import socket -import sys -from typing import Any -from urllib import error, parse, request - - -class RemoteRequestError(RuntimeError): - """Base error for remote request failures.""" - - -class RemoteTransportError(RemoteRequestError): - """Network/transport failure reaching remote server.""" - - -class RemoteResponseError(RemoteRequestError): - """Remote server responded with an HTTP/application error.""" - -from sim_db import _read_sim_db, add_sim_item, del_cases, init_sim_db, mark_done, resolve_case_ref, upd_cases - - -def _case_ref(*, case: str | None, job_id: str | None) -> str: - if bool(case) == bool(job_id): - raise ValueError("use exactly one of case or job_id") - return str(case or job_id) - - -class SimDbClient: - def __init__( - self, - base_url: str, - token: str, - timeout: float = 10.0, - local_db_path: str | None = None, - enable_local_write: bool = True, - ) -> None: - self.base_url = base_url.rstrip("/") - self.token = token - self.timeout = timeout - self.local_db_path = os.path.expanduser(local_db_path) if local_db_path else None - self.enable_local_write = enable_local_write - - def health(self) -> dict[str, Any]: - return self._request("GET", "/health") - - def init(self, db_path: str | None = None) -> dict[str, Any]: - payload: dict[str, Any] = {} - if db_path: - payload["db_path"] = db_path - return self._request("POST", "/init", payload) - - def create( - self, - *, - case: str, - bin_name: str, - status: str, - inp: str | None = None, - input_files: list[str] | None = None, - note: str | None = None, - work_dir: str | None = None, - extra_params: str | None = None, - db_path: str | None = None, - run_host: str | None = None, - ) -> dict[str, Any]: - payload: dict[str, Any] = { - "case": case, - "bin_name": bin_name, - "status": status, - "run_host": run_host or socket.gethostname(), - } - if inp is not None: - payload["inp"] = inp - if input_files is not None: - payload["input_files"] = input_files - if note is not None: - payload["note"] = note - if work_dir is not None: - payload["work_dir"] = work_dir - if extra_params is not None: - payload["extra_params"] = extra_params - if db_path is not None: - payload["db_path"] = db_path - return self._dual_write("create", payload) - - def add(self, **kwargs: Any) -> dict[str, Any]: - return self.create(**kwargs) - - def read(self, *, case: str | None = None, job_id: str | None = None, db_path: str | None = None) -> dict[str, Any]: - case_ref = _case_ref(case=case, job_id=job_id) - path = f"/cases/{parse.quote(case_ref, safe='')}" - if db_path: - path += "?" + parse.urlencode({"db_path": db_path}) - return self._request("GET", path) - - def done( - self, - *, - case: str | None = None, - job_id: str | None = None, - db_path: str | None = None, - run_host: str | None = None, - ) -> dict[str, Any]: - return self.update( - case=case, - job_id=job_id, - fields={"status": "done"}, - db_path=db_path, - run_host=run_host, - ) - - def update( - self, - *, - case: str | None = None, - job_id: str | None = None, - fields: dict[str, Any], - db_path: str | None = None, - run_host: str | None = None, - ) -> dict[str, Any]: - payload: dict[str, Any] = { - "case": case, - "job_id": job_id, - "fields": fields, - "run_host": run_host or socket.gethostname(), - } - if db_path is not None: - payload["db_path"] = db_path - return self._dual_write("update", payload) - - def delete( - self, - *, - case: str | None = None, - job_id: str | None = None, - db_path: str | None = None, - run_host: str | None = None, - ) -> dict[str, Any]: - payload: dict[str, Any] = { - "case": case, - "job_id": job_id, - "run_host": run_host or socket.gethostname(), - } - if db_path is not None: - payload["db_path"] = db_path - return self._dual_write("delete", payload) - - def list(self, db_path: str | None = None) -> dict[str, Any]: - path = "/cases" - if db_path: - path += "?" + parse.urlencode({"db_path": db_path}) - return self._request("GET", path) - - def summary( - self, - *, - db_path: str | None = None, - status: str | None = None, - run_host: str | None = None, - limit: int | None = None, - sort_by: str = "updated_at", - order: str = "desc", - ) -> dict[str, Any]: - query: dict[str, Any] = {"sort_by": sort_by, "order": order} - if db_path: - query["db_path"] = db_path - if status: - query["status"] = status - if run_host: - query["run_host"] = run_host - if limit is not None: - query["limit"] = limit - return self._request("GET", "/cases/summary?" + parse.urlencode(query)) - - def _dual_write(self, op: str, payload: dict[str, Any]) -> dict[str, Any]: - local_ok = None - local_error = None - if self.enable_local_write and self.local_db_path: - try: - self._apply_local(op, payload) - local_ok = True - except Exception as exc: - local_ok = False - local_error = str(exc) - - try: - remote = self._request_for_op(op, payload) - out = {"ok": True, "remote_ok": True, "remote": remote} - if local_ok is not None: - out["local_ok"] = local_ok - if local_error: - out["local_error"] = local_error - return out - except RemoteTransportError as exc: - if local_ok: - return { - "ok": True, - "remote_ok": False, - "remote_error": str(exc), - "local_ok": True, - "fallback": "local-only", - } - raise - - def _apply_local(self, op: str, payload: dict[str, Any]) -> None: - assert self.local_db_path is not None - init_sim_db(self.local_db_path) - run_host = payload.get("run_host") - - if op == "create": - if payload.get("job_id") not in (None, ""): - raise ValueError("field 'job_id' is auto-generated and cannot be set on create") - - case = payload["case"] - add_sim_item( - case=case, - inp=payload.get("inp"), - input_files=payload.get("input_files"), - bin_name=payload["bin_name"], - status=payload["status"], - db_path=self.local_db_path, - note=payload.get("note"), - work_dir=payload.get("work_dir"), - extra_params=payload.get("extra_params"), - ) - if run_host: - upd_cases(self.local_db_path, {case: {"run_host": str(run_host)}}) - return - - if op == "update": - _, rows = _read_sim_db(self.local_db_path) - case = resolve_case_ref(rows, _case_ref(case=payload.get("case"), job_id=payload.get("job_id"))) - fields = dict(payload.get("fields") or {}) - if run_host: - fields["run_host"] = str(run_host) - if fields.get("status") == "done": - mark_done(case=case, db_path=self.local_db_path) - fields.pop("status", None) - fields.pop("state_changed_at", None) - if fields: - upd_cases(self.local_db_path, {case: fields}) - return - - if op == "delete": - _, rows = _read_sim_db(self.local_db_path) - case = resolve_case_ref(rows, _case_ref(case=payload.get("case"), job_id=payload.get("job_id"))) - del_cases(self.local_db_path, [case]) - return - - raise ValueError(f"unsupported op: {op}") - - def _request_for_op(self, op: str, payload: dict[str, Any]) -> dict[str, Any]: - if op == "create": - return self._request("POST", "/cases", payload) - if op == "update": - case_ref = parse.quote(_case_ref(case=payload.get("case"), job_id=payload.get("job_id")), safe='') - remote_payload = {k: v for k, v in payload.items() if k not in {"case", "job_id"}} - return self._request("PATCH", f"/cases/{case_ref}", remote_payload) - if op == "delete": - case_ref = parse.quote(_case_ref(case=payload.get("case"), job_id=payload.get("job_id")), safe='') - db_path = payload.get("db_path") - path = f"/cases/{case_ref}" - if db_path: - path += "?" + parse.urlencode({"db_path": db_path}) - return self._request("DELETE", path) - raise ValueError(f"unsupported op: {op}") - - def _request(self, method: str, path: str, payload: dict[str, Any] | None = None) -> dict[str, Any]: - data = None - headers = {"Authorization": f"Bearer {self.token}"} - if payload is not None: - data = json.dumps(payload).encode("utf-8") - headers["Content-Type"] = "application/json" - - req = request.Request(self.base_url + path, method=method, headers=headers, data=data) - try: - with request.urlopen(req, timeout=self.timeout) as resp: - body = resp.read().decode("utf-8") - return json.loads(body) if body else {} - except error.HTTPError as exc: - body = exc.read().decode("utf-8") - msg = body or str(exc) - raise RemoteResponseError(f"HTTP {exc.code}: {msg}") from exc - except (error.URLError, TimeoutError) as exc: - raise RemoteTransportError(f"request failed: {exc}") from exc - - -def _parse_fields(pairs: list[str] | None) -> dict[str, str]: - fields: dict[str, str] = {} - for pair in pairs or []: - if "=" not in pair: - raise ValueError(f"invalid --field '{pair}', expected key=value") - key, value = pair.split("=", 1) - key = key.strip() - if not key: - raise ValueError(f"invalid --field '{pair}', empty key") - fields[key] = value - return fields - - -def _add_create_like_args(parser: argparse.ArgumentParser) -> None: - parser.add_argument("--case", required=True) - parser.add_argument("--bin", dest="bin_name", required=True) - parser.add_argument("--status", required=True) - parser.add_argument("--inp", default=None) - parser.add_argument("--input-file", action="append", default=None) - parser.add_argument("--note", default=None) - parser.add_argument("--work-dir", default=None) - parser.add_argument("--extra-params", default=None, help="Raw extra runtime params string (for example JSON)") - parser.add_argument("--db", default=None) - - -def _build_parser() -> argparse.ArgumentParser: - p = argparse.ArgumentParser(description="mini_sim_db REST client") - p.add_argument("--url", default="http://127.0.0.1:8765", help="Server base URL") - p.add_argument("--token", default=None, help="Bearer token (or SIM_DB_API_TOKEN)") - p.add_argument( - "--local-db", - default=os.path.expanduser("~/.sim_db_client_local.sqlite3"), - help="Local durable mirror DB for dual-write fallback (default: ~/.sim_db_client_local.sqlite3)", - ) - p.add_argument("--no-local-write", action="store_true", help="Disable local dual-write fallback") - sub = p.add_subparsers(dest="cmd", required=True) - - sub.add_parser("health") - - p_init = sub.add_parser("init") - p_init.add_argument("--db", default=None, help="Optional remote db_path override") - - p_create = sub.add_parser("create") - _add_create_like_args(p_create) - - p_add = sub.add_parser("add") - _add_create_like_args(p_add) - - p_read = sub.add_parser("read") - read_target = p_read.add_mutually_exclusive_group(required=True) - read_target.add_argument("--case") - read_target.add_argument("--job-id", dest="job_id") - p_read.add_argument("--db", default=None) - - p_update = sub.add_parser("update") - update_target = p_update.add_mutually_exclusive_group(required=True) - update_target.add_argument("--case") - update_target.add_argument("--job-id", dest="job_id") - p_update.add_argument("--field", action="append", default=None, help="key=value (repeatable)") - p_update.add_argument("--db", default=None) - - p_done = sub.add_parser("done") - done_target = p_done.add_mutually_exclusive_group(required=True) - done_target.add_argument("--case") - done_target.add_argument("--job-id", dest="job_id") - p_done.add_argument("--db", default=None) - - p_delete = sub.add_parser("delete") - delete_target = p_delete.add_mutually_exclusive_group(required=True) - delete_target.add_argument("--case") - delete_target.add_argument("--job-id", dest="job_id") - p_delete.add_argument("--db", default=None) - - p_list = sub.add_parser("list") - p_list.add_argument("--db", default=None) - - p_summary = sub.add_parser("summary") - p_summary.add_argument("--db", default=None) - p_summary.add_argument("--status", default=None) - p_summary.add_argument("--run-host", default=None) - p_summary.add_argument("--limit", type=int, default=None) - p_summary.add_argument("--sort-by", default="updated_at") - p_summary.add_argument("--order", choices=["asc", "desc"], default="desc") - - return p - - -def main(argv: list[str] | None = None) -> int: - args = _build_parser().parse_args(argv) - token = args.token or os.getenv("SIM_DB_API_TOKEN") - if not token: - print("Missing token. Set --token or SIM_DB_API_TOKEN", file=sys.stderr) - return 2 - - client = SimDbClient( - base_url=args.url, - token=token, - local_db_path=None if args.no_local_write else args.local_db, - enable_local_write=not args.no_local_write, - ) - - try: - if args.cmd == "health": - result = client.health() - elif args.cmd == "init": - result = client.init(db_path=args.db) - elif args.cmd in {"create", "add"}: - result = client.create( - case=args.case, - inp=args.inp, - input_files=args.input_file, - bin_name=args.bin_name, - status=args.status, - note=args.note, - work_dir=args.work_dir, - extra_params=args.extra_params, - db_path=args.db, - ) - elif args.cmd == "read": - result = client.read(case=args.case, job_id=args.job_id, db_path=args.db) - elif args.cmd == "update": - result = client.update(case=args.case, job_id=args.job_id, fields=_parse_fields(args.field), db_path=args.db) - elif args.cmd == "done": - result = client.done(case=args.case, job_id=args.job_id, db_path=args.db) - elif args.cmd == "delete": - result = client.delete(case=args.case, job_id=args.job_id, db_path=args.db) - elif args.cmd == "list": - result = client.list(db_path=args.db) - elif args.cmd == "summary": - result = client.summary( - db_path=args.db, - status=args.status, - run_host=args.run_host, - limit=args.limit, - sort_by=args.sort_by, - order=args.order, - ) - else: - return 1 - except (RuntimeError, ValueError) as exc: - print(str(exc), file=sys.stderr) - return 2 - - print(json.dumps(result, indent=2, sort_keys=True)) - return 0 +__all__ = [ + "RemoteRequestError", + "RemoteTransportError", + "RemoteResponseError", + "SimDbClient", + "main", +] if __name__ == "__main__": diff --git a/sim_db_server.py b/sim_db_server.py index d660c32..a0fc222 100644 --- a/sim_db_server.py +++ b/sim_db_server.py @@ -1,400 +1,8 @@ -"""HTTP host for mini_sim_db. +"""Backward-compatible entrypoint for remote_api server.""" -Stdlib-only JSON API around sim_db.py for centralized CRUD updates. -""" +from remote_api.server import SecurityPolicy, SimDbApiServer, SimDbRequestHandler, main -from __future__ import annotations - -import argparse -import json -import os -import threading -from datetime import datetime -from pathlib import Path -from http import HTTPStatus -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from typing import Any -from urllib.parse import parse_qs, unquote, urlsplit - -from sim_db import ( - ALLOWED_STATUS, - DEFAULT_DB_PATH, - add_sim_item, - del_cases, - init_sim_db, - list_items, - list_view, - upd_cases, -) - - -def _now_iso() -> str: - return datetime.now().isoformat(timespec="milliseconds") - - -class SecurityPolicy: - """Validates API token and requested DB paths.""" - - def __init__( - self, - token: str, - default_db_path: str, - allowed_db_path: str | None = None, - allowed_base_dir: str | None = None, - ) -> None: - self.token = token - self.default_db_path = _norm(default_db_path) - self.allowed_db_path = _norm(allowed_db_path) if allowed_db_path else None - self.allowed_base_dir = _norm(allowed_base_dir) if allowed_base_dir else None - - def is_authorized(self, auth_header: str | None) -> bool: - if not auth_header or not auth_header.startswith("Bearer "): - return False - return auth_header[len("Bearer ") :] == self.token - - def resolve_db_path(self, requested_db_path: str | None) -> str: - if not requested_db_path: - return self.default_db_path - - wanted = _norm(requested_db_path) - - if self.allowed_db_path is None and self.allowed_base_dir is None: - if wanted == self.default_db_path: - return wanted - raise ValueError("db_path is not allowed") - - if self.allowed_db_path and wanted == self.allowed_db_path: - return wanted - - if self.allowed_base_dir and _is_within_base(wanted, self.allowed_base_dir): - return wanted - - raise ValueError("db_path is outside allowed scope") - - -class SimDbApiServer(ThreadingHTTPServer): - def __init__(self, server_address: tuple[str, int], policy: SecurityPolicy) -> None: - self.policy = policy - self.mutation_lock = threading.Lock() - super().__init__(server_address, SimDbRequestHandler) - - -class SimDbRequestHandler(BaseHTTPRequestHandler): - server: SimDbApiServer - - def do_GET(self) -> None: # noqa: N802 - if self.path == "/health": - self._json(HTTPStatus.OK, {"ok": True}) - return - - if self.path.startswith("/cases"): - self._require_auth() - if not self._authorized: - return - - query = self._query_dict() - try: - db_path = self.server.policy.resolve_db_path(query.get("db_path")) - if urlsplit(self.path).path == "/cases/summary": - status = query.get("status") - run_host = query.get("run_host") - limit = int(query["limit"]) if query.get("limit") else None - sort_by = query.get("sort_by", "updated_at") - desc = query.get("order", "desc").lower() != "asc" - with self.server.mutation_lock: - rows = list_view(db_path=db_path, status=status, run_host=run_host, sort_by=sort_by, desc=desc, limit=limit) - self._json(HTTPStatus.OK, {"db_path": db_path, "count": len(rows), "items": rows}) - return - - case_ref = self._case_from_path() - with self.server.mutation_lock: - data = list_items(db_path) - if case_ref: - case = self._resolve_case_ref_from_table(data, case_ref) - self._json(HTTPStatus.OK, {"db_path": db_path, "case": case, "item": data[case]}) - return - self._json(HTTPStatus.OK, {"db_path": db_path, "cases": data}) - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) - return - - self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) - - def do_POST(self) -> None: # noqa: N802 - if self.path == "/init": - self._require_auth() - if not self._authorized: - return - payload = self._read_json_body() - if payload is None: - return - try: - db_path = self.server.policy.resolve_db_path(payload.get("db_path")) - with self.server.mutation_lock: - init_sim_db(db_path) - self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path}) - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) - return - - if self.path in {"/add", "/done", "/update", "/delete"}: - self._legacy_mutating_routes() - return - - if self.path == "/cases": - self._require_auth() - if not self._authorized: - return - payload = self._read_json_body() - if payload is None: - return - self._create_case(payload) - return - - self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) - - def do_PATCH(self) -> None: # noqa: N802 - if not self.path.startswith("/cases/"): - self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) - return - self._require_auth() - if not self._authorized: - return - - payload = self._read_json_body() - if payload is None: - return - - case = self._case_from_path() - if not case: - self._json(HTTPStatus.BAD_REQUEST, {"error": "missing case in path"}) - return - - self._update_case(case, payload) - - def do_DELETE(self) -> None: # noqa: N802 - if not self.path.startswith("/cases/"): - self._json(HTTPStatus.NOT_FOUND, {"error": "not found"}) - return - self._require_auth() - if not self._authorized: - return - - case_ref = self._case_from_path() - if not case_ref: - self._json(HTTPStatus.BAD_REQUEST, {"error": "missing case in path"}) - return - - query = self._query_dict() - req_db = query.get("db_path") - try: - db_path = self.server.policy.resolve_db_path(req_db) - with self.server.mutation_lock: - table = list_items(db_path) - case = self._resolve_case_ref_from_table(table, case_ref) - del_cases(db_path, [case]) - self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": case}) - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) - - def _legacy_mutating_routes(self) -> None: - self._require_auth() - if not self._authorized: - return - payload = self._read_json_body() - if payload is None: - return - - if self.path == "/add": - self._create_case(payload) - return - if self.path == "/done": - case_ref = payload.get("case") or payload.get("job_id") - if not case_ref: - self._json(HTTPStatus.BAD_REQUEST, {"error": "missing field: case or job_id"}) - return - self._update_case(case_ref, {"fields": {"status": "done"}, **payload}) - return - if self.path == "/update": - case_ref = payload.get("case") or payload.get("job_id") - if not case_ref: - self._json(HTTPStatus.BAD_REQUEST, {"error": "missing field: case or job_id"}) - return - self._update_case(case_ref, payload) - return - if self.path == "/delete": - case_ref = payload.get("case") or payload.get("job_id") - if not case_ref: - self._json(HTTPStatus.BAD_REQUEST, {"error": "missing field: case or job_id"}) - return - try: - db_path = self.server.policy.resolve_db_path(payload.get("db_path")) - with self.server.mutation_lock: - table = list_items(db_path) - case = self._resolve_case_ref_from_table(table, case_ref) - del_cases(db_path, [case]) - self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": case}) - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) - - def _create_case(self, payload: dict[str, Any]) -> None: - try: - if payload.get("job_id") not in (None, ""): - raise ValueError("field 'job_id' is auto-generated and cannot be set on create") - - db_path = self.server.policy.resolve_db_path(payload.get("db_path")) - with self.server.mutation_lock: - add_sim_item( - case=payload["case"], - inp=payload.get("inp"), - input_files=payload.get("input_files"), - bin_name=payload["bin_name"], - status=payload["status"], - db_path=db_path, - note=payload.get("note"), - work_dir=payload.get("work_dir"), - extra_params=payload.get("extra_params"), - ) - run_host = payload.get("run_host") - if run_host: - upd_cases(db_path, {payload["case"]: {"run_host": str(run_host)}}) - self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": payload["case"]}) - except KeyError as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": f"missing field: {exc.args[0]}"}) - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) - - @staticmethod - def _resolve_case_ref_from_table(data: dict[str, dict[str, Any]], case_ref: str) -> str: - if case_ref in data: - return case_ref - - matches = [case for case, item in data.items() if item.get("job_id") == case_ref] - if not matches: - raise ValueError(f"case/job_id not found: {case_ref}") - if len(matches) > 1: - joined = ", ".join(sorted(matches)) - raise ValueError(f"job_id matches multiple cases ({joined}), use case explicitly") - return matches[0] - - def _update_case(self, case_ref: str, payload: dict[str, Any]) -> None: - try: - db_path = self.server.policy.resolve_db_path(payload.get("db_path")) - fields = payload.get("fields") - if fields is None: - fields = {k: v for k, v in payload.items() if k not in {"case", "job_id", "db_path", "run_host"}} - if not isinstance(fields, dict) or not fields: - raise ValueError("fields must be a non-empty object") - if "case" in fields: - raise ValueError("field 'case' is immutable") - - status = fields.get("status") - if status is not None: - if status not in ALLOWED_STATUS: - allowed = ", ".join(sorted(ALLOWED_STATUS)) - raise ValueError(f"Invalid status '{status}'. Allowed: {allowed}") - fields["state_changed_at"] = _now_iso() - - fields["updated_at"] = _now_iso() - - run_host = payload.get("run_host") - if run_host: - fields["run_host"] = str(run_host) - - with self.server.mutation_lock: - table = list_items(db_path) - case = self._resolve_case_ref_from_table(table, case_ref) - upd_cases(db_path, {case: fields}) - self._json(HTTPStatus.OK, {"ok": True, "db_path": db_path, "case": case, "updated": sorted(fields.keys())}) - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": str(exc)}) - - def log_message(self, fmt: str, *args: Any) -> None: - return - - def _read_json_body(self) -> dict[str, Any] | None: - try: - raw_len = int(self.headers.get("Content-Length", "0")) - raw = self.rfile.read(raw_len) if raw_len > 0 else b"{}" - payload = json.loads(raw.decode("utf-8")) if raw else {} - if not isinstance(payload, dict): - raise ValueError("JSON body must be an object") - return payload - except Exception as exc: - self._json(HTTPStatus.BAD_REQUEST, {"error": f"invalid JSON: {exc}"}) - return None - - def _query_dict(self) -> dict[str, str]: - query = parse_qs(urlsplit(self.path).query, keep_blank_values=True) - out: dict[str, str] = {} - for k, vals in query.items(): - if vals: - out[k] = vals[0] - return out - - def _case_from_path(self) -> str | None: - path = urlsplit(self.path).path - parts = [p for p in path.split("/") if p] - if len(parts) >= 2 and parts[0] == "cases": - return unquote(parts[1]) - return None - - def _require_auth(self) -> None: - self._authorized = self.server.policy.is_authorized(self.headers.get("Authorization")) - if not self._authorized: - self._json(HTTPStatus.UNAUTHORIZED, {"error": "unauthorized"}) - - def _json(self, status: HTTPStatus, obj: dict[str, Any]) -> None: - data = json.dumps(obj, ensure_ascii=False).encode("utf-8") - self.send_response(int(status)) - self.send_header("Content-Type", "application/json; charset=utf-8") - self.send_header("Content-Length", str(len(data))) - self.end_headers() - self.wfile.write(data) - - -def _norm(path: str) -> str: - return str(Path(path).expanduser().resolve()) - - -def _is_within_base(path: str, base_dir: str) -> bool: - target = Path(path).expanduser().resolve() - base = Path(base_dir).expanduser().resolve() - try: - target.relative_to(base) - return True - except ValueError: - return False - - -def _build_parser() -> argparse.ArgumentParser: - p = argparse.ArgumentParser(description="mini_sim_db HTTP server") - p.add_argument("--host", default="127.0.0.1", help="Bind host (default: 127.0.0.1)") - p.add_argument("--port", type=int, default=8765, help="Bind port (default: 8765)") - p.add_argument("--db", default=DEFAULT_DB_PATH, help="Default DB path (default: ~/sim_db.csv)") - p.add_argument("--allowed-db-path", default=None, help="Optional exact writable DB path") - p.add_argument("--allowed-base-dir", default=None, help="Optional writable base directory") - p.add_argument("--token", default=None, help="Bearer token (or use SIM_DB_API_TOKEN env)") - return p - - -def main(argv: list[str] | None = None) -> int: - args = _build_parser().parse_args(argv) - token = args.token or os.getenv("SIM_DB_API_TOKEN") - if not token: - raise SystemExit("Missing token. Set --token or SIM_DB_API_TOKEN") - - policy = SecurityPolicy( - token=token, - default_db_path=args.db, - allowed_db_path=args.allowed_db_path, - allowed_base_dir=args.allowed_base_dir, - ) - server = SimDbApiServer((args.host, args.port), policy) - print(f"Serving mini_sim_db API at http://{args.host}:{args.port}") - print(f"Default DB path: {policy.default_db_path}") - server.serve_forever() - return 0 +__all__ = ["SecurityPolicy", "SimDbApiServer", "SimDbRequestHandler", "main"] if __name__ == "__main__": diff --git a/test_sim_db.py b/test_sim_db.py index a7aaeea..b1ee267 100644 --- a/test_sim_db.py +++ b/test_sim_db.py @@ -1,3 +1,4 @@ +import json import os import sqlite3 import sys @@ -19,6 +20,9 @@ list_view, mark_done, search_sim_db, + sync_export, + sync_import, + sync_status, upd_cases, ) @@ -213,5 +217,51 @@ def test_cli_table_view(self): self.assertIn('c2', r_list.stdout) +class TestLocalSync(unittest.TestCase): + def setUp(self): + self.tmp_dir = tempfile.TemporaryDirectory() + self.db_path = os.path.join(self.tmp_dir.name, 'sync.sqlite3') + self.sync_file = os.path.join(self.tmp_dir.name, 'sync.json') + init_sim_db(self.db_path) + + def tearDown(self): + self.tmp_dir.cleanup() + + def test_sync_export_and_pending_status(self): + add_sim_item(case='s1', inp='a.inp', bin_name='solver', status='start', db_path=self.db_path) + status_before = sync_status(self.db_path) + self.assertEqual(status_before['pending_cases'], 1) + + out = sync_export(self.db_path, self.sync_file) + self.assertEqual(out['exported'], 1) + + status_after = sync_status(self.db_path) + self.assertEqual(status_after['pending_cases'], 0) + + def test_sync_import_conflict_policy_local_newer_wins(self): + add_sim_item(case='s1', inp='a.inp', bin_name='solver', status='start', db_path=self.db_path) + sync_export(self.db_path, self.sync_file) + + data = list_items(self.db_path) + local = data['s1'] + older_remote = { + 'format': 'mini_sim_db_sync_v1', + 'items': [ + { + 'case': 's1', + **local, + 'updated_at': '2000-01-01T00:00:00.000', + 'note': 'remote older', + } + ], + } + with open(self.sync_file, 'w', encoding='utf-8') as f: + json.dump(older_remote, f) + + out = sync_import(self.db_path, self.sync_file) + self.assertEqual(len(out['conflicts']), 1) + self.assertEqual(out['conflicts'][0]['reason'], 'local_newer') + + if __name__ == '__main__': unittest.main() diff --git a/test_sim_db_rest.py b/test_sim_db_rest.py index 6bb9049..6c673be 100644 --- a/test_sim_db_rest.py +++ b/test_sim_db_rest.py @@ -6,10 +6,10 @@ import unittest from unittest import mock -import sim_db_server +import remote_api.server as sim_db_server from sim_db import list_items -from sim_db_client import SimDbClient -from sim_db_server import SecurityPolicy, SimDbApiServer +from remote_api.client import SimDbClient +from remote_api.server import SecurityPolicy, SimDbApiServer class RestServerTestCase(unittest.TestCase): @@ -197,7 +197,7 @@ def guarded_add(*args, **kwargs): gate.release() try: - with mock.patch("sim_db_server.add_sim_item", side_effect=guarded_add): + with mock.patch("remote_api.server.add_sim_item", side_effect=guarded_add): client = SimDbClient(base_url=url, token=self.token, enable_local_write=False) client.init() @@ -256,7 +256,7 @@ def guarded_list(db_path): client.init() client.create(case="c1", inp="a.inp", bin_name="solver", status="start") - with mock.patch("sim_db_server.list_items", side_effect=guarded_list): + with mock.patch("remote_api.server.list_items", side_effect=guarded_list): writer_done = threading.Event() def _update():