fhaviary 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
aviary/api.py ADDED
@@ -0,0 +1,195 @@
1
+ import os
2
+ import secrets
3
+ import uuid
4
+ from typing import TYPE_CHECKING, Annotated, Any
5
+
6
+ import httpx
7
+
8
+ from aviary.db import (
9
+ EnvironmentBase,
10
+ EnvironmentDB,
11
+ EnvironmentDBBackend,
12
+ EnvironmentDBSchema,
13
+ FrameDB,
14
+ FrameDBSchema,
15
+ )
16
+ from aviary.render import Frame
17
+
18
+ if TYPE_CHECKING:
19
+ from fastapi import FastAPI
20
+
21
+
22
+ class EnvDBClient:
23
+ """Interact with an Environment DB via REST API."""
24
+
25
+ def __init__(
26
+ self,
27
+ server_url: str,
28
+ request_headers: httpx._types.HeaderTypes | None = None,
29
+ request_timeout: float | None = None,
30
+ ):
31
+ self._request_url = server_url
32
+ self._request_headers = request_headers
33
+ self._request_timeout = request_timeout
34
+
35
+ async def write_environment_instance(self, name: str) -> uuid.UUID:
36
+ async with httpx.AsyncClient() as client:
37
+ response = await client.post(
38
+ f"{self._request_url}/environment_instance",
39
+ params={
40
+ "env_name": name,
41
+ },
42
+ headers=self._request_headers,
43
+ timeout=self._request_timeout,
44
+ )
45
+ response.raise_for_status()
46
+ return response.json()
47
+
48
+ async def write_environment_frame(
49
+ self,
50
+ environment_id: uuid.UUID,
51
+ frame: Frame,
52
+ ) -> uuid.UUID:
53
+ async with httpx.AsyncClient() as client:
54
+ frame_data = frame.model_dump()
55
+ response = await client.post(
56
+ f"{self._request_url}/environment_frame",
57
+ params={"environment_id": str(environment_id)},
58
+ json={
59
+ "state": frame_data.get("state"),
60
+ "metadata": frame_data.get("info"),
61
+ },
62
+ headers=self._request_headers,
63
+ timeout=self._request_timeout,
64
+ )
65
+ response.raise_for_status()
66
+ return response.json()
67
+
68
+ async def get_environment_instances(
69
+ self,
70
+ name: str | None = None,
71
+ environment_id: uuid.UUID | None = None,
72
+ ) -> list[EnvironmentDB]:
73
+ async with httpx.AsyncClient() as client:
74
+ params = (
75
+ {"env_name": name} if name else {"environment_id": str(environment_id)}
76
+ )
77
+ response = await client.get(
78
+ f"{self._request_url}/environment_instance",
79
+ params=params,
80
+ headers=self._request_headers,
81
+ timeout=self._request_timeout,
82
+ )
83
+ response.raise_for_status()
84
+ return [EnvironmentDB(**obj) for obj in response.json()]
85
+
86
+ async def get_environment_frames(
87
+ self,
88
+ environment_id: uuid.UUID | None = None,
89
+ frame_id: uuid.UUID | None = None,
90
+ ) -> list[FrameDB]:
91
+ async with httpx.AsyncClient() as client:
92
+ params = (
93
+ {"environment_id": str(environment_id)}
94
+ if environment_id
95
+ else {"frame_id": str(frame_id)}
96
+ )
97
+ response = await client.get(
98
+ f"{self._request_url}/environment_frame",
99
+ params=params,
100
+ headers=self._request_headers,
101
+ timeout=self._request_timeout,
102
+ )
103
+ response.raise_for_status()
104
+ return [FrameDB(**obj) for obj in response.json()]
105
+
106
+
107
+ def make_environment_db_server(render_docs: bool = False) -> "FastAPI": # noqa: C901
108
+ """Make a FastAPI app, an interface for the Environment DB."""
109
+ try:
110
+ from fastapi import Depends, FastAPI, HTTPException, status
111
+ from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
112
+ except ModuleNotFoundError as exc:
113
+ raise ImportError(
114
+ "Environment DB server requires the 'server' extra for 'fastapi'. Please:"
115
+ " `pip install aviary[server]`."
116
+ ) from exc
117
+
118
+ backend = EnvironmentDBBackend()
119
+
120
+ async def ensure_session():
121
+ if not backend.Session:
122
+ await backend.populate_session(
123
+ base=EnvironmentBase, uri=os.environ["ENV_DB_URI"]
124
+ )
125
+ yield
126
+
127
+ asgi_app = FastAPI(
128
+ title="Aviary Environment DB API",
129
+ description="CRUD operations for the Aviary Environment DB",
130
+ # Only render Swagger docs if local since we don't have a login here
131
+ docs_url="/docs" if render_docs else None,
132
+ redoc_url="/redoc" if render_docs else None,
133
+ )
134
+ auth_scheme = HTTPBearer()
135
+
136
+ async def validate_token(
137
+ token: Annotated[HTTPAuthorizationCredentials, Depends(auth_scheme)],
138
+ ) -> HTTPAuthorizationCredentials:
139
+ # NOTE: don't use os.environ.get() to avoid possible empty string matches, and
140
+ # to have clearer server failures if the AUTH_TOKEN env var isn't present
141
+ if not secrets.compare_digest(token.credentials, os.environ["AUTH_TOKEN"]):
142
+ raise HTTPException(
143
+ status_code=status.HTTP_401_UNAUTHORIZED,
144
+ detail="Incorrect bearer token",
145
+ headers={"WWW-Authenticate": "Bearer"},
146
+ )
147
+ return token
148
+
149
+ @asgi_app.post("/environment_instance")
150
+ async def add_environment_instance(
151
+ env_name: str,
152
+ _: Annotated[HTTPAuthorizationCredentials, Depends(validate_token)],
153
+ _db: Annotated[None, Depends(ensure_session)],
154
+ ) -> uuid.UUID:
155
+ return await backend.add_environment_instance(name=env_name)
156
+
157
+ @asgi_app.post("/environment_frame")
158
+ async def add_environment_frame(
159
+ environment_id: uuid.UUID,
160
+ state: dict[str, Any],
161
+ _: Annotated[HTTPAuthorizationCredentials, Depends(validate_token)],
162
+ _db: Annotated[None, Depends(ensure_session)],
163
+ metadata: dict[str, Any] | None = None,
164
+ ) -> uuid.UUID:
165
+ return await backend.add_environment_frame(
166
+ environment_id=environment_id, state=state, metadata=metadata
167
+ )
168
+
169
+ @asgi_app.get("/environment_instance")
170
+ async def get_environment_instances(
171
+ _: Annotated[HTTPAuthorizationCredentials, Depends(validate_token)],
172
+ _db: Annotated[None, Depends(ensure_session)],
173
+ env_name: str | None = None,
174
+ environment_id: uuid.UUID | None = None,
175
+ ) -> list[EnvironmentDBSchema]:
176
+ if not (env_name or environment_id):
177
+ raise HTTPException(400, "Must provide either a name or environment_id")
178
+ return await backend.get_environment_instances(
179
+ name=env_name, environment_id=environment_id
180
+ )
181
+
182
+ @asgi_app.get("/environment_frame")
183
+ async def get_environment_frames(
184
+ _: Annotated[HTTPAuthorizationCredentials, Depends(validate_token)],
185
+ _db: Annotated[None, Depends(ensure_session)],
186
+ environment_id: uuid.UUID | None = None,
187
+ frame_id: uuid.UUID | None = None,
188
+ ) -> list[FrameDBSchema]:
189
+ if not (frame_id or environment_id):
190
+ raise HTTPException(400, "Must provide either a frame_id or environment_id")
191
+ return await backend.get_environment_frames(
192
+ environment_id=environment_id, frame_id=frame_id
193
+ )
194
+
195
+ return asgi_app
aviary/db.py ADDED
@@ -0,0 +1,155 @@
1
+ from __future__ import annotations
2
+
3
+ import logging
4
+ import uuid
5
+ from datetime import datetime
6
+ from typing import Any, ClassVar
7
+
8
+ from pydantic import BaseModel, ConfigDict, Field, JsonValue
9
+ from sqlalchemy import Column, DateTime, ForeignKey, select
10
+ from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
11
+ from sqlalchemy.orm.decl_api import _TypeAnnotationMapType
12
+ from sqlalchemy.sql import func
13
+ from sqlalchemy.types import JSON
14
+
15
+ from aviary.utils import DBBackend
16
+ from aviary.version import __version__
17
+
18
+ logger = logging.getLogger(__name__)
19
+
20
+
21
+ class EnvironmentBase(DeclarativeBase):
22
+ # Allowing dict to correspond with arbitrary JSON
23
+ type_annotation_map: ClassVar[_TypeAnnotationMapType] = {JsonValue: JSON}
24
+
25
+
26
+ class EnvironmentDBBackend(DBBackend):
27
+ @classmethod
28
+ async def get_environment_instances(
29
+ cls, name: str | None = None, environment_id: uuid.UUID | None = None
30
+ ) -> list[EnvironmentDBSchema]:
31
+ """Given a name, pull associated environment instance data.
32
+
33
+ Args:
34
+ name: The name of the environment this instance is associated with
35
+ environment_id: The UUID of a specific environment instances
36
+
37
+ Returns:
38
+ A list of environment UUIDs.
39
+ """
40
+ async with cls.begin_session() as session:
41
+ if environment_id is not None:
42
+ query = select(EnvironmentDB).where(EnvironmentDB.id == environment_id)
43
+ elif name is not None:
44
+ query = select(EnvironmentDB).where(EnvironmentDB.name == name)
45
+ else:
46
+ raise ValueError("Must provide either a name or environment_id")
47
+
48
+ results = (await session.execute(query)).scalars().all()
49
+ return [EnvironmentDBSchema.model_validate(result) for result in results]
50
+
51
+ @classmethod
52
+ async def get_environment_frames(
53
+ cls, environment_id: uuid.UUID | None, frame_id: uuid.UUID | None
54
+ ) -> list[FrameDBSchema]:
55
+ """Given a environment_id and frame_id data, get environment frame data.
56
+
57
+ Args:
58
+ environment_id: The environment instance we are looking for
59
+ frame_id: The specific frame we are looking for
60
+
61
+ Returns:
62
+ A list of FrameDB.
63
+ """
64
+ async with cls.begin_session() as session:
65
+ if frame_id is not None:
66
+ query = select(FrameDB).where(FrameDB.id == frame_id)
67
+ elif environment_id is not None:
68
+ query = select(FrameDB).where(FrameDB.environment_id == environment_id)
69
+ else:
70
+ raise ValueError("Must provide either a environment_id or frame_id")
71
+ results = (await session.execute(query)).scalars().all()
72
+ return [FrameDBSchema.model_validate(result) for result in results]
73
+
74
+ @classmethod
75
+ async def add_environment_frame(
76
+ cls,
77
+ environment_id: uuid.UUID,
78
+ state: dict[str, Any],
79
+ metadata: dict[str, Any] | None = None,
80
+ ) -> uuid.UUID:
81
+ frame_id = uuid.uuid4()
82
+ async with cls.begin_session() as session:
83
+ session.add(
84
+ FrameDB(
85
+ id=frame_id,
86
+ environment_id=environment_id,
87
+ state=state,
88
+ supplemental_data=metadata,
89
+ )
90
+ )
91
+ return frame_id
92
+
93
+ @classmethod
94
+ async def add_environment_instance(
95
+ cls,
96
+ name: str,
97
+ ) -> uuid.UUID:
98
+ environment_id = uuid.uuid4()
99
+ async with cls.begin_session() as session:
100
+ session.add(
101
+ EnvironmentDB(id=environment_id, name=name, aviary_version=__version__)
102
+ )
103
+ return environment_id
104
+
105
+
106
+ class EnvironmentDB(EnvironmentBase):
107
+ """Rows here correspond to a single environment trajectory, a single named set of frames."""
108
+
109
+ __tablename__ = "environments"
110
+
111
+ id: Mapped[uuid.UUID] = mapped_column(primary_key=True)
112
+ created_at: Mapped[datetime] = mapped_column(default=func.now())
113
+ name: Mapped[str] = mapped_column(index=True)
114
+ aviary_version: Mapped[str] = mapped_column(default=__version__)
115
+
116
+
117
+ class FrameDB(EnvironmentBase):
118
+ """Rows are a snapshot of an environment after each step.
119
+
120
+ Note: corresponds to a frame in the renderer object
121
+
122
+ """
123
+
124
+ __tablename__ = "frames"
125
+
126
+ id: Mapped[uuid.UUID] = mapped_column(primary_key=True)
127
+ environment_id: Mapped[uuid.UUID] = mapped_column(
128
+ ForeignKey(EnvironmentDB.id), index=True
129
+ )
130
+ state: Mapped[JsonValue] # Just for human readability
131
+ supplemental_data: Mapped[JsonValue]
132
+ created_at = Column(DateTime(timezone=True), server_default=func.now())
133
+
134
+
135
+ class EnvironmentDBSchema(BaseModel):
136
+ """Sister model for EnvironmentDB."""
137
+
138
+ id: uuid.UUID
139
+ created_at: datetime = Field(default_factory=datetime.now)
140
+ name: str
141
+ aviary_version: str = __version__
142
+
143
+ model_config = ConfigDict(from_attributes=True)
144
+
145
+
146
+ class FrameDBSchema(BaseModel):
147
+ """Sister model for FrameDB."""
148
+
149
+ id: uuid.UUID
150
+ environment_id: uuid.UUID
151
+ state: JsonValue
152
+ supplemental_data: JsonValue
153
+ created_at: datetime = Field(default_factory=datetime.now)
154
+
155
+ model_config = ConfigDict(from_attributes=True)