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
1 change: 0 additions & 1 deletion backend/druks/browser/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@
PAYLOAD_WARNING_BYTES = 200 * 1024 * 1024

BROWSER_SESSION_NAME_MAX_LENGTH = 64
BROWSER_SESSION_NAME_PATTERN = r"^[a-z](?:[a-z0-9-]*[a-z0-9])?$"
SITE_MAX_LENGTH = 255

# TTL-only, no renewal: must outlast the longest borrow, bounded by the
Expand Down
4 changes: 2 additions & 2 deletions backend/druks/browser/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,8 +49,8 @@ def __init__(self, name: str) -> None:
class BrowserLoginWindowGoneError(BrowserApiError):
status_code = 410

def __init__(self, session_id: str) -> None:
super().__init__(f"Browser session {session_id!r} has no open login window.")
def __init__(self) -> None:
super().__init__("This login window is no longer open.")


class BrowserVncError(Exception):
Expand Down
48 changes: 24 additions & 24 deletions backend/druks/browser/login.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,15 +32,15 @@ class LoginWindow:
home) between them, and frees itself on the record's TTL if the operator
walks away."""

def __init__(self, session_id: str, host_id: str) -> None:
self.session_id = session_id
def __init__(self, session_name: str, host_id: str) -> None:
self.session_name = session_name
self.host_id = host_id

@classmethod
async def open(cls, session: StoredBrowserSession) -> "LoginWindow":
stale = await get_client().get(_key(session.id))
stale = await get_client().get(_key(session.name))
if stale:
await cls(session.id, json.loads(stale)["host_id"])._close()
await cls(session.name, json.loads(stale)["host_id"])._close()
settings = load_settings()
try:
browser = await sandbox_client.provision(
Expand All @@ -58,20 +58,20 @@ async def open(cls, session: StoredBrowserSession) -> "LoginWindow":
finally:
await browser.aclose()
await get_client().set(
_key(session.id),
_key(session.name),
json.dumps({"host_id": browser.id}),
ex=LOGIN_WINDOW_TTL_SECONDS,
)
return cls(session.id, browser.id)
return cls(session.name, browser.id)

@classmethod
async def get_for_session(cls, session_id: str) -> "LoginWindow":
async def get_for_session(cls, session_name: str) -> "LoginWindow":
"""The window the operator has open for this session; raises once it is
saved, cancelled, or aged out, which is every caller's cue to stop."""
record = await get_client().get(_key(session_id))
record = await get_client().get(_key(session_name))
if record:
return cls(session_id, json.loads(record)["host_id"])
raise exceptions.BrowserLoginWindowGoneError(session_id)
return cls(session_name, json.loads(record)["host_id"])
raise exceptions.BrowserLoginWindowGoneError

async def stream(self, websocket: WebSocket) -> None:
"""Put the operator's canvas in front of the container's screen until
Expand All @@ -98,28 +98,28 @@ async def save(self) -> StoredBrowserSession:
"""Store what the operator logged into as the session's payload, then
tear the window down. A login always captures a profile, so a session
imported as storage_state becomes a profile here."""
session = StoredBrowserSession.get_for_id(self.session_id)
if not session:
raise exceptions.BrowserSessionUnknownError(self.session_id)
try:
async with sandbox_client.attach(host_id=self.host_id) as browser:
payload = await _export(browser, session.name)
session.payload_format = BrowserSessionPayloadFormat.PROFILE_DIR.value
session.store_payload(payload)
return session
finally:
await self._close()
session = StoredBrowserSession.get_for_name(self.session_name)
if session:
try:
async with sandbox_client.attach(host_id=self.host_id) as browser:
payload = await _export(browser, session.name)
session.payload_format = BrowserSessionPayloadFormat.PROFILE_DIR.value
session.store_payload(payload)
return session
finally:
await self._close()
raise exceptions.BrowserSessionUnknownError(self.session_name)

async def cancel(self) -> None:
await self._close()

async def _close(self) -> None:
await sandbox_client.release(host_id=self.host_id)
await get_client().delete(_key(self.session_id))
await get_client().delete(_key(self.session_name))


def _key(session_id: str) -> str:
return f"{LOGIN_WINDOW_KEY_PREFIX}{session_id}"
def _key(session_name: str) -> str:
return f"{LOGIN_WINDOW_KEY_PREFIX}{session_name}"


def is_same_origin(websocket: WebSocket) -> bool:
Expand Down
30 changes: 14 additions & 16 deletions backend/druks/browser/models.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from datetime import datetime

from sqlalchemy import CheckConstraint, String, select
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.orm import Mapped, mapped_column

from druks.browser.constants import (
Expand Down Expand Up @@ -37,25 +38,26 @@ class StoredBrowserSession(Base, Uuid7Pk):
last_used_at: Mapped[datetime | None] = mapped_column(default=None)

@classmethod
def create(
def get_or_create(
cls,
*,
name: str,
payload_format: BrowserSessionPayloadFormat,
site: str,
):
browser_session = cls(
name=name,
payload_format=payload_format.value,
site=site,
"""Concurrency-safe lookup-or-create: two first actions racing on the
same session both INSERT with ON CONFLICT DO NOTHING, then converge on
the one row through the name lookup."""
browser_session = cls.get_for_name(name)
if browser_session:
return browser_session
session = db_session()
session.execute(
insert(cls)
.values(name=name, payload_format=payload_format.value, site=site)
.on_conflict_do_nothing(index_elements=["name"])
)
db_session().add(browser_session)
db_session().flush()
return browser_session

@classmethod
def get_for_id(cls, session_id: str):
return db_session().get(cls, session_id)
return session.scalars(select(cls).where(cls.name == name)).one()

@classmethod
def list_all(cls):
Expand All @@ -65,10 +67,6 @@ def list_all(cls):
def get_for_name(cls, name: str):
return db_session().scalar(select(cls).where(cls.name == name))

def rename(self, name: str) -> None:
self.name = name
db_session().flush()

def mark_stale(self) -> None:
self.status = BrowserSessionStatus.STALE.value
db_session().flush()
Expand Down
Loading