Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Empty file.
Empty file.
Empty file.
50 changes: 50 additions & 0 deletions miles/utils/workers/rpc/common/metadata.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
from __future__ import annotations

import dataclasses
import inspect
import typing
from collections.abc import Callable
from typing import Any

from miles.utils.workers.rpc.common.serialization import RpcSerializer


@dataclasses.dataclass(frozen=True)
class RpcMethodSpec:
name: str
serializer: RpcSerializer


def collect_rpc_method_specs(worker_cls: type) -> dict[str, RpcMethodSpec]:
specs: dict[str, RpcMethodSpec] = {}

for name in sorted(dir(worker_cls)):
if name.startswith("_"):
continue
static_attr = inspect.getattr_static(worker_cls, name)
if isinstance(static_attr, (classmethod, staticmethod, property)):
continue
if not callable(static_attr):
continue
specs[name] = _build_method_spec(worker_cls=worker_cls, name=name, fn=inspect.unwrap(static_attr))

return specs


def _build_method_spec(*, worker_cls: type, name: str, fn: Callable[..., Any]) -> RpcMethodSpec:
signature = inspect.signature(fn)
hints = typing.get_type_hints(fn, include_extras=True)

query_fields: dict[str, Any] = {}
for param in list(signature.parameters.values())[1:]:
default = ... if param.default is inspect.Parameter.empty else param.default
query_fields[param.name] = (hints[param.name], default)

return RpcMethodSpec(
name=name,
serializer=RpcSerializer.create(
query_model_name=f"{worker_cls.__name__}{name.title().replace('_', '')}Query",
query_fields=query_fields,
result_annotation=hints["return"],
),
)
30 changes: 30 additions & 0 deletions miles/utils/workers/rpc/common/protocol.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
from __future__ import annotations

from typing import Any, Literal

from miles.utils.pydantic_utils import StrictBaseModel

HEALTH_PATH = "/v1/health"
CALL_STATUS_PATH = "/v1/calls/{call_id}"
SUBMIT_PATH = "/v1/{method_name}"

DEFAULT_POLL_TIMEOUT_SECONDS = 30.0


class SubmitRequest(StrictBaseModel):
call_id: str
query: dict[str, Any]


class SubmitResponse(StrictBaseModel):
status: Literal["submitted"] = "submitted"


class CallStatusResponse(StrictBaseModel):
status: Literal["pending", "success", "failed"]
result: Any = None
error: str | None = None


class HealthResponse(StrictBaseModel):
status: Literal["ok"] = "ok"
29 changes: 29 additions & 0 deletions miles/utils/workers/rpc/common/serialization.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
from __future__ import annotations

import dataclasses
from typing import Any

from pydantic import BaseModel, ConfigDict, TypeAdapter, create_model


@dataclasses.dataclass(frozen=True)
class RpcSerializer:
query_model: type[BaseModel]
result_adapter: TypeAdapter[Any]

@classmethod
def create(cls, *, query_model_name: str, query_fields: dict[str, Any], result_annotation: Any) -> RpcSerializer:
query_model = create_model(query_model_name, __config__=ConfigDict(extra="forbid"), **query_fields)
return cls(query_model=query_model, result_adapter=TypeAdapter(result_annotation))

def encode_query(self, kwargs: dict[str, Any]) -> dict[str, Any]:
return self.query_model(**kwargs).model_dump(mode="json")

def decode_query(self, query: dict[str, Any]) -> dict[str, Any]:
return dict(self.query_model(**query))

def encode_result(self, result: Any) -> Any:
return self.result_adapter.dump_python(result, mode="json")

def decode_result(self, payload: Any) -> Any:
return self.result_adapter.validate_python(payload)
Empty file.
47 changes: 47 additions & 0 deletions miles/utils/workers/rpc/server/app.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
from __future__ import annotations

import logging

from fastapi import FastAPI, Query, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse

from miles.utils.workers.rpc.common.protocol import (
CALL_STATUS_PATH,
DEFAULT_POLL_TIMEOUT_SECONDS,
HEALTH_PATH,
SUBMIT_PATH,
CallStatusResponse,
HealthResponse,
SubmitRequest,
SubmitResponse,
)
from miles.utils.workers.rpc.server.core import RpcServer

logger = logging.getLogger(__name__)


def create_rpc_app(worker: object) -> FastAPI:
server = RpcServer(worker=worker)

app = FastAPI()

@app.exception_handler(RequestValidationError)
async def handle_malformed_request(request: Request, exc: RequestValidationError) -> JSONResponse:
return JSONResponse(status_code=400, content={"detail": str(exc)})

@app.get(HEALTH_PATH)
async def health() -> HealthResponse:
return HealthResponse()

@app.post(SUBMIT_PATH)
async def submit_call(method_name: str, request: SubmitRequest) -> SubmitResponse:
return server.submit_call(method_name=method_name, request=request)

@app.get(CALL_STATUS_PATH)
async def query_call(
call_id: str, timeout: float = Query(default=DEFAULT_POLL_TIMEOUT_SECONDS, ge=0.0)
) -> CallStatusResponse:
return await server.query_call(call_id=call_id, timeout=timeout)

return app
75 changes: 75 additions & 0 deletions miles/utils/workers/rpc/server/core.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
from __future__ import annotations

import functools
import logging
from typing import NoReturn

from fastapi import HTTPException
from pydantic import ValidationError

from miles.utils.tracking_utils.structured_log import log_structured
from miles.utils.workers.rpc.common.metadata import collect_rpc_method_specs
from miles.utils.workers.rpc.common.protocol import CallStatusResponse, SubmitRequest, SubmitResponse
from miles.utils.workers.rpc.server.executor import RpcCallExecutor
from miles.utils.workers.rpc.server.store import CallStore

logger = logging.getLogger(__name__)


class RpcServer:
def __init__(self, *, worker: object) -> None:
self._specs = collect_rpc_method_specs(type(worker))
self._store = CallStore()
self._executor = RpcCallExecutor(worker=worker)
log_structured(
logger.info,
tag="rpc",
op="server",
phase="boot",
worker=type(worker).__name__,
methods=len(self._specs),
)

def submit_call(self, *, method_name: str, request: SubmitRequest) -> SubmitResponse:
def reject(*, status_code: int, reason: str, detail: str) -> NoReturn:
log_structured(
logger.warning,
tag="rpc",
op="submit",
phase="reject",
reason=reason,
method=method_name,
call=request.call_id,
error=detail,
)
raise HTTPException(status_code=status_code, detail=detail)

spec = self._specs.get(method_name)
if spec is None:
reject(status_code=404, reason="unknown_method", detail=f"unknown rpc method {method_name!r}")

try:
kwargs = spec.serializer.decode_query(request.query)
except ValidationError as e:
reject(status_code=400, reason="invalid_query", detail=str(e))

self._store.begin(call_id=request.call_id)

self._executor.start(
spec=spec,
kwargs=kwargs,
call_id=request.call_id,
finish=functools.partial(self._store.finish, call_id=request.call_id),
)

return SubmitResponse()

async def query_call(self, *, call_id: str, timeout: float) -> CallStatusResponse:
if not self._store.contains(call_id):
log_structured(logger.warning, tag="rpc", op="poll", phase="reject", reason="unknown_call", call=call_id)
raise HTTPException(status_code=404, detail=f"unknown call id {call_id!r}")

outcome = await self._store.wait(call_id=call_id, timeout=timeout)
if outcome is None:
return CallStatusResponse(status="pending")
return CallStatusResponse(status=outcome.status, result=outcome.result, error=outcome.error)
47 changes: 47 additions & 0 deletions miles/utils/workers/rpc/server/executor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
from __future__ import annotations

import asyncio
import logging
import time
import traceback
from collections.abc import Callable
from typing import Any

from miles.utils.tracking_utils.structured_log import log_structured
from miles.utils.workers.rpc.common.metadata import RpcMethodSpec
from miles.utils.workers.rpc.common.protocol import CallStatusResponse

logger = logging.getLogger(__name__)


class RpcCallExecutor:
def __init__(self, *, worker: object) -> None:
self._worker = worker
self._background_tasks: set[asyncio.Task[None]] = set()

def start(self, *, spec: RpcMethodSpec, kwargs: dict[str, Any], call_id: str, finish: Callable[..., None]) -> None:
task = asyncio.create_task(self._run(spec=spec, kwargs=kwargs, call_id=call_id, finish=finish))
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)

async def _run(
self, *, spec: RpcMethodSpec, kwargs: dict[str, Any], call_id: str, finish: Callable[..., None]
) -> None:
started_at = time.monotonic()
log_fields = {"tag": "rpc", "op": "execute", "method": spec.name, "call": call_id}
log_structured(logger.debug, phase="start", **log_fields)

try:
result = await self._call_worker(spec=spec, kwargs=kwargs)
outcome = CallStatusResponse(status="success", result=spec.serializer.encode_result(result))
log_structured(
logger.debug, phase="end", ok=True, **log_fields, elapsed_s=round(time.monotonic() - started_at, 3)
)
except Exception as e:
log_structured(logger.error, phase="end", ok=False, **log_fields, exc_info=True)
outcome = CallStatusResponse(status="failed", error="".join(traceback.format_exception(e)))

finish(outcome=outcome)

async def _call_worker(self, *, spec: RpcMethodSpec, kwargs: dict[str, Any]) -> Any:
return await getattr(self._worker, spec.name)(**kwargs)
69 changes: 69 additions & 0 deletions miles/utils/workers/rpc/server/store.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
from __future__ import annotations

import asyncio
import contextlib
import dataclasses
import logging
import time

from miles.utils.tracking_utils.structured_log import log_structured
from miles.utils.workers.rpc.common.protocol import CallStatusResponse

logger = logging.getLogger(__name__)

RETRIEVED_TTL_SECONDS = 300.0


class CallStore:
def __init__(self, *, retrieved_ttl_seconds: float = RETRIEVED_TTL_SECONDS) -> None:
self._retrieved_ttl_seconds = retrieved_ttl_seconds
self._records: dict[str, _CallRecord] = {}

def begin(self, *, call_id: str) -> None:
self._purge_expired()

self._records[call_id] = _CallRecord(finished_event=asyncio.Event())
log_structured(
logger.debug, tag="rpc", op="call_store", phase="accept", call=call_id, tracked=len(self._records)
)

def finish(self, *, call_id: str, outcome: CallStatusResponse) -> None:
record = self._records[call_id]
if record.outcome is not None:
raise RuntimeError(f"call {call_id} finished twice")
record.outcome = outcome
record.finished_event.set()
log_structured(logger.debug, tag="rpc", op="call_store", phase="finish", call=call_id, status=outcome.status)

async def wait(self, *, call_id: str, timeout: float) -> CallStatusResponse | None:
record = self._records[call_id]

with contextlib.suppress(TimeoutError, asyncio.TimeoutError):
await asyncio.wait_for(record.finished_event.wait(), timeout=timeout)

if record.outcome is not None and record.first_retrieved_at is None:
record.first_retrieved_at = time.monotonic()
return record.outcome

def contains(self, call_id: str) -> bool:
return call_id in self._records

def _purge_expired(self) -> None:
now = time.monotonic()
retained = {
call_id: record
for call_id, record in self._records.items()
if record.first_retrieved_at is None or now - record.first_retrieved_at <= self._retrieved_ttl_seconds
}
if len(retained) != len(self._records):
log_structured(
logger.debug, tag="rpc", op="call_store", phase="purge", purged=len(self._records) - len(retained)
)
self._records = retained


@dataclasses.dataclass
class _CallRecord:
finished_event: asyncio.Event
outcome: CallStatusResponse | None = None
first_retrieved_at: float | None = None
Empty file.
36 changes: 36 additions & 0 deletions miles/utils/workers/serving/serve_inner.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
from __future__ import annotations

import argparse
import sys

import uvicorn

from miles.utils.function_registry import load_function
from miles.utils.workers.rpc.server.app import create_rpc_app
from miles.utils.workers.serving.utils import split_worker_argv

DEFAULT_HOST = "0.0.0.0"
DEFAULT_PORT = 8000


def main() -> None:
own_argv, worker_argv = split_worker_argv(sys.argv[1:])
args = parse_own_args(own_argv)

factory = load_function(args.worker)
worker = factory(worker_argv)

app = create_rpc_app(worker)
uvicorn.run(app, host=args.host, port=args.port)


def parse_own_args(own_argv: list[str]) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Serve a worker over rpc")
parser.add_argument("--worker", required=True, help="Worker factory as 'package.module.callable'")
parser.add_argument("--host", default=DEFAULT_HOST)
parser.add_argument("--port", type=int, default=DEFAULT_PORT)
return parser.parse_args(own_argv)


if __name__ == "__main__":
main()
Loading
Loading