mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-09-24 17:12:20 +02:00
Squash Odysseus development history
This commit is contained in:
+38
-14
@@ -15,29 +15,53 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
|
||||
def atomic_write_json(path: str, data: Any, *, indent: Optional[int] = None) -> None:
|
||||
"""Atomically persist `data` as JSON at `path`.
|
||||
|
||||
The temp file uses the live PID as a suffix so two processes saving the
|
||||
same file (e.g. unit tests) don't collide on the rename target.
|
||||
The temp file uses a random suffix so two concurrent writers saving the
|
||||
same file don't collide on the rename target. A PID suffix does not do
|
||||
this: the PID is constant for the life of a process, so two writers on
|
||||
the same path within one process (or one single-process container, where
|
||||
the PID never changes at all) still race for the same temp file.
|
||||
"""
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
tmp = f"{path}.tmp.{os.getpid()}"
|
||||
with open(tmp, "w") as f:
|
||||
json.dump(data, f, indent=indent)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, path)
|
||||
tmp = f"{path}.tmp.{uuid.uuid4().hex}"
|
||||
|
||||
try:
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=indent)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, path)
|
||||
finally:
|
||||
# Directly unlink to avoid a check-then-act race condition.
|
||||
# Swallows FileNotFoundError (on success path) and other cleanup OSErrors.
|
||||
try:
|
||||
os.unlink(tmp)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def atomic_write_text(path: str, text: str) -> None:
|
||||
if not isinstance(text, str):
|
||||
raise TypeError("atomic_write_text expects a string")
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
tmp = f"{path}.tmp.{os.getpid()}"
|
||||
with open(tmp, "w") as f:
|
||||
f.write(text)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, path)
|
||||
tmp = f"{path}.tmp.{uuid.uuid4().hex}"
|
||||
|
||||
try:
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, path)
|
||||
finally:
|
||||
# Directly unlink to avoid a check-then-act race condition.
|
||||
# Swallows FileNotFoundError (on success path) and other cleanup OSErrors.
|
||||
try:
|
||||
os.unlink(tmp)
|
||||
except OSError:
|
||||
pass
|
||||
+327
-73
@@ -3,6 +3,7 @@ Authentication module — multi-user password hashing, session tokens, config pe
|
||||
Config stored in data/auth.json. Uses bcrypt directly.
|
||||
"""
|
||||
|
||||
import enum
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
@@ -30,16 +31,44 @@ DEFAULT_PRIVILEGES = {
|
||||
"can_manage_memory": True,
|
||||
"max_messages_per_day": 0,
|
||||
"allowed_models": [],
|
||||
"allowed_models_restricted": False,
|
||||
# Explicit "block every model" sentinel. An empty `allowed_models` list is
|
||||
# ambiguous — it's also what gets sent when the admin clicks "[All]" — so
|
||||
# we need a dedicated flag to express "this user may use no models at all"
|
||||
# distinctly from "this user has no restriction".
|
||||
"block_all_models": False,
|
||||
}
|
||||
|
||||
# Admins get everything
|
||||
ADMIN_PRIVILEGES = {k: (True if isinstance(v, bool) else (0 if isinstance(v, int) else [])) for k, v in DEFAULT_PRIVILEGES.items()}
|
||||
ADMIN_PRIVILEGES["allowed_models_restricted"] = False
|
||||
# Admins must never be blocked from using models — the generic dict
|
||||
# comprehension above flips every boolean default to True, which would be
|
||||
# backwards for this sentinel.
|
||||
ADMIN_PRIVILEGES["block_all_models"] = False
|
||||
|
||||
DEFAULT_AUTH_PATH = os.path.join(
|
||||
Path(__file__).parent.parent, "data", "auth.json"
|
||||
)
|
||||
from src.constants import AUTH_FILE, PASSWORD_MIN_LENGTH
|
||||
from src.owner_identity import RESERVED_AUTH_USERNAMES
|
||||
DEFAULT_AUTH_PATH = AUTH_FILE
|
||||
TOKEN_TTL = 60 * 60 * 24 * 7 # 7 days
|
||||
|
||||
# Usernames the auth + middleware layer reserves for request sentinels and
|
||||
# internal storage owners; they must never belong to a real login account.
|
||||
# "internal-tool" is the most dangerous because `core.middleware.require_admin`
|
||||
# treats it as the in-process tool loopback. "api" collides with bearer-token
|
||||
# attribution. "demo"/"system" are synthetic owners already special-cased by
|
||||
# scheduler/assistant/research paths. The Default/Local owner is a storage
|
||||
# bucket for explicit auth-disabled no-login mode, not a login username.
|
||||
RESERVED_USERNAMES = frozenset(RESERVED_AUTH_USERNAMES)
|
||||
|
||||
|
||||
def normalize_known_username(users: Dict[str, Any], username: str | None) -> Optional[str]:
|
||||
"""Return a normalized username only when it exists in the auth user map."""
|
||||
key = str(username or "").strip().lower()
|
||||
if not key or key not in users:
|
||||
return None
|
||||
return key
|
||||
|
||||
|
||||
def _hash_password(password: str) -> str:
|
||||
return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8")
|
||||
@@ -49,6 +78,15 @@ def _verify_password(password: str, hashed: str) -> bool:
|
||||
return bcrypt.checkpw(password.encode("utf-8"), hashed.encode("utf-8"))
|
||||
|
||||
|
||||
class SetAdminResult(enum.Enum):
|
||||
"""Outcome of AuthManager.set_admin, so callers can map each case to a
|
||||
precise response instead of guessing from a bare bool."""
|
||||
OK = "ok"
|
||||
USER_NOT_FOUND = "user_not_found"
|
||||
NOT_AUTHORIZED = "not_authorized" # requester is not an admin
|
||||
LAST_ADMIN = "last_admin" # would remove the last remaining admin
|
||||
|
||||
|
||||
class AuthManager:
|
||||
"""Manages multi-user password + session-token auth system."""
|
||||
|
||||
@@ -60,16 +98,33 @@ class AuthManager:
|
||||
# Guards mutations of self._sessions and the on-disk sessions.json.
|
||||
# Validate/create/revoke run concurrently from the FastAPI threadpool.
|
||||
self._sessions_lock = threading.RLock()
|
||||
# Guards all mutations of self._config and the on-disk auth.json so
|
||||
# concurrent create/delete/rename/privilege operations don't interleave
|
||||
# and corrupt the user database.
|
||||
self._config_lock = threading.Lock()
|
||||
# Guards the first-run setup check-and-write so concurrent requests
|
||||
# cannot both observe is_configured==False and both create admin accounts.
|
||||
self._setup_lock = threading.Lock()
|
||||
self._load()
|
||||
self._load_sessions()
|
||||
self._migrate_single_user()
|
||||
self._drop_reserved_loaded_users()
|
||||
self._migrate_legacy_admin_role()
|
||||
|
||||
def _load(self):
|
||||
try:
|
||||
if os.path.exists(self.auth_path):
|
||||
with open(self.auth_path, "r") as f:
|
||||
with open(self.auth_path, "r", encoding="utf-8") as f:
|
||||
self._config = json.load(f)
|
||||
# Normalize all stored usernames to lowercase so they match
|
||||
# the .strip().lower() applied at login/verify time. Fixes
|
||||
# "Invalid credentials" when auth.json was written with
|
||||
# mixed-case keys (e.g. via manual edit or a future migration).
|
||||
if "users" in self._config:
|
||||
self._config["users"] = {
|
||||
k.strip().lower(): v
|
||||
for k, v in self._config["users"].items()
|
||||
}
|
||||
logger.info("Auth config loaded")
|
||||
else:
|
||||
self._config = {}
|
||||
@@ -82,7 +137,7 @@ class AuthManager:
|
||||
"""Load persisted session tokens from disk, pruning expired ones."""
|
||||
try:
|
||||
if os.path.exists(self._sessions_path):
|
||||
with open(self._sessions_path, "r") as f:
|
||||
with open(self._sessions_path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
now = time.time()
|
||||
self._sessions = {k: v for k, v in data.items() if v.get("expiry", 0) > now}
|
||||
@@ -106,20 +161,52 @@ class AuthManager:
|
||||
def _migrate_single_user(self):
|
||||
"""Migrate old single-user format to multi-user format."""
|
||||
if "password_hash" in self._config and "users" not in self._config:
|
||||
old_user = self._config.get("username", "admin")
|
||||
old_user = str(self._config.get("username", "admin") or "admin").strip().lower()
|
||||
if old_user in RESERVED_USERNAMES:
|
||||
logger.warning(
|
||||
"Migrating legacy single-user reserved username '%s' to 'admin'",
|
||||
old_user,
|
||||
)
|
||||
old_user = "admin"
|
||||
old_hash = self._config["password_hash"]
|
||||
self._config = {
|
||||
"users": {
|
||||
old_user: {
|
||||
"password_hash": old_hash,
|
||||
"created": time.time(),
|
||||
"is_admin": True,
|
||||
with self._config_lock:
|
||||
self._config = {
|
||||
"users": {
|
||||
old_user: {
|
||||
"password_hash": old_hash,
|
||||
"created": time.time(),
|
||||
"is_admin": True,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
self._save()
|
||||
self._save()
|
||||
logger.info(f"Migrated single-user auth to multi-user (admin: {old_user})")
|
||||
|
||||
def _drop_reserved_loaded_users(self):
|
||||
"""Fail closed for legacy/manual auth rows that collide with sentinels."""
|
||||
users = self._config.get("users")
|
||||
if not isinstance(users, dict):
|
||||
return
|
||||
normalized = {}
|
||||
removed = []
|
||||
for username, data in users.items():
|
||||
key = str(username or "").strip().lower()
|
||||
if not key:
|
||||
continue
|
||||
if key in RESERVED_USERNAMES:
|
||||
removed.append(key)
|
||||
continue
|
||||
normalized[key] = data
|
||||
if removed or normalized != users:
|
||||
with self._config_lock:
|
||||
self._config["users"] = normalized
|
||||
self._save()
|
||||
if removed:
|
||||
logger.warning(
|
||||
"Removed reserved username(s) from auth config: %s",
|
||||
", ".join(sorted(set(removed))),
|
||||
)
|
||||
|
||||
def _migrate_legacy_admin_role(self):
|
||||
"""Normalize setup.py's old role='admin' marker to is_admin=True."""
|
||||
changed = False
|
||||
@@ -144,37 +231,54 @@ class AuthManager:
|
||||
|
||||
@signup_enabled.setter
|
||||
def signup_enabled(self, value: bool):
|
||||
self._config["signup_enabled"] = value
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
self._config["signup_enabled"] = value
|
||||
self._save()
|
||||
|
||||
@property
|
||||
def is_configured(self) -> bool:
|
||||
return len(self.users) > 0
|
||||
|
||||
def policy(self) -> dict:
|
||||
"""Return public auth policy constants for the frontend."""
|
||||
return {
|
||||
"password_min_length": PASSWORD_MIN_LENGTH,
|
||||
"reserved_usernames": sorted(RESERVED_USERNAMES),
|
||||
"signup_enabled": self.signup_enabled,
|
||||
"session_days": TOKEN_TTL // 86400,
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Account management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def setup(self, username: str, password: str) -> bool:
|
||||
"""First-run admin setup. Only works if no users exist."""
|
||||
if self.is_configured:
|
||||
return False
|
||||
return self.create_user(username, password, is_admin=True)
|
||||
with self._setup_lock:
|
||||
if self.is_configured:
|
||||
return False
|
||||
return self.create_user(username, password, is_admin=True)
|
||||
|
||||
def create_user(self, username: str, password: str, is_admin: bool = False) -> bool:
|
||||
"""Create a new user account."""
|
||||
username = username.strip().lower()
|
||||
if username in self.users:
|
||||
if not username:
|
||||
return False
|
||||
if "users" not in self._config:
|
||||
self._config["users"] = {}
|
||||
self._config["users"][username] = {
|
||||
"password_hash": _hash_password(password),
|
||||
"created": time.time(),
|
||||
"is_admin": is_admin,
|
||||
"privileges": dict(ADMIN_PRIVILEGES if is_admin else DEFAULT_PRIVILEGES),
|
||||
}
|
||||
self._save()
|
||||
if username in RESERVED_USERNAMES:
|
||||
logger.warning("Refused to create reserved username '%s'", username)
|
||||
return False
|
||||
with self._config_lock:
|
||||
if username in self.users:
|
||||
return False
|
||||
if "users" not in self._config:
|
||||
self._config["users"] = {}
|
||||
self._config["users"][username] = {
|
||||
"password_hash": _hash_password(password),
|
||||
"created": time.time(),
|
||||
"is_admin": is_admin,
|
||||
"privileges": dict(ADMIN_PRIVILEGES if is_admin else DEFAULT_PRIVILEGES),
|
||||
}
|
||||
self._save()
|
||||
logger.info(f"Created user '{username}' (admin={is_admin})")
|
||||
return True
|
||||
|
||||
@@ -187,14 +291,31 @@ class AuthManager:
|
||||
their cookie expired naturally (default ~30 days).
|
||||
"""
|
||||
username = username.strip().lower()
|
||||
if username not in self.users:
|
||||
return False
|
||||
if username == requesting_user:
|
||||
return False
|
||||
if not self.users.get(requesting_user, {}).get("is_admin"):
|
||||
return False
|
||||
del self._config["users"][username]
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
if username not in self.users:
|
||||
return False
|
||||
if username == requesting_user:
|
||||
return False
|
||||
if not self.users.get(requesting_user, {}).get("is_admin"):
|
||||
return False
|
||||
# Revoke API bearer tokens before removing the auth row. The bearer
|
||||
# path authenticates from ApiToken rows and does not require the
|
||||
# owner to still exist, so a successful delete must not leave active
|
||||
# rows behind. If the token store is unavailable, fail closed and
|
||||
# keep the user/session state intact so the admin can retry.
|
||||
try:
|
||||
from core.database import get_db_session, ApiToken
|
||||
with get_db_session() as db:
|
||||
removed_tokens = db.query(ApiToken).filter(ApiToken.owner == username).delete()
|
||||
if removed_tokens:
|
||||
logger.info(
|
||||
f"Revoked {removed_tokens} API token(s) owned by deleted user '{username}'"
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(f"Failed to revoke API tokens for deleted user '{username}'")
|
||||
return False
|
||||
del self._config["users"][username]
|
||||
self._save()
|
||||
# Purge all sessions belonging to this user. validate_token doesn't
|
||||
# cross-check `self.users`, so without this step a deleted user's
|
||||
# cookie keeps authenticating.
|
||||
@@ -210,6 +331,41 @@ class AuthManager:
|
||||
logger.info(f"Deleted user '{username}' (by {requesting_user}); revoked {revoked} active session(s)")
|
||||
return True
|
||||
|
||||
def rename_user(self, old_username: str, new_username: str, requesting_user: str) -> bool:
|
||||
"""Rename a user in auth config and active sessions. Admin only."""
|
||||
old_username = old_username.strip().lower()
|
||||
new_username = new_username.strip().lower()
|
||||
requesting_user = (requesting_user or "").strip().lower()
|
||||
if not old_username or not new_username:
|
||||
return False
|
||||
if new_username in RESERVED_USERNAMES:
|
||||
logger.warning("Refused to rename '%s' into reserved username '%s'", old_username, new_username)
|
||||
return False
|
||||
with self._config_lock:
|
||||
if old_username not in self.users:
|
||||
return False
|
||||
if new_username in self.users:
|
||||
return False
|
||||
if not self.users.get(requesting_user, {}).get("is_admin"):
|
||||
return False
|
||||
self._config.setdefault("users", {})[new_username] = self._config["users"].pop(old_username)
|
||||
self._save()
|
||||
|
||||
renamed_sessions = 0
|
||||
with self._sessions_lock:
|
||||
for sess in self._sessions.values():
|
||||
sess_user = str((sess or {}).get("username") or "").strip().lower()
|
||||
if sess_user == old_username:
|
||||
sess["username"] = new_username
|
||||
renamed_sessions += 1
|
||||
if renamed_sessions:
|
||||
self._save_sessions()
|
||||
logger.info(
|
||||
"Renamed user '%s' -> '%s' (by %s); updated %d active session(s)",
|
||||
old_username, new_username, requesting_user, renamed_sessions,
|
||||
)
|
||||
return True
|
||||
|
||||
def is_admin(self, username: str) -> bool:
|
||||
return self.users.get(username, {}).get("is_admin", False)
|
||||
|
||||
@@ -231,28 +387,93 @@ class AuthManager:
|
||||
def set_privileges(self, username: str, privileges: Dict[str, Any]) -> bool:
|
||||
"""Update privileges for a user. Can't modify admin privileges."""
|
||||
username = username.strip().lower()
|
||||
if username not in self.users:
|
||||
return False
|
||||
if self.users[username].get("is_admin"):
|
||||
return False # admins always have full access
|
||||
# Only allow known privilege keys
|
||||
current = self.get_privileges(username)
|
||||
for k, v in privileges.items():
|
||||
if k in DEFAULT_PRIVILEGES:
|
||||
current[k] = v
|
||||
self._config["users"][username]["privileges"] = current
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
if username not in self.users:
|
||||
return False
|
||||
if self.users[username].get("is_admin"):
|
||||
return False # admins always have full access
|
||||
# Only allow known privilege keys
|
||||
current = self.get_privileges(username)
|
||||
for k, v in privileges.items():
|
||||
if k in DEFAULT_PRIVILEGES:
|
||||
current[k] = v
|
||||
self._config["users"][username]["privileges"] = current
|
||||
self._save()
|
||||
logger.info(f"Updated privileges for '{username}': {current}")
|
||||
return True
|
||||
|
||||
def set_admin(self, username: str, is_admin: bool,
|
||||
requesting_user: str) -> SetAdminResult:
|
||||
"""Promote/demote an existing user to/from admin. Admin only.
|
||||
|
||||
Refuses to remove the last remaining admin so the instance can never
|
||||
be locked out of admin access; self-demotion is allowed as long as
|
||||
another admin remains. Admin status is re-checked live on every
|
||||
request, so unlike delete/rename no session or token revocation is
|
||||
needed — a demoted admin simply fails the next is_admin() gate.
|
||||
|
||||
Promotion stashes the user's current privilege map and demotion
|
||||
restores it, so a temporary admin stint can't silently broaden a
|
||||
user's non-admin access; users without a stash (created as admin,
|
||||
or promoted before stashing existed) demote to DEFAULT_PRIVILEGES.
|
||||
|
||||
Counting admins and flipping the flag happen in one critical section
|
||||
so two concurrent demotions can't race the admin count to zero.
|
||||
"""
|
||||
username = (username or "").strip().lower()
|
||||
requesting_user = (requesting_user or "").strip().lower()
|
||||
is_admin = bool(is_admin)
|
||||
with self._config_lock:
|
||||
target = self._config.get("users", {}).get(username)
|
||||
if target is None:
|
||||
return SetAdminResult.USER_NOT_FOUND
|
||||
if not self.users.get(requesting_user, {}).get("is_admin"):
|
||||
return SetAdminResult.NOT_AUTHORIZED
|
||||
currently_admin = bool(target.get("is_admin"))
|
||||
if currently_admin == is_admin:
|
||||
return SetAdminResult.OK # no-op; leave privileges untouched
|
||||
if currently_admin and not is_admin:
|
||||
admin_count = sum(1 for d in self.users.values() if d.get("is_admin"))
|
||||
if admin_count <= 1:
|
||||
return SetAdminResult.LAST_ADMIN
|
||||
# Write order matters for lock-free readers: get_privileges()
|
||||
# reads without _config_lock and trusts is_admin, so the admin
|
||||
# flag must be flipped while the stored map is safe to expose —
|
||||
# before writing admin privileges on promote, after restoring
|
||||
# the pre-admin map on demote.
|
||||
if is_admin:
|
||||
target["is_admin"] = True
|
||||
# Stash the pre-admin map so a later demotion can restore it.
|
||||
# While is_admin is set the stored map is inert: get_privileges
|
||||
# short-circuits to ADMIN_PRIVILEGES and set_privileges refuses
|
||||
# admins, so only set_admin ever touches the stash.
|
||||
target["privileges_before_admin"] = dict(
|
||||
target.get("privileges") or DEFAULT_PRIVILEGES
|
||||
)
|
||||
target["privileges"] = dict(ADMIN_PRIVILEGES)
|
||||
else:
|
||||
# Restore the stashed pre-admin map. Fall back to defaults for
|
||||
# users created as admins (their stored map is ADMIN_PRIVILEGES,
|
||||
# which must not leak past demotion — e.g. can_use_bash) and
|
||||
# for admins promoted before the stash existed.
|
||||
target["privileges"] = dict(
|
||||
target.pop("privileges_before_admin", None)
|
||||
or DEFAULT_PRIVILEGES
|
||||
)
|
||||
target["is_admin"] = False
|
||||
self._save()
|
||||
logger.info("Set is_admin=%s for '%s' (by '%s')", is_admin, username, requesting_user)
|
||||
return SetAdminResult.OK
|
||||
|
||||
def change_password(self, username: str, current_password: str, new_password: str) -> bool:
|
||||
username = username.strip().lower()
|
||||
if username not in self.users:
|
||||
return False
|
||||
if not _verify_password(current_password, self.users[username]["password_hash"]):
|
||||
return False
|
||||
self._config["users"][username]["password_hash"] = _hash_password(new_password)
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
self._config["users"][username]["password_hash"] = _hash_password(new_password)
|
||||
self._save()
|
||||
return True
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -270,8 +491,9 @@ class AuthManager:
|
||||
if username not in self.users:
|
||||
return None
|
||||
secret = pyotp.random_base32()
|
||||
self._config["users"][username]["totp_secret_pending"] = secret
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
self._config["users"][username]["totp_secret_pending"] = secret
|
||||
self._save()
|
||||
return secret
|
||||
|
||||
def totp_get_provisioning_uri(self, username: str, secret: str) -> str:
|
||||
@@ -290,13 +512,14 @@ class AuthManager:
|
||||
if not totp.verify(code, valid_window=1):
|
||||
return False
|
||||
# Enable 2FA
|
||||
self._config["users"][username]["totp_secret"] = secret
|
||||
self._config["users"][username]["totp_enabled"] = True
|
||||
self._config["users"][username].pop("totp_secret_pending", None)
|
||||
# Generate backup codes
|
||||
backup = [secrets.token_hex(4) for _ in range(8)]
|
||||
self._config["users"][username]["totp_backup_codes"] = backup
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
self._config["users"][username]["totp_secret"] = secret
|
||||
self._config["users"][username]["totp_enabled"] = True
|
||||
self._config["users"][username].pop("totp_secret_pending", None)
|
||||
# Generate backup codes
|
||||
backup = [secrets.token_hex(4) for _ in range(8)]
|
||||
self._config["users"][username]["totp_backup_codes"] = backup
|
||||
self._save()
|
||||
logger.info(f"2FA enabled for '{username}'")
|
||||
return True
|
||||
|
||||
@@ -308,13 +531,17 @@ class AuthManager:
|
||||
return True # 2FA not enabled, always pass
|
||||
secret = user.get("totp_secret")
|
||||
if not secret:
|
||||
return True
|
||||
# 2FA is enabled but no secret is stored (corrupt/partially-written
|
||||
# auth.json). Fail closed — returning True here bypassed the second
|
||||
# factor entirely.
|
||||
return False
|
||||
# Check backup codes first
|
||||
backup = user.get("totp_backup_codes", [])
|
||||
if code in backup:
|
||||
backup.remove(code)
|
||||
self._config["users"][username]["totp_backup_codes"] = backup
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
backup.remove(code)
|
||||
self._config["users"][username]["totp_backup_codes"] = backup
|
||||
self._save()
|
||||
logger.info(f"Backup code used for '{username}' ({len(backup)} remaining)")
|
||||
return True
|
||||
totp = pyotp.TOTP(secret)
|
||||
@@ -325,11 +552,12 @@ class AuthManager:
|
||||
username = username.strip().lower()
|
||||
if not self.verify_password(username, password):
|
||||
return False
|
||||
self._config["users"][username].pop("totp_secret", None)
|
||||
self._config["users"][username].pop("totp_secret_pending", None)
|
||||
self._config["users"][username].pop("totp_backup_codes", None)
|
||||
self._config["users"][username]["totp_enabled"] = False
|
||||
self._save()
|
||||
with self._config_lock:
|
||||
self._config["users"][username].pop("totp_secret", None)
|
||||
self._config["users"][username].pop("totp_secret_pending", None)
|
||||
self._config["users"][username].pop("totp_backup_codes", None)
|
||||
self._config["users"][username]["totp_enabled"] = False
|
||||
self._save()
|
||||
logger.info(f"2FA disabled for '{username}'")
|
||||
return True
|
||||
|
||||
@@ -348,12 +576,22 @@ class AuthManager:
|
||||
username = username.strip().lower()
|
||||
if not self.verify_password(username, password):
|
||||
return None
|
||||
return self.create_session_trusted(username)
|
||||
|
||||
def create_session_trusted(self, username: str) -> Optional[str]:
|
||||
"""Issue a session token for an already-verified user.
|
||||
Call only after verify_password (and TOTP if enabled) have passed."""
|
||||
username = username.strip().lower()
|
||||
token = secrets.token_hex(32)
|
||||
with self._sessions_lock:
|
||||
self._sessions[token] = {
|
||||
"username": username,
|
||||
"expiry": time.time() + TOKEN_TTL,
|
||||
}
|
||||
with self._config_lock:
|
||||
if username not in self.users:
|
||||
logger.warning("Refused to issue session for missing user '%s'", username)
|
||||
return None
|
||||
with self._sessions_lock:
|
||||
self._sessions[token] = {
|
||||
"username": username,
|
||||
"expiry": time.time() + TOKEN_TTL,
|
||||
}
|
||||
self._save_sessions()
|
||||
return token
|
||||
|
||||
@@ -412,6 +650,22 @@ class AuthManager:
|
||||
self._sessions.pop(token, None)
|
||||
self._save_sessions()
|
||||
|
||||
def revoke_user_sessions(self, username: str, except_token: Optional[str] = None) -> int:
|
||||
"""Revoke active browser sessions for a user, optionally preserving one."""
|
||||
username = username.strip().lower()
|
||||
revoked = 0
|
||||
with self._sessions_lock:
|
||||
to_drop = [
|
||||
token for token, session in self._sessions.items()
|
||||
if token != except_token and (session or {}).get("username") == username
|
||||
]
|
||||
for token in to_drop:
|
||||
self._sessions.pop(token, None)
|
||||
revoked += 1
|
||||
if revoked:
|
||||
self._save_sessions()
|
||||
return revoked
|
||||
|
||||
def status(self, token: Optional[str]) -> Dict[str, Any]:
|
||||
username = self.get_username_for_token(token)
|
||||
authenticated = username is not None
|
||||
|
||||
+11
-39
@@ -1,40 +1,12 @@
|
||||
# src/constants.py
|
||||
"""Application-wide constants and configuration values."""
|
||||
import os
|
||||
# core/constants.py
|
||||
"""Backward-compatible shim — the single source of truth is src/constants.py.
|
||||
|
||||
APP_VERSION = "0.9.1"
|
||||
|
||||
# Base paths
|
||||
BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + "/"
|
||||
STATIC_DIR = os.path.join(BASE_DIR, "static")
|
||||
DATA_DIR = os.path.join(BASE_DIR, "data")
|
||||
|
||||
# Data file paths
|
||||
SESSIONS_FILE = os.path.join(DATA_DIR, "sessions.json")
|
||||
MEMORY_FILE = os.path.join(DATA_DIR, "memory.json")
|
||||
MEMORY_DOC = os.path.join(DATA_DIR, "memory_doc.md")
|
||||
PERSONAL_DIR = os.path.join(DATA_DIR, "personal_docs")
|
||||
RUNBOOK_DIR = os.path.join(PERSONAL_DIR, "runbook")
|
||||
UPLOAD_DIR = os.path.join(DATA_DIR, "uploads")
|
||||
FEATURES_FILE = os.path.join(DATA_DIR, "features.json")
|
||||
SETTINGS_FILE = os.path.join(DATA_DIR, "settings.json")
|
||||
|
||||
# API Configuration
|
||||
MAX_CONTEXT_MESSAGES = 90
|
||||
REQUEST_TIMEOUT = 20
|
||||
OPENAI_COMPAT_PATH = "/v1/chat/completions"
|
||||
|
||||
# Environment variables with defaults
|
||||
DEFAULT_HOST = os.getenv("LLM_HOST", "localhost")
|
||||
LLM_HOSTS = [h.strip() for h in os.getenv("LLM_HOSTS", "").split(",") if h.strip()]
|
||||
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
|
||||
SEARXNG_INSTANCE = os.getenv('SEARXNG_INSTANCE', 'http://localhost:8080')
|
||||
|
||||
|
||||
# Cleanup configuration
|
||||
CLEANUP_ENABLED = os.getenv("CLEANUP_ENABLED", "True").lower() == "true"
|
||||
CLEANUP_INTERVAL_HOURS = int(os.getenv("CLEANUP_INTERVAL_HOURS", "24"))
|
||||
|
||||
# Default parameters
|
||||
DEFAULT_TEMPERATURE = 1.0
|
||||
DEFAULT_MAX_TOKENS = 0
|
||||
Historically there were two copies of this module (this one lagged behind at
|
||||
APP_VERSION 0.9.1 and was missing the consolidated tool-output constants). To
|
||||
kill the drift, this now simply re-exports everything from src.constants so
|
||||
there is exactly one place that defines paths and reads ODYSSEUS_DATA_DIR.
|
||||
internal_api_base() also lives in src.constants now and is re-exported here so
|
||||
existing `from core.constants import internal_api_base` callers keep working.
|
||||
"""
|
||||
from src.constants import * # noqa: F401,F403
|
||||
from src.constants import internal_api_base # noqa: F401 (explicit: functions aren't covered by some linters' * checks)
|
||||
|
||||
+1302
-107
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -1,4 +1,4 @@
|
||||
# src/exceptions.py
|
||||
# core/exceptions.py
|
||||
"""Custom exceptions for the application."""
|
||||
|
||||
class SessionNotFoundError(Exception):
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Helpers for keeping sensitive data out of logs.
|
||||
|
||||
Endpoint URLs configured by admins can embed credentials in the userinfo
|
||||
(``https://user:pass@host``) or query string (``?api_key=...``). Logging them
|
||||
raw leaks those secrets, so route/diagnostic logs run URLs through
|
||||
``redact_url`` first. Reconstructing the URL without userinfo/query/fragment
|
||||
also doubles as a sanitizer barrier for CodeQL's clear-text-logging query.
|
||||
"""
|
||||
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
|
||||
def redact_url(url: str) -> str:
|
||||
"""Return a URL safe for logs by removing userinfo and query/fragment.
|
||||
|
||||
Keeps scheme, host, port and path so logs stay useful for debugging.
|
||||
"""
|
||||
try:
|
||||
parsed = urlparse(url or "")
|
||||
host = parsed.hostname or ""
|
||||
if ":" in host: # IPv6 literal — re-bracket so host:port stays unambiguous
|
||||
host = f"[{host}]"
|
||||
if parsed.port:
|
||||
host = f"{host}:{parsed.port}"
|
||||
return urlunparse((parsed.scheme, host, parsed.path, "", "", ""))
|
||||
except Exception:
|
||||
return "<endpoint>"
|
||||
+60
-8
@@ -3,10 +3,14 @@
|
||||
|
||||
import os
|
||||
import secrets
|
||||
from collections.abc import Mapping
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.responses import Response
|
||||
from starlette.routing import get_route_path
|
||||
|
||||
from src.owner_identity import INTERNAL_TOOL_USER, auth_disabled
|
||||
|
||||
|
||||
# Per-process token that lets the in-app tool layer hit admin-gated
|
||||
@@ -17,6 +21,39 @@ INTERNAL_TOOL_TOKEN = os.environ.get("ODYSSEUS_INTERNAL_TOKEN") or secrets.token
|
||||
INTERNAL_TOOL_HEADER = "X-Odysseus-Internal-Token"
|
||||
|
||||
|
||||
def get_application_route_path(scope: Mapping[str, object]) -> str:
|
||||
"""Return the application-relative path used by Starlette routing.
|
||||
|
||||
Uvicorn prefixes ``scope["path"]`` with a configured ASGI ``root_path``;
|
||||
Starlette removes that prefix before matching routes. Middleware policy
|
||||
must use the same path form or a deployment prefix can change which policy
|
||||
applies to an otherwise unchanged application route.
|
||||
"""
|
||||
return get_route_path(scope)
|
||||
|
||||
|
||||
def with_asgi_root_path(scope: Mapping[str, object], path: str) -> str:
|
||||
"""Prefix an application path for a client-facing redirect target."""
|
||||
root_path = scope.get("root_path", "")
|
||||
if not isinstance(root_path, str) or not root_path:
|
||||
return path
|
||||
return f"{root_path.rstrip('/')}{path}"
|
||||
|
||||
|
||||
def path_is_route_or_child(path: str, prefix: str) -> bool:
|
||||
"""Return whether ``path`` is exactly ``prefix`` or below that route."""
|
||||
return path == prefix or path.startswith(prefix + "/")
|
||||
|
||||
|
||||
def is_cors_preflight(method: str, headers) -> bool:
|
||||
"""True for a genuine CORS preflight: an OPTIONS request carrying the
|
||||
Access-Control-Request-Method header. Such requests are credential-less by
|
||||
design and must reach CORSMiddleware to be answered -- gating them on auth
|
||||
401s the preflight and breaks every cross-origin browser/WebView client.
|
||||
Pure so it can be unit-tested without standing up the app."""
|
||||
return method == "OPTIONS" and "access-control-request-method" in headers
|
||||
|
||||
|
||||
def require_admin(request: Request):
|
||||
"""Raise 403 if the current user isn't an admin.
|
||||
Allows access when auth is explicitly disabled, or when the request carries
|
||||
@@ -27,15 +64,16 @@ def require_admin(request: Request):
|
||||
# (b) the auth middleware already validated the token and stamped
|
||||
# request.state.current_user = "internal-tool".
|
||||
try:
|
||||
if request.headers.get(INTERNAL_TOOL_HEADER) == INTERNAL_TOOL_TOKEN:
|
||||
hdr = request.headers.get(INTERNAL_TOOL_HEADER)
|
||||
if hdr and secrets.compare_digest(hdr, INTERNAL_TOOL_TOKEN):
|
||||
return
|
||||
if getattr(request.state, "current_user", None) == "internal-tool":
|
||||
if getattr(request.state, "current_user", None) == INTERNAL_TOOL_USER:
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
auth_mgr = getattr(request.app.state, "auth_manager", None)
|
||||
if os.getenv("AUTH_ENABLED", "true").lower() == "false":
|
||||
if auth_disabled():
|
||||
return
|
||||
if not auth_mgr or not auth_mgr.is_configured:
|
||||
raise HTTPException(403, "Admin only")
|
||||
@@ -55,13 +93,23 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
||||
response = await call_next(request)
|
||||
path = request.url.path
|
||||
|
||||
# Tool render endpoints are served inside iframes — allow framing by self
|
||||
# Tool render endpoints
|
||||
is_tool_render = path.startswith("/api/tools/") and path.endswith("/render")
|
||||
# Document library PDF preview endpoint
|
||||
is_document_pdf_preview = path.startswith("/api/document/") and path.endswith("/render-pdf")
|
||||
# Visual report pages are self-contained HTML — need inline scripts + external images
|
||||
is_report = path.startswith("/api/research/report/")
|
||||
|
||||
response.headers["X-Content-Type-Options"] = "nosniff"
|
||||
response.headers["Referrer-Policy"] = "no-referrer"
|
||||
response.headers["Permissions-Policy"] = "camera=(), microphone=(self), geolocation=()"
|
||||
|
||||
is_https = (
|
||||
request.url.scheme == "https"
|
||||
or request.headers.get("X-Forwarded-Proto") == "https"
|
||||
)
|
||||
if is_https:
|
||||
response.headers["Strict-Transport-Security"] = "max-age=31536000; includeSubDomains"
|
||||
|
||||
if is_report:
|
||||
response.headers["Content-Security-Policy"] = (
|
||||
@@ -74,10 +122,14 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
||||
"frame-ancestors 'none'"
|
||||
)
|
||||
elif is_tool_render:
|
||||
# Tool iframe content: skip all framing headers — the iframe's
|
||||
# sandbox="allow-scripts" attribute provides isolation.
|
||||
# Don't overwrite the route's own restrictive CSP either.
|
||||
# Skip framing headers for tools.
|
||||
pass
|
||||
elif is_document_pdf_preview:
|
||||
response.headers["X-Frame-Options"] = "SAMEORIGIN"
|
||||
response.headers["Content-Security-Policy"] = (
|
||||
"default-src 'none'; "
|
||||
"frame-ancestors 'self'"
|
||||
)
|
||||
else:
|
||||
response.headers["X-Frame-Options"] = "DENY"
|
||||
# NOTE: `style-src 'unsafe-inline'` is intentionally retained.
|
||||
@@ -91,7 +143,7 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
||||
f"script-src 'self' 'nonce-{nonce}' https://cdn.jsdelivr.net; "
|
||||
"style-src 'self' 'unsafe-inline' https://cdn.jsdelivr.net; "
|
||||
"font-src 'self' https://cdn.jsdelivr.net; "
|
||||
"img-src 'self' data: blob:; "
|
||||
"img-src 'self' data: blob: https:; "
|
||||
"media-src 'self' blob:; "
|
||||
"connect-src 'self'; "
|
||||
"frame-src 'self'; "
|
||||
|
||||
+136
-15
@@ -8,17 +8,61 @@ These are simple datacontainers. All persistence is handled by SessionManager.
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Any, Optional, TYPE_CHECKING
|
||||
|
||||
from src.tool_approval_scopes import (
|
||||
CHAT_SESSION_APPROVAL_CONTEXT_MARKER,
|
||||
CHAT_SESSION_APPROVAL_DECISION,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .session_manager import SessionManager
|
||||
|
||||
# Module-level session manager reference (set at app startup)
|
||||
_session_manager: Optional["SessionManager"] = None
|
||||
# Module-level session manager singleton (single source of truth)
|
||||
_SESSION_MANAGER_INSTANCE: Optional["SessionManager"] = None
|
||||
|
||||
|
||||
def set_session_manager(manager: "SessionManager"):
|
||||
"""Set the global session manager reference."""
|
||||
global _session_manager
|
||||
_session_manager = manager
|
||||
def set_session_manager_instance(manager: "SessionManager"):
|
||||
"""Set the global SessionManager singleton."""
|
||||
global _SESSION_MANAGER_INSTANCE
|
||||
_SESSION_MANAGER_INSTANCE = manager
|
||||
|
||||
|
||||
def get_session_manager_instance() -> Optional["SessionManager"]:
|
||||
"""Get the global SessionManager singleton."""
|
||||
return _SESSION_MANAGER_INSTANCE
|
||||
|
||||
|
||||
# Keep legacy name for backward compatibility
|
||||
set_session_manager = set_session_manager_instance
|
||||
get_session_manager = get_session_manager_instance
|
||||
|
||||
|
||||
def _history_grants_chat_session_approval(
|
||||
history: List["ChatMessage"],
|
||||
session_id: str,
|
||||
) -> bool:
|
||||
"""Return whether this exact chat has a resolved session-scope grant."""
|
||||
|
||||
expected_session = str(session_id or "")
|
||||
if not expected_session:
|
||||
return False
|
||||
for message in reversed(history or []):
|
||||
metadata = getattr(message, "metadata", None)
|
||||
if not isinstance(metadata, dict):
|
||||
continue
|
||||
tool_events = metadata.get("tool_events")
|
||||
if not isinstance(tool_events, list):
|
||||
continue
|
||||
for event in reversed(tool_events):
|
||||
ask_user = event.get("ask_user") if isinstance(event, dict) else None
|
||||
if not isinstance(ask_user, dict):
|
||||
continue
|
||||
if (
|
||||
ask_user.get("kind") == "tool_approval"
|
||||
and ask_user.get("resolved") == CHAT_SESSION_APPROVAL_DECISION
|
||||
and str(ask_user.get("session_id") or "") == expected_session
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -42,7 +86,17 @@ class ChatMessage:
|
||||
|
||||
@dataclass
|
||||
class Session:
|
||||
"""A chat session — pure data container."""
|
||||
"""A chat session — pure data container.
|
||||
|
||||
``.history`` is the authoritative mutable message list. Callers may
|
||||
read, append, pop, or reassign it directly — these changes take
|
||||
effect immediately. ``_history`` remains a compatibility alias that
|
||||
always resolves to the authoritative ``history`` list.
|
||||
|
||||
Each session gets its own unique history list at construction time
|
||||
(the dataclass default is never shared between instances).
|
||||
"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
endpoint_url: str
|
||||
@@ -54,31 +108,98 @@ class Session:
|
||||
owner: Optional[str] = None
|
||||
is_important: bool = False
|
||||
message_count: int = 0
|
||||
memory_extraction_enabled: bool = True
|
||||
skill_injection_enabled: bool = True
|
||||
thinking_mode: str = "off"
|
||||
temperature_override: Optional[float] = None
|
||||
max_tokens_override: Optional[int] = None
|
||||
cwd: Optional[str] = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.history is None:
|
||||
self.history = []
|
||||
if self.headers is None:
|
||||
self.headers = {}
|
||||
# Ensure each session gets its OWN list (not the shared dataclass default)
|
||||
if self.history is None:
|
||||
self.history = []
|
||||
|
||||
@property
|
||||
def _history(self) -> List[ChatMessage]:
|
||||
"""Compatibility alias for callers that still reference ``_history``."""
|
||||
return self.history
|
||||
|
||||
@_history.setter
|
||||
def _history(self, messages: List[ChatMessage]):
|
||||
self.history = messages
|
||||
|
||||
def add_message(self, message: ChatMessage):
|
||||
"""
|
||||
Add a message to this session.
|
||||
|
||||
Delegates to SessionManager for persistence if available,
|
||||
otherwise just appends to history.
|
||||
Appends to the authoritative history list and increments
|
||||
message_count. Delegates to SessionManager for persistence
|
||||
if available.
|
||||
"""
|
||||
self.history.append(message)
|
||||
self.message_count = len(self.history)
|
||||
|
||||
# Delegate to session manager for persistence
|
||||
if _session_manager:
|
||||
_session_manager._persist_message(self.id, message)
|
||||
if _SESSION_MANAGER_INSTANCE:
|
||||
_SESSION_MANAGER_INSTANCE._persist_message(self.id, message)
|
||||
|
||||
def get_context_messages(self) -> List[Dict[str, Any]]:
|
||||
"""Get messages in format for LLM API."""
|
||||
return [msg.to_dict() for msg in self.history]
|
||||
"""Get messages in format for LLM API.
|
||||
|
||||
Slash-command / setup replies are persisted to history so they render
|
||||
in the transcript, but they are UI chatter (e.g. ``/setup ...`` and its
|
||||
status lines) the user never meant as conversation. They carry
|
||||
``metadata.source == "slash"``; exclude them here so they never reach
|
||||
the model. Display/history-load paths use the raw ``history`` and are
|
||||
unaffected.
|
||||
"""
|
||||
messages = [
|
||||
msg.to_dict()
|
||||
for msg in self.history
|
||||
if (msg.metadata or {}).get("source") != "slash"
|
||||
]
|
||||
from src.background_tool_jobs import background_result_context
|
||||
messages = [part for message in messages for part in (
|
||||
*background_result_context(message.get('metadata')), message,
|
||||
)]
|
||||
# Resume an interrupted thinking-only response from its actual model
|
||||
# reasoning channel. Restrict this to the latest assistant message so
|
||||
# old traces do not accumulate in context or cause reasoning loops.
|
||||
for index in range(len(messages) - 1, -1, -1):
|
||||
message = messages[index]
|
||||
if message.get("role") != "assistant":
|
||||
continue
|
||||
metadata = message.get("metadata") or {}
|
||||
thinking = str(metadata.get("thinking") or "").strip()
|
||||
if metadata.get("stopped") and thinking:
|
||||
resumed = dict(message)
|
||||
resumed["reasoning_content"] = thinking
|
||||
messages[index] = resumed
|
||||
break
|
||||
if not _history_grants_chat_session_approval(self.history, self.id):
|
||||
return messages
|
||||
|
||||
# Keep the grant close to the latest user request so route-neutral
|
||||
# compaction/trimming preserves it. Copy the metadata instead of
|
||||
# mutating the durable transcript object.
|
||||
for index in range(len(messages) - 1, -1, -1):
|
||||
if messages[index].get("role") != "user":
|
||||
continue
|
||||
message = dict(messages[index])
|
||||
metadata = dict(message.get("metadata") or {})
|
||||
metadata[CHAT_SESSION_APPROVAL_CONTEXT_MARKER] = True
|
||||
message["metadata"] = metadata
|
||||
messages[index] = message
|
||||
break
|
||||
return messages
|
||||
|
||||
def get(self, key: str, default=None):
|
||||
"""Dict-like access for compatibility."""
|
||||
return getattr(self, key, default)
|
||||
|
||||
def __getitem__(self, key: str):
|
||||
"""Allow session['field'] syntax."""
|
||||
return getattr(self, key)
|
||||
|
||||
@@ -0,0 +1,452 @@
|
||||
"""Cross-platform OS compatibility helpers.
|
||||
|
||||
Odysseus began as a Linux/macOS/Docker-only app. This module centralizes the
|
||||
small set of OS differences needed to run it *natively* on Windows so the rest
|
||||
of the codebase can stay platform-agnostic. Import from here instead of
|
||||
sprinkling ``os.name == "nt"`` checks (and POSIX-only calls) across modules.
|
||||
|
||||
Design rules:
|
||||
* Stdlib + ctypes only — no new third-party deps (no psutil/pywinpty).
|
||||
* POSIX behaviour is unchanged; Windows gets a faithful equivalent or a
|
||||
safe, documented no-op.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import ntpath
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
import sys
|
||||
from typing import List, Optional
|
||||
import platform
|
||||
|
||||
IS_WINDOWS = os.name == "nt"
|
||||
IS_POSIX = not IS_WINDOWS
|
||||
# Allows APFEL support and ARM-native binary recommendations on Apple Silicon Macs.
|
||||
IS_APPLE_SILICON = (
|
||||
IS_POSIX
|
||||
and platform.system() == "Darwin"
|
||||
and platform.machine().lower()
|
||||
in {
|
||||
"arm64",
|
||||
"aarch64",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
# ── File permissions ────────────────────────────────────────────────────────
|
||||
def safe_chmod(path, mode: int) -> bool:
|
||||
"""``os.chmod`` that is a harmless no-op on Windows.
|
||||
|
||||
On POSIX we apply the mode — used to lock secret/key files down to 0o600.
|
||||
Windows has no POSIX permission bits; files under the user profile are
|
||||
already ACL-restricted to that user, so we skip rather than raise. Returns
|
||||
True when the mode was actually applied.
|
||||
"""
|
||||
if IS_WINDOWS:
|
||||
return False
|
||||
try:
|
||||
os.chmod(path, mode)
|
||||
return True
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
# ── Process detach / liveness / teardown ────────────────────────────────────
|
||||
def detached_popen_kwargs() -> dict:
|
||||
"""Keyword args for :class:`subprocess.Popen` that fully detach a child so
|
||||
it outlives the request/stream that launched it.
|
||||
|
||||
POSIX: ``start_new_session=True`` (setsid) — new session + process group.
|
||||
Windows: ``CREATE_NEW_PROCESS_GROUP | DETACHED_PROCESS`` — the child gets
|
||||
its own process group (so it isn't killed when the parent's console closes)
|
||||
and is detached from any console.
|
||||
"""
|
||||
if IS_WINDOWS:
|
||||
flags = getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0x00000200) | getattr(
|
||||
subprocess, "DETACHED_PROCESS", 0x00000008
|
||||
)
|
||||
return {"creationflags": flags}
|
||||
return {"start_new_session": True}
|
||||
|
||||
|
||||
def pid_alive(pid: Optional[int]) -> bool:
|
||||
"""True if a process with ``pid`` is currently running.
|
||||
|
||||
POSIX uses the classic ``os.kill(pid, 0)`` probe. That is **unsafe on
|
||||
Windows**: CPython's ``os.kill`` calls ``TerminateProcess(handle, sig)`` for
|
||||
any signal other than CTRL_C/CTRL_BREAK, so ``os.kill(pid, 0)`` would *kill*
|
||||
the process it is checking. We instead open the process and read its exit
|
||||
code via the Win32 API.
|
||||
"""
|
||||
if not pid:
|
||||
return False
|
||||
if IS_WINDOWS:
|
||||
import ctypes
|
||||
from ctypes import wintypes
|
||||
|
||||
PROCESS_QUERY_LIMITED_INFORMATION = 0x1000
|
||||
STILL_ACTIVE = 259
|
||||
kernel32 = ctypes.windll.kernel32
|
||||
handle = kernel32.OpenProcess(
|
||||
PROCESS_QUERY_LIMITED_INFORMATION, False, int(pid)
|
||||
)
|
||||
if not handle:
|
||||
return False
|
||||
try:
|
||||
code = wintypes.DWORD()
|
||||
if kernel32.GetExitCodeProcess(handle, ctypes.byref(code)):
|
||||
return code.value == STILL_ACTIVE
|
||||
return False
|
||||
finally:
|
||||
kernel32.CloseHandle(handle)
|
||||
try:
|
||||
os.kill(pid, 0)
|
||||
return True
|
||||
except (OSError, ProcessLookupError):
|
||||
return False
|
||||
|
||||
|
||||
def kill_process_tree(pid: Optional[int]) -> None:
|
||||
"""Terminate ``pid`` and all of its descendants.
|
||||
|
||||
POSIX: signal the whole process group (``killpg``), falling back to a plain
|
||||
``kill`` if the pid isn't a group leader.
|
||||
Windows: ``taskkill /T /F`` walks and kills the child tree (there is no
|
||||
process-group signalling).
|
||||
"""
|
||||
if not pid:
|
||||
return
|
||||
if IS_WINDOWS:
|
||||
try:
|
||||
subprocess.run(
|
||||
["taskkill", "/F", "/T", "/PID", str(pid)],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return
|
||||
import signal
|
||||
|
||||
try:
|
||||
os.killpg(os.getpgid(pid), signal.SIGTERM)
|
||||
except Exception:
|
||||
try:
|
||||
os.kill(pid, signal.SIGTERM)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ── Shell / executable resolution ───────────────────────────────────────────
|
||||
_BASH_CACHE: Optional[str] = None
|
||||
_BASH_PROBED = False
|
||||
|
||||
# Common Git-for-Windows install locations to probe when bash isn't on PATH.
|
||||
_WINDOWS_BASH_ROOT_ENV_VARS = (
|
||||
"ProgramFiles",
|
||||
"ProgramW6432",
|
||||
"ProgramFiles(x86)",
|
||||
"LocalAppData",
|
||||
)
|
||||
_WINDOWS_BASH_DEFAULT_ROOTS = (
|
||||
r"C:\Program Files\Git",
|
||||
r"C:\Program Files (x86)\Git",
|
||||
)
|
||||
_WINDOWS_BASH_RELATIVE_PATHS = (
|
||||
("bin", "bash.exe"),
|
||||
("usr", "bin", "bash.exe"),
|
||||
)
|
||||
|
||||
# Paths to add to the remote SSH probe command to find tools like nvidia-smi that may not be on PATH.
|
||||
_SSH_PATH_MEMBERS = (
|
||||
"/usr/bin",
|
||||
"/usr/local/bin",
|
||||
"/usr/local/cuda/bin",
|
||||
"/usr/lib/wsl/lib"
|
||||
)
|
||||
# Fallback locations for nvidia-smi on WSL and other Linux distros where it may not be on PATH.
|
||||
NVIDIA_PATH_CANDIDATES = (
|
||||
"/usr/bin/nvidia-smi",
|
||||
"/usr/local/bin/nvidia-smi",
|
||||
"/usr/local/cuda/bin/nvidia-smi",
|
||||
"/usr/lib/wsl/lib/nvidia-smi",
|
||||
)
|
||||
|
||||
|
||||
def _ssh_path_override() -> str:
|
||||
"""Build the PATH export snippet used for remote SSH shell probes."""
|
||||
return f"export PATH=\"$PATH:{':'.join(_SSH_PATH_MEMBERS)}\"; "
|
||||
|
||||
|
||||
SSH_PATH_OVERRIDE = _ssh_path_override()
|
||||
|
||||
|
||||
def _windows_bash_fallbacks() -> List[str]:
|
||||
roots: List[str] = []
|
||||
for env_name in _WINDOWS_BASH_ROOT_ENV_VARS:
|
||||
base = os.environ.get(env_name)
|
||||
if base:
|
||||
roots.append(ntpath.join(base, "Git"))
|
||||
if env_name == "LocalAppData":
|
||||
roots.append(ntpath.join(base, "Programs", "Git"))
|
||||
roots.extend(_WINDOWS_BASH_DEFAULT_ROOTS)
|
||||
|
||||
paths: List[str] = []
|
||||
seen = set()
|
||||
for root in roots:
|
||||
for rel in _WINDOWS_BASH_RELATIVE_PATHS:
|
||||
path = ntpath.join(root, *rel)
|
||||
key = path.lower()
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
paths.append(path)
|
||||
return paths
|
||||
|
||||
|
||||
def _is_windows_bash_stub(path: str) -> bool:
|
||||
lowered = path.lower()
|
||||
return (
|
||||
"system32\\bash.exe" in lowered
|
||||
or "sysnative\\bash.exe" in lowered
|
||||
or "windowsapps\\bash.exe" in lowered
|
||||
)
|
||||
|
||||
|
||||
def git_bash_path(path: str | Path) -> str:
|
||||
"""Convert a path to POSIX style suitable for Git Bash on Windows.
|
||||
|
||||
Transforms drive letters (e.g., 'C:\\path') to POSIX '/c/path',
|
||||
and uses forward slashes.
|
||||
"""
|
||||
p = Path(path)
|
||||
p_str = p.as_posix()
|
||||
if IS_WINDOWS and len(p_str) >= 2 and p_str[1] == ":":
|
||||
drive = p_str[0].lower()
|
||||
return f"/{drive}{p_str[2:]}"
|
||||
return p_str
|
||||
|
||||
|
||||
|
||||
def find_bash() -> Optional[str]:
|
||||
"""Locate a real ``bash`` interpreter, or None.
|
||||
|
||||
On Windows this is typically Git Bash / WSL. Many Odysseus features (the
|
||||
agent ``bash`` tool, background jobs, Cookbook scripts) emit bash syntax, so
|
||||
when a bash is present we use it and keep full parity with POSIX. Result is
|
||||
cached.
|
||||
"""
|
||||
global _BASH_CACHE, _BASH_PROBED
|
||||
if _BASH_PROBED:
|
||||
return _BASH_CACHE
|
||||
_BASH_PROBED = True
|
||||
found = which_tool("bash")
|
||||
if found and IS_WINDOWS and _is_windows_bash_stub(found):
|
||||
found = None
|
||||
if not found and IS_WINDOWS:
|
||||
for cand in _windows_bash_fallbacks():
|
||||
if os.path.exists(cand):
|
||||
found = cand
|
||||
break
|
||||
_BASH_CACHE = found
|
||||
return found
|
||||
|
||||
|
||||
def has_bash() -> bool:
|
||||
return find_bash() is not None
|
||||
|
||||
|
||||
def which_tool(name: str) -> Optional[str]:
|
||||
"""``shutil.which`` that also tries Windows executable suffixes.
|
||||
|
||||
On Windows, Node/npm shims are ``npx.cmd``/``npm.cmd`` and binaries end in
|
||||
``.exe``; a bare ``which("npx")`` can miss them depending on PATHEXT. We try
|
||||
the bare name first, then the common suffixes.
|
||||
"""
|
||||
found = shutil.which(name)
|
||||
if found:
|
||||
return found
|
||||
if IS_WINDOWS:
|
||||
for ext in (".cmd", ".exe", ".bat"):
|
||||
found = shutil.which(name + ext)
|
||||
if found:
|
||||
return found
|
||||
return None
|
||||
|
||||
|
||||
def run_script_argv(script_path) -> List[str]:
|
||||
"""argv to execute a shell *script file*.
|
||||
|
||||
Prefers bash (so existing ``.sh`` wrappers work verbatim, including on
|
||||
Windows via Git Bash). On Windows with no bash available, falls back to
|
||||
``cmd.exe /c`` — simple commands still run, but bash-specific syntax won't.
|
||||
Callers that need guaranteed bash should check :func:`has_bash` first and
|
||||
surface a clear "install Git Bash" message.
|
||||
"""
|
||||
bash = find_bash()
|
||||
if bash:
|
||||
return [bash, str(script_path)]
|
||||
if IS_WINDOWS:
|
||||
comspec = os.environ.get("ComSpec", "cmd.exe")
|
||||
return [comspec, "/c", str(script_path)]
|
||||
return ["sh", str(script_path)]
|
||||
|
||||
|
||||
def is_wsl() -> bool:
|
||||
"""True if running inside Windows Subsystem for Linux (WSL)."""
|
||||
import sys
|
||||
if sys.platform.startswith("linux") or os.name == "posix":
|
||||
try:
|
||||
with open("/proc/version", "r", encoding="utf-8", errors="ignore") as f:
|
||||
if "microsoft" in f.read().lower():
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
|
||||
def translate_path(path_str: str) -> str:
|
||||
"""Translate a path (possibly a Windows path) to the current OS format.
|
||||
|
||||
Particularly handles Windows paths (e.g. C:\\foo or C:/foo) when running
|
||||
under WSL, translating them to /mnt/c/foo.
|
||||
Also handles standard path normalization to avoid string breakages.
|
||||
"""
|
||||
if not path_str:
|
||||
return path_str
|
||||
|
||||
if is_wsl():
|
||||
path_str = path_str.replace("\\", "/")
|
||||
import re
|
||||
m = re.match(r"^([a-zA-Z]):(.*)", path_str)
|
||||
if m:
|
||||
drive = m.group(1).lower()
|
||||
rest = m.group(2)
|
||||
if not rest.startswith("/"):
|
||||
rest = "/" + rest
|
||||
return f"/mnt/{drive}{rest}"
|
||||
|
||||
try:
|
||||
return str(Path(path_str).resolve())
|
||||
except Exception:
|
||||
return path_str
|
||||
|
||||
|
||||
def get_wsl_windows_user_profile() -> Optional[str]:
|
||||
"""Retrieve the Windows host User Profile path from inside WSL."""
|
||||
if not is_wsl():
|
||||
return None
|
||||
try:
|
||||
r = run_wsl_windows_powershell("Write-Output $env:USERPROFILE", timeout=5)
|
||||
if r.returncode == 0 and r.stdout.strip():
|
||||
return translate_path(r.stdout.strip())
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
users_dir = "/mnt/c/Users"
|
||||
if os.path.isdir(users_dir):
|
||||
for entry in os.listdir(users_dir):
|
||||
if entry not in ("All Users", "Default", "Default User", "desktop.ini", "Public"):
|
||||
path = os.path.join(users_dir, entry)
|
||||
if os.path.isdir(path):
|
||||
return path
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def _ssh_exec_argv(
|
||||
remote: str,
|
||||
ssh_port: str | None,
|
||||
*,
|
||||
remote_cmd: str | None = None,
|
||||
connect_timeout: int | None = None,
|
||||
strict_host_key_checking: bool | None = None,
|
||||
) -> list[str]:
|
||||
"""Build a consistent ssh argv for remote command execution."""
|
||||
remote_value = str(remote or "").strip()
|
||||
remote_host = remote_value.rsplit("@", 1)[-1]
|
||||
if not remote_value or remote_value.startswith("-") or not remote_host or remote_host.startswith("-"):
|
||||
raise ValueError("Invalid SSH remote host")
|
||||
argv = ["ssh"]
|
||||
if connect_timeout is not None:
|
||||
argv.extend(["-o", f"ConnectTimeout={int(connect_timeout)}"])
|
||||
if strict_host_key_checking is not None:
|
||||
argv.extend(
|
||||
[
|
||||
"-o",
|
||||
"StrictHostKeyChecking=yes"
|
||||
if strict_host_key_checking
|
||||
else "StrictHostKeyChecking=no",
|
||||
]
|
||||
)
|
||||
if ssh_port and ssh_port != "22":
|
||||
argv.extend(["-p", str(ssh_port)])
|
||||
argv.append(remote)
|
||||
if remote_cmd is not None:
|
||||
argv.append(remote_cmd)
|
||||
return argv
|
||||
|
||||
|
||||
def run_ssh_command(
|
||||
remote: str,
|
||||
ssh_port: str | None,
|
||||
remote_cmd: str,
|
||||
*,
|
||||
timeout: float,
|
||||
connect_timeout: int | None = None,
|
||||
strict_host_key_checking: bool | None = None,
|
||||
text: bool = True,
|
||||
) -> subprocess.CompletedProcess:
|
||||
"""Run an ssh command with centralized timeout and stderr/stdout capture."""
|
||||
return subprocess.run(
|
||||
_ssh_exec_argv(
|
||||
remote,
|
||||
ssh_port,
|
||||
remote_cmd=remote_cmd,
|
||||
connect_timeout=connect_timeout,
|
||||
strict_host_key_checking=strict_host_key_checking,
|
||||
),
|
||||
timeout=timeout,
|
||||
capture_output=True,
|
||||
text=text,
|
||||
)
|
||||
|
||||
|
||||
def _windows_powershell_argv(
|
||||
command: str,
|
||||
*,
|
||||
no_profile: bool = True,
|
||||
non_interactive: bool = True,
|
||||
) -> List[str]:
|
||||
argv: List[str] = ["powershell.exe"]
|
||||
if no_profile:
|
||||
argv.append("-NoProfile")
|
||||
if non_interactive:
|
||||
argv.append("-NonInteractive")
|
||||
argv.extend(["-Command", command])
|
||||
return argv
|
||||
|
||||
|
||||
def run_wsl_windows_powershell(
|
||||
command: str,
|
||||
*,
|
||||
timeout: float = 5,
|
||||
) -> subprocess.CompletedProcess[str]:
|
||||
"""Run a PowerShell command on the Windows host from WSL.
|
||||
|
||||
Raises ``RuntimeError`` when called outside WSL.
|
||||
"""
|
||||
|
||||
if not is_wsl():
|
||||
raise RuntimeError("run_wsl_windows_powershell is only supported in WSL")
|
||||
return subprocess.run(
|
||||
_windows_powershell_argv(command),
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
+287
-45
@@ -14,12 +14,54 @@ import logging
|
||||
from datetime import datetime, timezone, timedelta
|
||||
from typing import Dict, Optional
|
||||
|
||||
from .database import Session as DbSession, ChatMessage as DbChatMessage, Document as DbDocument, SessionLocal
|
||||
from sqlalchemy import func
|
||||
|
||||
from .database import Session as DbSession, ChatMessage as DbChatMessage, Document as DbDocument, SessionLocal, utcnow_naive
|
||||
from .models import Session, ChatMessage
|
||||
from src.attachment_refs import persistable_message_content
|
||||
from src.upload_handler import reserve_message_upload_references
|
||||
|
||||
# Re-export singleton accessors from models for convenience
|
||||
from .models import set_session_manager_instance, get_session_manager_instance
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _message_timestamp_iso(value: Optional[datetime]) -> Optional[str]:
|
||||
"""Return a stable ISO timestamp for chat message metadata."""
|
||||
if not value:
|
||||
return None
|
||||
if value.tzinfo is None:
|
||||
value = value.replace(tzinfo=timezone.utc)
|
||||
return value.isoformat().replace("+00:00", "Z")
|
||||
|
||||
|
||||
def _parse_msg_content(raw):
|
||||
"""Parse message content from DB — deserialises JSON arrays back to lists
|
||||
(multimodal content with image/audio attachments)."""
|
||||
if isinstance(raw, list):
|
||||
return raw
|
||||
if isinstance(raw, str) and raw.startswith('[{') and '"type"' in raw:
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
# Only treat as serialized multimodal content when EVERY element is
|
||||
# a dict whose "type" is a recognized content-block kind. Otherwise a
|
||||
# plain text message that merely *looks* like a JSON array of objects
|
||||
# (e.g. a user pasting an API schema/sample with a "type" field) was
|
||||
# silently parsed back into a list, destroying the original string.
|
||||
_BLOCK_TYPES = {
|
||||
"text", "image", "image_url", "audio", "input_audio",
|
||||
"input_image", "document", "file",
|
||||
}
|
||||
if (isinstance(parsed, list) and parsed
|
||||
and all(isinstance(p, dict) and p.get("type") in _BLOCK_TYPES
|
||||
for p in parsed)):
|
||||
return parsed
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
return raw
|
||||
|
||||
|
||||
class SessionManager:
|
||||
"""
|
||||
Manages chat sessions with database persistence.
|
||||
@@ -34,6 +76,7 @@ class SessionManager:
|
||||
def __init__(self, sessions_file: str = None):
|
||||
# sessions_file kept for backward compat, not used
|
||||
self.sessions: Dict[str, Session] = {}
|
||||
self.upload_handler = None
|
||||
self.load_sessions()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -51,14 +94,28 @@ class SessionManager:
|
||||
try:
|
||||
db_sessions = db.query(DbSession).filter(
|
||||
DbSession.archived == False,
|
||||
DbSession.message_count > 0,
|
||||
DbSession.messages.any(),
|
||||
).order_by(DbSession.last_accessed.desc()).limit(100).all()
|
||||
|
||||
# message_count is derived metadata and can drift after interrupted
|
||||
# or legacy writes. Count only the bounded discovery set so startup
|
||||
# remains metadata-only while lazy hydration sees an authoritative
|
||||
# positive count for every discovered non-empty session.
|
||||
message_counts = {}
|
||||
if db_sessions:
|
||||
message_counts = dict(
|
||||
db.query(DbChatMessage.session_id, func.count(DbChatMessage.id))
|
||||
.filter(DbChatMessage.session_id.in_([row.id for row in db_sessions]))
|
||||
.group_by(DbChatMessage.session_id)
|
||||
.all()
|
||||
)
|
||||
|
||||
loaded_count = 0
|
||||
for db_session in db_sessions:
|
||||
try:
|
||||
session = self._db_to_session_meta(db_session)
|
||||
if session is not None:
|
||||
session.message_count = message_counts[db_session.id]
|
||||
self.sessions[db_session.id] = session
|
||||
loaded_count += 1
|
||||
except Exception as e:
|
||||
@@ -93,6 +150,12 @@ class SessionManager:
|
||||
history=[],
|
||||
owner=getattr(db_session, "owner", None),
|
||||
is_important=getattr(db_session, "is_important", False) or False,
|
||||
memory_extraction_enabled=getattr(db_session, "memory_extraction_enabled", True) is not False,
|
||||
skill_injection_enabled=getattr(db_session, "skill_injection_enabled", True) is not False,
|
||||
thinking_mode=getattr(db_session, "thinking_mode", "") or "off",
|
||||
temperature_override=getattr(db_session, "temperature_override", None),
|
||||
max_tokens_override=getattr(db_session, "max_tokens_override", None),
|
||||
cwd=getattr(db_session, "cwd", None) or None,
|
||||
)
|
||||
session.message_count = getattr(db_session, "message_count", 0) or 0
|
||||
return session
|
||||
@@ -107,9 +170,10 @@ class SessionManager:
|
||||
meta = json.loads(db_msg.meta_data) if db_msg.meta_data else {}
|
||||
if meta is None: meta = {}
|
||||
meta['_db_id'] = db_msg.id
|
||||
meta.setdefault('timestamp', _message_timestamp_iso(db_msg.timestamp))
|
||||
history.append(ChatMessage(
|
||||
role=db_msg.role,
|
||||
content=db_msg.content,
|
||||
content=_parse_msg_content(db_msg.content),
|
||||
metadata=meta,
|
||||
))
|
||||
else:
|
||||
@@ -121,9 +185,10 @@ class SessionManager:
|
||||
meta = json.loads(db_msg.meta_data) if db_msg.meta_data else {}
|
||||
if meta is None: meta = {}
|
||||
meta['_db_id'] = db_msg.id
|
||||
meta.setdefault('timestamp', _message_timestamp_iso(db_msg.timestamp))
|
||||
history.append(ChatMessage(
|
||||
role=db_msg.role,
|
||||
content=db_msg.content,
|
||||
content=_parse_msg_content(db_msg.content),
|
||||
metadata=meta,
|
||||
))
|
||||
|
||||
@@ -149,9 +214,20 @@ class SessionManager:
|
||||
history=history,
|
||||
owner=getattr(db_session, 'owner', None),
|
||||
is_important=getattr(db_session, 'is_important', False) or False,
|
||||
memory_extraction_enabled=getattr(db_session, 'memory_extraction_enabled', True) is not False,
|
||||
skill_injection_enabled=getattr(db_session, 'skill_injection_enabled', True) is not False,
|
||||
thinking_mode=getattr(db_session, "thinking_mode", "") or "off",
|
||||
temperature_override=getattr(db_session, "temperature_override", None),
|
||||
max_tokens_override=getattr(db_session, "max_tokens_override", None),
|
||||
cwd=getattr(db_session, "cwd", None) or None,
|
||||
)
|
||||
|
||||
session.message_count = getattr(db_session, 'message_count', len(history))
|
||||
# The rows just loaded are the whole transcript, so they — not the
|
||||
# denormalized sessions.message_count column — are the truth for this
|
||||
# cached object. get_session's hydration gate compares against this
|
||||
# number; seeding it from a drifted column would ask for a reload that
|
||||
# can never close the gap.
|
||||
session.message_count = len(history)
|
||||
return session
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -162,12 +238,17 @@ class SessionManager:
|
||||
"""
|
||||
Add a message to a session and persist to database.
|
||||
|
||||
Updates the authoritative history list and persists through this
|
||||
manager directly so tests and temporary managers do not depend on the
|
||||
process-wide session-manager singleton.
|
||||
|
||||
Args:
|
||||
session_id: Session ID
|
||||
message: ChatMessage to add
|
||||
"""
|
||||
session = self.get_session(session_id)
|
||||
session.history.append(message)
|
||||
session._history = session.history
|
||||
session.message_count = len(session.history)
|
||||
|
||||
self._persist_message(session_id, message)
|
||||
@@ -176,31 +257,59 @@ class SessionManager:
|
||||
"""Persist a single message to the database."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session is None:
|
||||
# A stream/tool callback can outlive a session delete. Do not
|
||||
# create a chat_messages row with no parent session; also drop
|
||||
# any stale cached session so later writes fail closed too.
|
||||
self.sessions.pop(session_id, None)
|
||||
logger.warning("Dropping message for deleted session %s", session_id)
|
||||
return
|
||||
|
||||
missing_upload_id = reserve_message_upload_references(
|
||||
getattr(self, "upload_handler", None),
|
||||
getattr(db_session, "owner", None),
|
||||
message.content,
|
||||
message.metadata,
|
||||
)
|
||||
if missing_upload_id:
|
||||
raise ValueError(
|
||||
f"Referenced upload is no longer available: {missing_upload_id}"
|
||||
)
|
||||
|
||||
msg_id = str(uuid.uuid4())
|
||||
msg_time = datetime.utcnow()
|
||||
if message.metadata is None:
|
||||
message.metadata = {}
|
||||
message.metadata.setdefault('timestamp', _message_timestamp_iso(msg_time))
|
||||
# Multimodal content may contain provider data URLs for the live
|
||||
# model call. Persist only readable text plus attachment references
|
||||
# so chat_messages/FTS do not duplicate upload bytes.
|
||||
_content = persistable_message_content(message.content, message.metadata)
|
||||
db_message = DbChatMessage(
|
||||
id=msg_id,
|
||||
session_id=session_id,
|
||||
role=message.role,
|
||||
content=message.content,
|
||||
meta_data=json.dumps(message.metadata) if message.metadata else None
|
||||
content=_content,
|
||||
meta_data=json.dumps(message.metadata) if message.metadata else None,
|
||||
timestamp=msg_time,
|
||||
)
|
||||
db.add(db_message)
|
||||
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session:
|
||||
db_session.message_count = len(self.sessions.get(session_id, {}).history) if session_id in self.sessions else 0
|
||||
_now = datetime.now(timezone.utc)
|
||||
db_session.last_accessed = _now
|
||||
# Clean "last conversation" timestamp — only bumped here on a
|
||||
# real message persist, so it powers an accurate "Last active"
|
||||
# sort that ignores renames / model swaps / mere opens.
|
||||
db_session.last_message_at = _now
|
||||
if session_id in self.sessions:
|
||||
db_session.message_count = len(self.sessions[session_id].history)
|
||||
else:
|
||||
db_session.message_count = 0
|
||||
_now = datetime.now(timezone.utc)
|
||||
db_session.last_accessed = _now
|
||||
# Clean "last conversation" timestamp — only bumped here on a
|
||||
# real message persist, so it powers an accurate "Last active"
|
||||
# sort that ignores renames / model swaps / mere opens.
|
||||
db_session.last_message_at = _now
|
||||
|
||||
db.commit()
|
||||
|
||||
# Store DB ID on the in-memory message for edit/delete by ID
|
||||
if message.metadata is None:
|
||||
message.metadata = {}
|
||||
message.metadata['_db_id'] = msg_id
|
||||
|
||||
logger.debug(f"Persisted message to session {session_id}")
|
||||
@@ -231,13 +340,17 @@ class SessionManager:
|
||||
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session:
|
||||
db_session.message_count = keep_count
|
||||
# keep_count can exceed the real message total (e.g. the AI tool
|
||||
# defaults to keep_count=10 on a short session); message_count must
|
||||
# track the rows that actually remain, not the requested cap.
|
||||
db_session.message_count = min(keep_count, len(db_messages))
|
||||
db_session.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
db.commit()
|
||||
|
||||
# Update in-memory
|
||||
session.history = session.history[:keep_count]
|
||||
session._history = session.history
|
||||
|
||||
logger.info(f"Truncated session {session_id} to {keep_count} messages")
|
||||
return True
|
||||
@@ -254,6 +367,28 @@ class SessionManager:
|
||||
session = self.get_session(session_id)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session is None:
|
||||
logger.warning("Cannot replace history for missing session %s", session_id)
|
||||
return False
|
||||
|
||||
# Reserve every incoming attachment before removing any durable
|
||||
# message row. reserve_upload() shares the upload lifecycle lock
|
||||
# with cleanup, so an upload cannot be deleted between this
|
||||
# ownership check/access touch and the replacement transaction.
|
||||
# A failed reservation must leave the existing transcript intact.
|
||||
for message in messages:
|
||||
missing_upload_id = reserve_message_upload_references(
|
||||
getattr(self, "upload_handler", None),
|
||||
getattr(db_session, "owner", None),
|
||||
message.content,
|
||||
message.metadata,
|
||||
)
|
||||
if missing_upload_id:
|
||||
raise ValueError(
|
||||
f"Referenced upload is no longer available: {missing_upload_id}"
|
||||
)
|
||||
|
||||
db.query(DbChatMessage).filter(DbChatMessage.session_id == session_id).delete()
|
||||
now = datetime.now(timezone.utc)
|
||||
for i, message in enumerate(messages):
|
||||
@@ -262,7 +397,9 @@ class SessionManager:
|
||||
id=msg_id,
|
||||
session_id=session_id,
|
||||
role=message.role,
|
||||
content=message.content,
|
||||
# Mirrors _persist_message: keep raw media bytes out of the
|
||||
# persisted transcript and search index.
|
||||
content=persistable_message_content(message.content, message.metadata),
|
||||
meta_data=json.dumps(message.metadata) if message.metadata else None,
|
||||
timestamp=now + timedelta(microseconds=i),
|
||||
)
|
||||
@@ -271,15 +408,14 @@ class SessionManager:
|
||||
message.metadata = {}
|
||||
message.metadata["_db_id"] = msg_id
|
||||
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session:
|
||||
db_session.message_count = len(messages)
|
||||
db_session.updated_at = now
|
||||
db_session.last_accessed = now
|
||||
db_session.last_message_at = now
|
||||
db_session.message_count = len(messages)
|
||||
db_session.updated_at = now
|
||||
db_session.last_accessed = now
|
||||
db_session.last_message_at = now
|
||||
|
||||
db.commit()
|
||||
session.history = list(messages)
|
||||
session._history = session.history
|
||||
session.message_count = len(messages)
|
||||
logger.info("Replaced session %s history with %d messages", session_id, len(messages))
|
||||
return True
|
||||
@@ -295,24 +431,85 @@ class SessionManager:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_session(self, session_id: str) -> Session:
|
||||
"""Get a session by ID, loading from DB if needed.
|
||||
"""Get a session by ID, loading complete DB history when needed.
|
||||
|
||||
Sessions seeded by `load_sessions` start with empty history. The
|
||||
first read here hydrates them with the message rows.
|
||||
Sessions seeded by ``load_sessions`` start with empty history, and a
|
||||
cached session can also become partially stale. Refresh metadata first,
|
||||
then hydrate whenever the cached transcript is short of the stored rows.
|
||||
Model-send routes enter through this method before building context,
|
||||
while paginated display history reads SQLite directly.
|
||||
|
||||
The gate compares against ``sync_session_metadata``'s reconciled count
|
||||
(the real ``chat_messages`` total), never the denormalized column, so a
|
||||
hydrate always closes the gap and the next read is a cache hit.
|
||||
"""
|
||||
if session_id not in self.sessions:
|
||||
self._load_session_from_db(session_id)
|
||||
else:
|
||||
cached = self.sessions[session_id]
|
||||
# Lazy hydrate: metadata-only entries get their messages on first read.
|
||||
if not cached.history and getattr(cached, "message_count", 0) > 0:
|
||||
self._load_session_from_db(session_id)
|
||||
|
||||
# Keep model/endpoint metadata fresh. Endpoint deletion can clear the
|
||||
# DB row while a session object is still cached in RAM. Refreshing first
|
||||
# also exposes the authoritative message count before completeness is
|
||||
# checked.
|
||||
self.sync_session_metadata(session_id)
|
||||
|
||||
cached = self.sessions[session_id]
|
||||
cached_count = len(cached.history or [])
|
||||
stored_count = int(getattr(cached, "message_count", 0) or 0)
|
||||
if cached_count < stored_count:
|
||||
self._load_session_from_db(session_id)
|
||||
|
||||
# Update last_accessed
|
||||
self._touch_session(session_id)
|
||||
|
||||
return self.sessions[session_id]
|
||||
|
||||
def sync_session_metadata(self, session_id: str) -> bool:
|
||||
"""Refresh non-message session fields from the DB into the cached object.
|
||||
|
||||
``message_count`` is reconciled against the real ``chat_messages`` rows
|
||||
rather than copied from the denormalized ``sessions.message_count``
|
||||
column. That column drifts in normal operation — ``_persist_message``
|
||||
swallows a failed insert but ``add_message`` has already appended in
|
||||
memory, so the next successful persist writes rows+1, and a persist for
|
||||
an uncached session writes 0. Hydration keys off this number: a
|
||||
drifted-high column would reload the whole transcript on every warm
|
||||
read, and a drifted-low one would leave the model a truncated one.
|
||||
"""
|
||||
session = self.sessions.get(session_id)
|
||||
if session is None:
|
||||
return False
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session is None:
|
||||
return False
|
||||
headers = db_session.headers
|
||||
if isinstance(headers, str):
|
||||
try:
|
||||
headers = json.loads(headers)
|
||||
except json.JSONDecodeError:
|
||||
headers = {}
|
||||
session.name = db_session.name
|
||||
session.endpoint_url = db_session.endpoint_url or ""
|
||||
session.model = db_session.model or ""
|
||||
session.headers = headers or {}
|
||||
session.rag = db_session.rag
|
||||
session.archived = db_session.archived
|
||||
session.owner = getattr(db_session, "owner", None)
|
||||
session.is_important = getattr(db_session, "is_important", False) or False
|
||||
session.cwd = getattr(db_session, "cwd", None) or None
|
||||
session.message_count = (
|
||||
db.query(DbChatMessage)
|
||||
.filter(DbChatMessage.session_id == session_id)
|
||||
.count()
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Error syncing session metadata {session_id}: {e}")
|
||||
return False
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def _load_session_from_db(self, session_id: str):
|
||||
"""Hydrate a single session (with messages) from the database."""
|
||||
db = SessionLocal()
|
||||
@@ -361,9 +558,12 @@ class SessionManager:
|
||||
endpoint_url: str,
|
||||
model: str,
|
||||
rag: bool = False,
|
||||
owner: str = None
|
||||
owner: str = None,
|
||||
cwd: str = None,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
) -> Session:
|
||||
"""Create a new session and save to database."""
|
||||
session_headers = dict(headers or {})
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db_session = DbSession(
|
||||
@@ -372,8 +572,9 @@ class SessionManager:
|
||||
endpoint_url=endpoint_url,
|
||||
model=model,
|
||||
rag=rag,
|
||||
headers={},
|
||||
headers=session_headers,
|
||||
owner=owner,
|
||||
cwd=cwd or None,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc)
|
||||
)
|
||||
@@ -386,8 +587,9 @@ class SessionManager:
|
||||
endpoint_url=endpoint_url,
|
||||
model=model,
|
||||
rag=rag,
|
||||
headers={},
|
||||
headers=session_headers,
|
||||
owner=owner,
|
||||
cwd=cwd or None,
|
||||
)
|
||||
|
||||
self.sessions[session_id] = session
|
||||
@@ -404,6 +606,12 @@ class SessionManager:
|
||||
"""Permanently delete a session and all its messages."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
try:
|
||||
from src.session_image_cleanup import cleanup_session_images
|
||||
cleanup_session_images(session_id, db=db)
|
||||
except Exception as e:
|
||||
logger.warning(f"Image cleanup failed while deleting session {session_id}: {e}")
|
||||
|
||||
# Detach documents so they survive as orphans in the library
|
||||
db.query(DbDocument).filter(DbDocument.session_id == session_id).update(
|
||||
{DbDocument.session_id: None}, synchronize_session=False
|
||||
@@ -416,11 +624,17 @@ class SessionManager:
|
||||
db_session = db.query(DbSession).filter(DbSession.id == session_id).first()
|
||||
if db_session:
|
||||
db.delete(db_session)
|
||||
|
||||
# Drop the in-memory copy even when there is no DB row. A "ghost"
|
||||
# session lives only here (never persisted, or its row was removed
|
||||
# out-of-band); without this it can never be cleared and keeps
|
||||
# 404ing on every operation (issue #1044).
|
||||
removed_in_memory = self.sessions.pop(session_id, None) is not None
|
||||
|
||||
if db_session or removed_in_memory:
|
||||
# Commit the document-detach / message-delete above (a no-op when
|
||||
# the ghost had no rows) together with the session delete.
|
||||
db.commit()
|
||||
|
||||
if session_id in self.sessions:
|
||||
del self.sessions[session_id]
|
||||
|
||||
logger.info(f"Deleted session {session_id}")
|
||||
return True
|
||||
return False
|
||||
@@ -513,24 +727,52 @@ class SessionManager:
|
||||
def save_sessions(self):
|
||||
"""No-op for DB compatibility."""
|
||||
|
||||
def ensure_task_session(self, session_id: str, name: str, endpoint_url: str, model: str, owner: str = None, task: object = None) -> Session:
|
||||
"""Create a task session if it doesn't exist, or return the existing one.
|
||||
|
||||
Unlike create_session, this checks the cache first and does NOT
|
||||
overwrite an existing in-memory session. The task scheduler must
|
||||
use this instead of direct dict assignment.
|
||||
"""
|
||||
if session_id in self.sessions:
|
||||
return self.sessions[session_id]
|
||||
|
||||
session = self.create_session(session_id, name, endpoint_url, model, owner=owner)
|
||||
if task is not None:
|
||||
task.session_id = session_id
|
||||
return session
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Cleanup
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def cleanup_empty_sessions(self, auto_archive_days: int = 30) -> dict:
|
||||
"""Clean up empty and old sessions."""
|
||||
def cleanup_empty_sessions(self, auto_archive_days: int = 30, min_age_hours: int = 1) -> dict:
|
||||
"""Clean up empty and old sessions.
|
||||
|
||||
Args:
|
||||
auto_archive_days: Age in days before non-important sessions are archived.
|
||||
min_age_hours: Minimum age in hours before an empty session can be deleted.
|
||||
Prevents deleting sessions that were just created.
|
||||
"""
|
||||
db = SessionLocal()
|
||||
stats = {'deleted_empty': 0, 'archived_old': 0, 'total_checked': 0}
|
||||
|
||||
try:
|
||||
all_sessions = db.query(DbSession).all()
|
||||
cutoff_date = datetime.now(timezone.utc) - timedelta(days=auto_archive_days)
|
||||
cutoff_date = utcnow_naive() - timedelta(days=auto_archive_days)
|
||||
min_age = utcnow_naive() - timedelta(hours=min_age_hours)
|
||||
|
||||
for db_session in all_sessions:
|
||||
stats['total_checked'] += 1
|
||||
|
||||
# Delete empty sessions
|
||||
# Delete empty sessions only if older than min_age_hours
|
||||
if db_session.message_count == 0:
|
||||
if db_session.created_at is not None:
|
||||
created = db_session.created_at
|
||||
if created.tzinfo is None:
|
||||
created = created.replace(tzinfo=timezone.utc)
|
||||
if created > min_age:
|
||||
continue # Too young to delete
|
||||
if db_session.id in self.sessions:
|
||||
del self.sessions[db_session.id]
|
||||
db.delete(db_session)
|
||||
|
||||
Reference in New Issue
Block a user