272b1eb883
Přidal jsem AUTH_MODE=google|mixed|local: google: jen Google login, lokální/Gitea nejdou ani přes přímý endpoint mixed: Google + lokální/Gitea fallback local: jen lokální login Správa uživatelů zůstává admin-only, guard posledního admina platí pro odebrání role i deaktivaci. Přidal jsem do ní i přehled oprávnění rolí a zobrazení provider/avatar metadat. DB migrace teď přes PRAGMA table_info(users) doplňuje chybějící sloupce bez mazání/recreate existující tabulky; pokud users vůbec neexistuje, bezpečně ji vytvoří. Podporované sloupce zahrnují email, display_name, role, is_enabled, created_at, updated_at, last_login_at, auth_provider, provider_subject, plus kompatibilní sloupce pro lokální login.
292 lines
8.0 KiB
Python
292 lines
8.0 KiB
Python
from typing import Any
|
|
|
|
from fastapi import HTTPException, Request, status
|
|
from passlib.context import CryptContext
|
|
|
|
from app.db.database import get_connection
|
|
from app.db.migrations import run_migrations
|
|
|
|
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
|
|
|
|
|
def _table_columns(con, table_name: str) -> set[str]:
|
|
rows = con.execute(f"PRAGMA table_info({table_name})").fetchall()
|
|
return {row["name"] for row in rows}
|
|
|
|
|
|
def _user_select_columns(con) -> str:
|
|
columns = _table_columns(con, "users")
|
|
enabled_select = "COALESCE(is_enabled, 1) AS is_enabled" if "is_enabled" in columns else "1 AS is_enabled"
|
|
optional_columns = []
|
|
for column in ("last_login_at", "auth_provider", "provider_subject", "avatar_url"):
|
|
if column in columns:
|
|
optional_columns.append(column)
|
|
else:
|
|
optional_columns.append(f"NULL AS {column}")
|
|
return (
|
|
"id, username, display_name, email, password_hash, role, "
|
|
f"is_active, {enabled_select}, created_at, updated_at, "
|
|
+ ", ".join(optional_columns)
|
|
)
|
|
|
|
|
|
def mark_last_login(user_id: int) -> None:
|
|
run_migrations()
|
|
con = get_connection()
|
|
columns = _table_columns(con, "users")
|
|
if "last_login_at" in columns:
|
|
con.execute(
|
|
"""
|
|
UPDATE users
|
|
SET last_login_at = CURRENT_TIMESTAMP,
|
|
updated_at = CURRENT_TIMESTAMP
|
|
WHERE id = ?
|
|
""",
|
|
(user_id,),
|
|
)
|
|
con.commit()
|
|
con.close()
|
|
|
|
|
|
def hash_password(password: str) -> str:
|
|
return pwd_context.hash(password)
|
|
|
|
|
|
def verify_password(plain_password: str, password_hash: str) -> bool:
|
|
if not password_hash:
|
|
return False
|
|
|
|
return pwd_context.verify(plain_password, password_hash)
|
|
|
|
|
|
def get_user_by_username(username: str) -> dict[str, Any] | None:
|
|
run_migrations()
|
|
con = get_connection()
|
|
|
|
row = con.execute(
|
|
f"""
|
|
SELECT {_user_select_columns(con)}
|
|
FROM users
|
|
WHERE username = ?
|
|
""",
|
|
(username,),
|
|
).fetchone()
|
|
|
|
con.close()
|
|
return dict(row) if row else None
|
|
|
|
|
|
def find_or_create_oauth_user(username: str, display_name: str, email: str) -> dict[str, Any]:
|
|
username = username.strip()
|
|
display_name = (display_name or username).strip() or username
|
|
email = (email or "").strip()
|
|
|
|
run_migrations()
|
|
con = get_connection()
|
|
row = con.execute(
|
|
f"""
|
|
SELECT {_user_select_columns(con)}
|
|
FROM users
|
|
WHERE username = ?
|
|
""",
|
|
(username,),
|
|
).fetchone()
|
|
|
|
if row:
|
|
con.execute(
|
|
"""
|
|
UPDATE users
|
|
SET display_name = ?,
|
|
email = ?,
|
|
updated_at = CURRENT_TIMESTAMP
|
|
WHERE id = ?
|
|
""",
|
|
(display_name, email, row["id"]),
|
|
)
|
|
else:
|
|
con.execute(
|
|
"""
|
|
INSERT INTO users (
|
|
username,
|
|
display_name,
|
|
email,
|
|
password_hash,
|
|
role,
|
|
is_active,
|
|
created_at,
|
|
updated_at
|
|
)
|
|
VALUES (?, ?, ?, '', 'developer', 1, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
|
|
""",
|
|
(username, display_name, email),
|
|
)
|
|
|
|
con.commit()
|
|
user_id = row["id"] if row else con.execute("SELECT last_insert_rowid() AS id").fetchone()["id"]
|
|
user = con.execute(
|
|
f"""
|
|
SELECT {_user_select_columns(con)}
|
|
FROM users
|
|
WHERE id = ?
|
|
""",
|
|
(user_id,),
|
|
).fetchone()
|
|
|
|
con.close()
|
|
return dict(user)
|
|
|
|
|
|
def find_or_create_google_user(email: str, display_name: str, provider_subject: str, avatar_url: str = "") -> tuple[dict[str, Any], bool]:
|
|
email = email.strip().lower()
|
|
display_name = (display_name or email).strip() or email
|
|
provider_subject = (provider_subject or "").strip()
|
|
avatar_url = (avatar_url or "").strip()
|
|
|
|
run_migrations()
|
|
con = get_connection()
|
|
columns = _table_columns(con, "users")
|
|
params: list[Any] = [email, email]
|
|
provider_clause = ""
|
|
if provider_subject:
|
|
provider_clause = " OR (auth_provider = 'google' AND provider_subject = ?)"
|
|
params.append(provider_subject)
|
|
|
|
row = con.execute(
|
|
f"""
|
|
SELECT {_user_select_columns(con)}
|
|
FROM users
|
|
WHERE LOWER(COALESCE(email, '')) = ?
|
|
OR LOWER(COALESCE(username, '')) = ?
|
|
{provider_clause}
|
|
ORDER BY
|
|
CASE
|
|
WHEN LOWER(COALESCE(email, '')) = ? THEN 0
|
|
WHEN LOWER(COALESCE(username, '')) = ? THEN 1
|
|
ELSE 2
|
|
END,
|
|
id
|
|
LIMIT 1
|
|
""",
|
|
(*params, email, email),
|
|
).fetchone()
|
|
|
|
if row:
|
|
updates = [
|
|
"display_name = ?",
|
|
"email = ?",
|
|
"auth_provider = 'google'",
|
|
"provider_subject = ?",
|
|
"updated_at = CURRENT_TIMESTAMP",
|
|
]
|
|
values: list[Any] = [display_name, email, provider_subject]
|
|
if "avatar_url" in columns:
|
|
updates.append("avatar_url = ?")
|
|
values.append(avatar_url)
|
|
values.append(row["id"])
|
|
con.execute(f"UPDATE users SET {', '.join(updates)} WHERE id = ?", values)
|
|
created = False
|
|
user_id = row["id"]
|
|
else:
|
|
role = "admin" if con.execute("SELECT COUNT(*) AS count FROM users").fetchone()["count"] == 0 else "viewer"
|
|
insert_columns = [
|
|
"username",
|
|
"display_name",
|
|
"email",
|
|
"password_hash",
|
|
"role",
|
|
"is_active",
|
|
"is_enabled",
|
|
"auth_provider",
|
|
"provider_subject",
|
|
"created_at",
|
|
"updated_at",
|
|
]
|
|
values = [email, display_name, email, "", role, 1, 1, "google", provider_subject]
|
|
placeholders = ["?", "?", "?", "?", "?", "?", "?", "?", "?", "CURRENT_TIMESTAMP", "CURRENT_TIMESTAMP"]
|
|
if "avatar_url" in columns:
|
|
insert_columns.append("avatar_url")
|
|
values.append(avatar_url)
|
|
placeholders.append("?")
|
|
con.execute(
|
|
f"""
|
|
INSERT INTO users ({', '.join(insert_columns)})
|
|
VALUES ({', '.join(placeholders)})
|
|
""",
|
|
values,
|
|
)
|
|
created = True
|
|
user_id = con.execute("SELECT last_insert_rowid() AS id").fetchone()["id"]
|
|
|
|
con.commit()
|
|
user = con.execute(
|
|
f"""
|
|
SELECT {_user_select_columns(con)}
|
|
FROM users
|
|
WHERE id = ?
|
|
""",
|
|
(user_id,),
|
|
).fetchone()
|
|
|
|
con.close()
|
|
return dict(user), created
|
|
|
|
|
|
def get_user_by_id(user_id: int) -> dict[str, Any] | None:
|
|
run_migrations()
|
|
con = get_connection()
|
|
|
|
row = con.execute(
|
|
f"""
|
|
SELECT {_user_select_columns(con)}
|
|
FROM users
|
|
WHERE id = ?
|
|
""",
|
|
(user_id,),
|
|
).fetchone()
|
|
|
|
con.close()
|
|
return dict(row) if row else None
|
|
|
|
|
|
def authenticate_user(username: str, password: str) -> dict[str, Any] | None:
|
|
user = get_user_by_username(username.strip())
|
|
|
|
if not user or not user.get("is_active") or not user.get("is_enabled", 1):
|
|
return None
|
|
|
|
if not verify_password(password, user.get("password_hash") or ""):
|
|
return None
|
|
|
|
mark_last_login(int(user["id"]))
|
|
return user
|
|
|
|
|
|
def current_user(request: Request) -> dict[str, Any] | None:
|
|
user_id = request.session.get("user_id")
|
|
if not user_id:
|
|
return None
|
|
|
|
try:
|
|
user_id = int(user_id)
|
|
except (TypeError, ValueError):
|
|
request.session.pop("user_id", None)
|
|
return None
|
|
|
|
user = get_user_by_id(user_id)
|
|
if not user or not user.get("is_active") or not user.get("is_enabled", 1):
|
|
request.session.pop("user_id", None)
|
|
return None
|
|
|
|
return user
|
|
|
|
|
|
def require_user(request: Request) -> dict[str, Any]:
|
|
user = current_user(request)
|
|
if user:
|
|
return user
|
|
|
|
raise HTTPException(
|
|
status_code=status.HTTP_303_SEE_OTHER,
|
|
headers={"Location": "/portal/login"},
|
|
)
|