emrekuruu's picture
feat: expose native typed action tools
6fc8ce2 verified
Raw History Blame Contribute Delete
2.51 kB
"""FastAPI and WebSocket entrypoint for MATH OpenEnv."""
from __future__ import annotations
from collections.abc import Callable
import os
from fastapi import FastAPI
from openenv.core.env_server.http_server import create_app
from openenv.core.env_server.interfaces import Environment
from math_env.models import (
MathAction,
MathObservation,
native_tool_specs,
)
from math_env.server.environment import MathEnvironment
def _configured_capacity() -> int:
raw = os.getenv("MAX_CONCURRENT_ENVS", "1")
try:
capacity = int(raw)
except ValueError as error:
raise ValueError(f"MAX_CONCURRENT_ENVS must be an integer, got {raw!r}") from error
if capacity < 1:
raise ValueError("MAX_CONCURRENT_ENVS must be at least 1")
return capacity
def _configured_default_task_id() -> int | None:
value = os.getenv("DEFAULT_TASK_ID")
if value is None or not value.strip():
return None
try:
return int(value)
except ValueError as error:
raise ValueError(f"DEFAULT_TASK_ID must be an integer, got {value!r}") from error
def _configured_environment_class(
default_task_id: int | None,
) -> type[MathEnvironment]:
if default_task_id is None:
return MathEnvironment
class ConfiguredMathEnvironment(MathEnvironment):
def __init__(self) -> None:
super().__init__(default_task_id=default_task_id)
ConfiguredMathEnvironment.__name__ = "MathEnvironment"
return ConfiguredMathEnvironment
def create_math_app(
max_concurrent_envs: int | None = None,
environment_factory: Callable[[], Environment] | None = None,
) -> FastAPI:
"""Create the pinned OpenEnv app without constructing runtime data eagerly."""
capacity = _configured_capacity() if max_concurrent_envs is None else max_concurrent_envs
if capacity < 1:
raise ValueError("MAX_CONCURRENT_ENVS must be at least 1")
factory = environment_factory or _configured_environment_class(
_configured_default_task_id()
)
app = create_app(
factory,
MathAction,
MathObservation,
env_name="math_env",
max_concurrent_envs=capacity,
)
app.add_api_route("/tools", native_tool_specs, methods=["GET"])
return app
app = create_math_app()
def main(host: str = "0.0.0.0", port: int = 8000) -> None:
"""Run the standalone OpenEnv server."""
import uvicorn
uvicorn.run(app, host=host, port=port)
if __name__ == "__main__":
main()