Skip to content
Open
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
28 changes: 28 additions & 0 deletions migrations/shard-core-0002-users.sql
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
-- shard-core-0002-users
-- depends: shard-core-0001-init

-- Keep in sync with Role in shard_core/data_model/user.py
CREATE TYPE user_role AS ENUM ('owner', 'member');

CREATE TABLE users (
id BIGSERIAL PRIMARY KEY,
username TEXT UNIQUE NOT NULL,
display_name TEXT NOT NULL,
email TEXT,
role user_role NOT NULL DEFAULT 'member',
password_hash TEXT,
disabled BOOLEAN NOT NULL DEFAULT FALSE,
created TIMESTAMPTZ NOT NULL DEFAULT now()
);

-- Existing shards have their default identity at migration time; create the
-- owner user and bind existing terminals right here so user_id can be
-- NOT NULL from the start. Fresh shards have empty tables — the owner user
-- is created at startup (service.user.ensure_owner_user) before any pairing.
INSERT INTO users (username, display_name, email, role)
SELECT 'owner', COALESCE(name, 'Shard Owner'), email, 'owner'
FROM identities WHERE is_default = TRUE;

ALTER TABLE terminals ADD COLUMN user_id BIGINT REFERENCES users (id);
UPDATE terminals SET user_id = (SELECT id FROM users WHERE role = 'owner');
ALTER TABLE terminals ALTER COLUMN user_id SET NOT NULL;
2 changes: 2 additions & 0 deletions shard_core/app_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
backup,
disk,
telemetry,
user,
)
from .service.app_installation.util import (
write_traefik_dyn_config,
Expand Down Expand Up @@ -79,6 +80,7 @@ def configure_logging():
async def lifespan(_):
await database.init_database()
await identity.init_default_identity()
await user.ensure_owner_user()

await write_traefik_dyn_config()
await render_all_docker_compose_templates()
Expand Down
4 changes: 3 additions & 1 deletion shard_core/data_model/terminal.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,16 +21,18 @@ class Terminal(BaseModel):
name: str
icon: Icon = Icon.UNKNOWN
last_connection: Optional[datetime] = None
user_id: Optional[int] = None

def __str__(self):
return f"Terminal[{self.id}, {self.name}]"

@classmethod
def create(cls, name: str) -> "Terminal":
def create(cls, name: str, user_id: Optional[int] = None) -> "Terminal":
return Terminal(
id=human_encoding.random_string(6),
name=name,
last_connection=datetime.now(timezone.utc),
user_id=user_id,
)


Expand Down
24 changes: 24 additions & 0 deletions shard_core/data_model/user.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
from datetime import datetime
from enum import Enum
from typing import Optional

from pydantic import BaseModel


class Role(str, Enum):
# Keep in sync with the user_role enum in migrations/shard-core-0002-users.sql
OWNER = "owner"
MEMBER = "member"


class User(BaseModel):
id: int
username: str
display_name: str
email: Optional[str] = None
role: Role = Role.MEMBER
disabled: bool = False
created: Optional[datetime] = None

def __str__(self):
return f"User[{self.id}, {self.username}]"
4 changes: 2 additions & 2 deletions shard_core/database/terminals.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,8 @@ async def get_by_name(conn: AsyncConnection, name: str) -> dict | None:


async def insert(conn: AsyncConnection, terminal: dict) -> dict:
sql: LiteralString = """INSERT INTO terminals (id, name, icon, last_connection)
VALUES (%(id)s, %(name)s, %(icon)s, %(last_connection)s)
sql: LiteralString = """INSERT INTO terminals (id, name, icon, last_connection, user_id)
VALUES (%(id)s, %(name)s, %(icon)s, %(last_connection)s, %(user_id)s)
RETURNING *"""
async with conn.cursor(row_factory=dict_row) as cur:
await cur.execute(sql, terminal)
Expand Down
28 changes: 23 additions & 5 deletions shard_core/database/tinydb_migration.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,8 +72,9 @@ async def migrate_tinydb_data():
async with db_conn() as conn:
await _migrate_kv_store(conn, data.get("_default", {}))
await _migrate_identities(conn, data.get("identities", {}))
owner_id = await _ensure_owner_user(conn)
await _migrate_installed_apps(conn, data.get("installed_apps", {}))
await _migrate_terminals(conn, data.get("terminals", {}))
await _migrate_terminals(conn, data.get("terminals", {}), owner_id)
await _migrate_peers(conn, data.get("peers", {}))
await _migrate_backups(conn, data.get("backups", {}))
await _migrate_tours(conn, data.get("tours", {}))
Expand Down Expand Up @@ -121,15 +122,32 @@ async def _migrate_installed_apps(conn: AsyncConnection, records: dict):
log.info(f"migrated {len(records)} installed apps")


async def _migrate_terminals(conn: AsyncConnection, records: dict):
async def _ensure_owner_user(conn: AsyncConnection) -> int | None:
"""Terminals require a user (NOT NULL); create the owner from the just-
migrated default identity, mirroring the 0002 migration's backfill."""
cur = await conn.execute("SELECT id FROM users WHERE role = 'owner'")
row = await cur.fetchone()
if row:
return row[0]
cur = await conn.execute("""INSERT INTO users (username, display_name, email, role)
SELECT 'owner', COALESCE(name, 'Shard Owner'), email, 'owner'
FROM identities WHERE is_default = TRUE
RETURNING id""")
row = await cur.fetchone()
return row[0] if row else None


async def _migrate_terminals(
conn: AsyncConnection, records: dict, owner_id: int | None
):
for record in records.values():
filtered = _filter_keys(record, _TERMINAL_COLUMNS)
filtered.setdefault("icon", "unknown")
await conn.execute(
"""INSERT INTO terminals (id, name, icon, last_connection)
VALUES (%(id)s, %(name)s, %(icon)s, %(last_connection)s)
"""INSERT INTO terminals (id, name, icon, last_connection, user_id)
VALUES (%(id)s, %(name)s, %(icon)s, %(last_connection)s, %(user_id)s)
ON CONFLICT (id) DO NOTHING""",
filtered,
{**filtered, "user_id": owner_id},
)
log.info(f"migrated {len(records)} terminals")

Expand Down
54 changes: 54 additions & 0 deletions shard_core/database/users.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
from typing import LiteralString

from psycopg import AsyncConnection
from psycopg.rows import class_row

from shard_core.data_model.user import User

_UPDATABLE_COLUMNS = {"username", "display_name", "email", "role", "disabled"}


async def get_by_id(conn: AsyncConnection, id: int) -> User | None:
sql: LiteralString = "SELECT * FROM users WHERE id = %s"
async with conn.cursor(row_factory=class_row(User)) as cur:
await cur.execute(sql, (id,))
return await cur.fetchone()


async def get_owner(conn: AsyncConnection) -> User | None:
sql: LiteralString = "SELECT * FROM users WHERE role = 'owner'"
async with conn.cursor(row_factory=class_row(User)) as cur:
await cur.execute(sql)
return await cur.fetchone()


async def insert(conn: AsyncConnection, user: dict) -> User:
sql: LiteralString = """INSERT INTO users (username, display_name, email, role)
VALUES (%(username)s, %(display_name)s, %(email)s, %(role)s)
RETURNING *"""
async with conn.cursor(row_factory=class_row(User)) as cur:
await cur.execute(sql, user)
return await cur.fetchone()


async def update(conn: AsyncConnection, id: int, data: dict) -> User | None:
set_clauses = []
params = {"_id": id}
for key, value in data.items():
if key not in _UPDATABLE_COLUMNS:
raise ValueError(f"Invalid column: {key}")
set_clauses.append(f"{key} = %({key})s")
params[key] = value
if not set_clauses:
return await get_by_id(conn, id)
sql = f"UPDATE users SET {', '.join(set_clauses)} WHERE id = %(_id)s RETURNING *"
async with conn.cursor(row_factory=class_row(User)) as cur:
await cur.execute(sql, params)
return await cur.fetchone()


async def count(conn: AsyncConnection) -> int:
sql: LiteralString = "SELECT COUNT(*) FROM users"
async with conn.cursor() as cur:
await cur.execute(sql)
return (await cur.fetchone())[0]
42 changes: 42 additions & 0 deletions shard_core/service/user.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
import logging

from shard_core.data_model.identity import Identity
from shard_core.data_model.user import Role, User
from shard_core.database.connection import db_conn
from shard_core.database import identities as db_identities
from shard_core.database import users as db_users

log = logging.getLogger(__name__)


async def ensure_owner_user() -> User:
"""Ensure the shard owner exists as a user.

The owner user is created by the 0002 migration on shards that already
have an identity; on fresh shards it is created here, right after the
default identity. Also backfills the email for migration-created owners —
OIDC clients need an email-shaped identifier to auto-provision accounts.
Idempotent — called on every startup, before any pairing can happen.
"""
async with db_conn() as conn:
owner = await db_users.get_owner(conn)
if owner is None:
identity_row = await db_identities.get_default(conn)
identity = Identity(**identity_row)
owner = await db_users.insert(
conn,
{
"username": "owner",
"display_name": identity.name,
"email": identity.email or f"owner@{identity.domain}",
"role": Role.OWNER.value,
},
)
log.info(f"created owner user {owner.id}")
elif owner.email is None:
identity_row = await db_identities.get_default(conn)
identity = Identity(**identity_row)
owner = await db_users.update(
conn, owner.id, {"email": f"owner@{identity.domain}"}
)
return owner
5 changes: 4 additions & 1 deletion shard_core/web/public/pair.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from shard_core.database.connection import db_conn
from shard_core.database import terminals as db_terminals
from shard_core.database import identities as db_identities
from shard_core.database import users as db_users
from shard_core.data_model.identity import Identity
from shard_core.data_model.terminal import Terminal, InputTerminal
from shard_core.service import pairing
Expand Down Expand Up @@ -38,8 +39,10 @@ async def add_terminal(code: str, terminal: InputTerminal, response: Response):
detail="This pairing code is not valid.",
) from e

new_terminal = Terminal.create(terminal.name)
async with db_conn() as conn:
# the owner user always exists here — created at startup, before pairing
owner = await db_users.get_owner(conn)
new_terminal = Terminal.create(terminal.name, user_id=owner.id)
await db_terminals.insert(conn, new_terminal.model_dump())
is_first_terminal = await db_terminals.count(conn) == 1

Expand Down
3 changes: 2 additions & 1 deletion tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -218,9 +218,10 @@ async def app_client(mocker) -> AsyncGenerator[AsyncClient]:
# Initialize the database (migrations + pool) and create default identity
await database.init_database()
try:
from shard_core.service import identity
from shard_core.service import identity, user

await identity.init_default_identity()
await user.ensure_owner_user()

app = app_factory.create_app()

Expand Down
75 changes: 75 additions & 0 deletions tests/test_users.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
from httpx import AsyncClient

from shard_core.data_model.user import Role, User
from shard_core.database.connection import db_conn
from shard_core.database import identities as db_identities
from shard_core.database import terminals as db_terminals
from shard_core.database import users as db_users
from shard_core.service import identity, user
from tests.util import pair_new_terminal


async def test_ensure_owner_user_creates_owner_from_default_identity(db):
default_identity = await identity.init_default_identity()

owner = await user.ensure_owner_user()

assert isinstance(owner, User)
assert isinstance(owner.id, int)
assert owner.role == Role.OWNER
assert owner.username == "owner"
assert owner.display_name == default_identity.name
assert owner.email == f"owner@{default_identity.domain}"
assert owner.disabled is False


async def test_ensure_owner_user_keeps_identity_email(db):
default_identity = await identity.init_default_identity()
async with db_conn() as conn:
await db_identities.update(
conn, default_identity.id, {"email": "max@freeshard.net"}
)

owner = await user.ensure_owner_user()

assert owner.email == "max@freeshard.net"


async def test_ensure_owner_user_is_idempotent(db):
await identity.init_default_identity()

first = await user.ensure_owner_user()
second = await user.ensure_owner_user()

assert first.id == second.id
async with db_conn() as conn:
assert await db_users.count(conn) == 1


async def test_ensure_owner_user_backfills_missing_email(db):
"""Migration-created owners (existing shards) start without a synthesized
email; ensure_owner_user fills it on the next startup."""
default_identity = await identity.init_default_identity()
async with db_conn() as conn:
await db_users.insert(
conn,
{
"username": "owner",
"display_name": default_identity.name,
"email": None,
"role": Role.OWNER.value,
},
) # simulates the 0002 SQL backfill (no synthesized email)

owner = await user.ensure_owner_user()

assert owner.email == f"owner@{default_identity.domain}"


async def test_pairing_binds_terminal_to_owner(app_client: AsyncClient):
await pair_new_terminal(app_client, "T1")

async with db_conn() as conn:
owner = await db_users.get_owner(conn)
terminal = await db_terminals.get_by_name(conn, "T1")
assert terminal["user_id"] == owner.id
Loading