yashu2000/TemporalBenchEnv
0
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7"""FastAPI application for TemporalBenchEnv."""8 9import os10from pathlib import Path11 12try:13 from env.config import EnvConfig14 from env.temporal_bench_env import TemporalBenchEnvironment15 from models import TemporalBenchAction, TemporalBenchObservation16except ImportError:17 from ..env.config import EnvConfig18 from ..env.temporal_bench_env import TemporalBenchEnvironment19 from ..models import TemporalBenchAction, TemporalBenchObservation20 21try:22 from openenv.core.env_server import create_app23except ImportError:24 create_app = None # type: ignore25 26 27def _env_factory():28 """Create a fresh environment instance per WebSocket session."""29 bank_dir = os.environ.get("TEMPORALBENCH_QUESTION_BANK_DIR")30 if not bank_dir:31 default = Path(__file__).resolve().parents[1] / "tests" / "fixtures" / "banks"32 if default.is_dir():33 bank_dir = str(default)34 cfg = EnvConfig(question_bank_path=bank_dir) if bank_dir else EnvConfig()35 return TemporalBenchEnvironment(config=cfg)36 37 38if create_app is not None:39 app = create_app(40 _env_factory,41 TemporalBenchAction,42 TemporalBenchObservation,43 env_name="temporal-bench-env",44 max_concurrent_envs=64,45 )46else:47 from fastapi import FastAPI48 49 app = FastAPI(title="temporal-bench-env")50 app.get("/health")(lambda: {"status": "ok"})51 52 53def main(host: str | None = None, port: int | None = None) -> None:54 """55 Entry point for `uv run server` and OpenEnv multi-mode validation.56 57 OpenEnv's validator does a naive substring check for ``main()`` in this58 file, so the ``if __name__ == "__main__"`` block must call ``main()`` with59 no arguments; CLI flags are parsed here via ``parse_known_args``.60 """61 import argparse62 63 import uvicorn64 65 if host is None or port is None:66 parser = argparse.ArgumentParser()67 parser.add_argument("--host", type=str, default="0.0.0.0")68 parser.add_argument("--port", type=int, default=8000)69 ns, _ = parser.parse_known_args()70 if host is None:71 host = ns.host72 if port is None:73 port = ns.port74 75 uvicorn.run(app, host=host, port=port)76 77 78if __name__ == "__main__":79 main()80 