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 +195 -0
- aviary/db.py +155 -0
- aviary/env.py +443 -0
- aviary/env_client.py +58 -0
- aviary/message.py +132 -0
- aviary/py.typed +0 -0
- aviary/render.py +130 -0
- aviary/tools/__init__.py +33 -0
- aviary/tools/argref.py +172 -0
- aviary/tools/base.py +401 -0
- aviary/utils.py +98 -0
- aviary/version.py +16 -0
- fhaviary-0.1.0.dist-info/LICENSE +201 -0
- fhaviary-0.1.0.dist-info/METADATA +250 -0
- fhaviary-0.1.0.dist-info/RECORD +17 -0
- fhaviary-0.1.0.dist-info/WHEEL +5 -0
- fhaviary-0.1.0.dist-info/top_level.txt +1 -0
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)
|