mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-09 08:22:19 +02:00
Squash Odysseus development history
This commit is contained in:
+246
-12
@@ -5,7 +5,7 @@ from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from urllib.parse import unquote, urlparse
|
||||
from sqlalchemy import DDL, event, create_engine, Column, String, Text, Boolean, DateTime, Integer, ForeignKey, JSON, Index, func, inspect, text
|
||||
from sqlalchemy import DDL, event, create_engine, Column, String, Text, Boolean, DateTime, Integer, Float, ForeignKey, JSON, Index, func, inspect, text
|
||||
from sqlalchemy.engine import Engine, make_url
|
||||
from sqlalchemy.types import TypeDecorator
|
||||
from sqlalchemy.ext.declarative import declarative_base, declared_attr
|
||||
@@ -75,7 +75,7 @@ DATABASE_URL = _normalize_sqlite_url(os.getenv("DATABASE_URL", _default_database
|
||||
# Create engine
|
||||
engine = create_engine(
|
||||
DATABASE_URL,
|
||||
connect_args={"check_same_thread": False} if "sqlite" in DATABASE_URL else {}
|
||||
connect_args={"check_same_thread": False, "timeout": 30} if "sqlite" in DATABASE_URL else {}
|
||||
)
|
||||
|
||||
|
||||
@@ -144,6 +144,8 @@ def set_sqlite_pragma(dbapi_connection, connection_record):
|
||||
if isinstance(dbapi_connection, sqlite3.Connection):
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.execute("PRAGMA busy_timeout=30000")
|
||||
cursor.execute("PRAGMA journal_mode=WAL")
|
||||
cursor.close()
|
||||
|
||||
|
||||
@@ -191,9 +193,15 @@ class Session(TimestampMixin, Base):
|
||||
# Configuration flags
|
||||
rag = Column(Boolean, default=False)
|
||||
archived = Column(Boolean, default=False)
|
||||
memory_extraction_enabled = Column(Boolean, default=True)
|
||||
skill_injection_enabled = Column(Boolean, default=True)
|
||||
thinking_mode = Column(String, nullable=True, default="off")
|
||||
temperature_override = Column(Float, nullable=True, default=None)
|
||||
max_tokens_override = Column(Integer, nullable=True, default=None)
|
||||
|
||||
# Organization
|
||||
folder = Column(String, nullable=True, default=None)
|
||||
cwd = Column(String, nullable=True, default=None)
|
||||
|
||||
# Headers stored as JSON
|
||||
headers = Column(JSON, default=dict)
|
||||
@@ -219,6 +227,7 @@ class Session(TimestampMixin, Base):
|
||||
message_count = Column(Integer, default=0)
|
||||
total_input_tokens = Column(Integer, default=0)
|
||||
total_output_tokens = Column(Integer, default=0)
|
||||
total_cost_usd = Column(Float, default=0.0)
|
||||
mode = Column(String, nullable=True) # 'agent', 'chat', or 'research'
|
||||
crew_member_id = Column(String, nullable=True) # links to crew_members.id
|
||||
|
||||
@@ -239,6 +248,11 @@ class Session(TimestampMixin, Base):
|
||||
'endpoint_url': self.endpoint_url,
|
||||
'rag': self.rag,
|
||||
'archived': self.archived,
|
||||
'memory_extraction_enabled': self.memory_extraction_enabled is not False,
|
||||
'skill_injection_enabled': self.skill_injection_enabled is not False,
|
||||
'thinking_mode': self.thinking_mode or '',
|
||||
'temperature_override': self.temperature_override,
|
||||
'max_tokens_override': self.max_tokens_override,
|
||||
'created_at': self.created_at.isoformat() if self.created_at else None,
|
||||
'updated_at': self.updated_at.isoformat() if self.updated_at else None,
|
||||
'last_accessed': self.last_accessed.isoformat() if self.last_accessed else None,
|
||||
@@ -248,6 +262,7 @@ class Session(TimestampMixin, Base):
|
||||
'folder': self.folder,
|
||||
'total_input_tokens': self.total_input_tokens or 0,
|
||||
'total_output_tokens': self.total_output_tokens or 0,
|
||||
'total_cost_usd': self.total_cost_usd or 0.0,
|
||||
'crew_member_id': self.crew_member_id,
|
||||
}
|
||||
|
||||
@@ -280,6 +295,22 @@ class ChatMessage(Base):
|
||||
Index('ix_messages_session_time', 'session_id', 'timestamp'), # Composite for efficient message retrieval
|
||||
)
|
||||
|
||||
class BackgroundToolJob(Base):
|
||||
"""Durable origin and once-only chat delivery for background tool work."""
|
||||
__tablename__ = "background_tool_jobs"
|
||||
id = Column(String, primary_key=True)
|
||||
session_id = Column(String, ForeignKey("sessions.id", ondelete="CASCADE"), nullable=False, index=True)
|
||||
owner = Column(String, nullable=False, index=True)
|
||||
tool = Column(String, nullable=False)
|
||||
query = Column(Text, nullable=False)
|
||||
rounds = Column(Integer, nullable=True)
|
||||
status = Column(String, nullable=False, default="running", index=True)
|
||||
payload = Column(Text, nullable=True)
|
||||
summary = Column(Text, nullable=True)
|
||||
message_id = Column(String, nullable=True)
|
||||
created_at = Column(DateTime, default=utcnow_naive)
|
||||
|
||||
|
||||
class Document(TimestampMixin, Base):
|
||||
"""Living document that the AI can create and edit in-place."""
|
||||
__tablename__ = "documents"
|
||||
@@ -544,6 +575,9 @@ class ModelEndpoint(TimestampMixin, Base):
|
||||
# can be toggled per-endpoint in the UI. NULL = unknown, falls
|
||||
# back to the model-name keyword heuristic in agent_loop.py.
|
||||
supports_tools = Column(Boolean, nullable=True, default=None)
|
||||
# JSON object: model id -> native tool schema surface preference.
|
||||
# Values: none, compact, full. Missing key = legacy automatic behavior.
|
||||
model_tool_modes = Column(Text, nullable=True)
|
||||
# Per-user ownership. NULL = legacy/shared (visible to every user) — this
|
||||
# is the historical default. When non-null, the model picker only shows
|
||||
# the endpoint to that user (admins always see everything).
|
||||
@@ -830,6 +864,23 @@ class TaskRun(Base):
|
||||
)
|
||||
|
||||
|
||||
class NotificationLog(Base):
|
||||
"""Persisted task notifications, including completion and error text."""
|
||||
__tablename__ = "notification_logs"
|
||||
|
||||
id = Column(String, primary_key=True, index=True)
|
||||
owner = Column(String, nullable=True, index=True)
|
||||
task_name = Column(String, nullable=False)
|
||||
task_id = Column(String, nullable=True, index=True)
|
||||
status = Column(String, nullable=False, default="success")
|
||||
body = Column(Text, nullable=True)
|
||||
timestamp = Column(DateTime, nullable=False, default=utcnow_naive, index=True)
|
||||
|
||||
__table_args__ = (
|
||||
Index('ix_notification_logs_owner_time', 'owner', 'timestamp'),
|
||||
)
|
||||
|
||||
|
||||
class Memory(Base):
|
||||
"""
|
||||
SQLAlchemy model for Memory table.
|
||||
@@ -910,6 +961,74 @@ def _migrate_add_last_message_at_column():
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_memory_extraction_enabled_column():
|
||||
"""Add per-session auto memory extraction toggle."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
columns = [row[1] for row in conn.execute("PRAGMA table_info(sessions)").fetchall()]
|
||||
if "memory_extraction_enabled" not in columns:
|
||||
conn.execute("ALTER TABLE sessions ADD COLUMN memory_extraction_enabled BOOLEAN DEFAULT 1")
|
||||
conn.commit()
|
||||
logging.getLogger(__name__).info("Migrated: added memory_extraction_enabled to sessions")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"memory_extraction_enabled migration failed: {e}")
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_skill_injection_enabled_column():
|
||||
"""Add per-session skill injection toggle."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
columns = [row[1] for row in conn.execute("PRAGMA table_info(sessions)").fetchall()]
|
||||
if "skill_injection_enabled" not in columns:
|
||||
conn.execute("ALTER TABLE sessions ADD COLUMN skill_injection_enabled BOOLEAN DEFAULT 1")
|
||||
conn.commit()
|
||||
logging.getLogger(__name__).info("Migrated: added skill_injection_enabled to sessions")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"skill_injection_enabled migration failed: {e}")
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_session_generation_settings_columns():
|
||||
"""Add per-chat model generation controls."""
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
columns = {row[1] for row in conn.execute("PRAGMA table_info(sessions)").fetchall()}
|
||||
additions = {
|
||||
"thinking_mode": "VARCHAR DEFAULT 'off'",
|
||||
"temperature_override": "FLOAT",
|
||||
"max_tokens_override": "INTEGER",
|
||||
}
|
||||
for name, sql_type in additions.items():
|
||||
if name not in columns:
|
||||
conn.execute(f"ALTER TABLE sessions ADD COLUMN {name} {sql_type}")
|
||||
conn.commit()
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"session generation settings migration failed: {e}")
|
||||
finally:
|
||||
if conn is not None:
|
||||
conn.close()
|
||||
|
||||
def _migrate_add_document_archived_column():
|
||||
"""Add `archived` to documents (soft-archive flag). Guarded + idempotent."""
|
||||
import sqlite3
|
||||
@@ -1159,6 +1278,30 @@ def _migrate_add_supports_tools_column():
|
||||
pass
|
||||
|
||||
|
||||
def _migrate_add_model_tool_modes_column():
|
||||
"""Add per-model tool-surface preferences to model_endpoints if missing."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.execute("PRAGMA table_info(model_endpoints)")
|
||||
columns = [row[1] for row in cursor.fetchall()]
|
||||
if columns and "model_tool_modes" not in columns:
|
||||
conn.execute("ALTER TABLE model_endpoints ADD COLUMN model_tool_modes TEXT")
|
||||
conn.commit()
|
||||
logging.getLogger(__name__).info("Migrated: added 'model_tool_modes' column to model_endpoints")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"model_tool_modes migration failed: {e}")
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _migrate_add_cached_models_column():
|
||||
"""Add cached_models column to model_endpoints if it doesn't exist."""
|
||||
import sqlite3
|
||||
@@ -1282,6 +1425,29 @@ def _migrate_add_folder_column():
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_session_cwd_column():
|
||||
"""Add cwd column to sessions table if it doesn't exist."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.execute("PRAGMA table_info(sessions)")
|
||||
columns = [row[1] for row in cursor.fetchall()]
|
||||
if "cwd" not in columns:
|
||||
conn.execute("ALTER TABLE sessions ADD COLUMN cwd TEXT")
|
||||
conn.commit()
|
||||
logging.getLogger(__name__).info("Migrated: added 'cwd' column to sessions")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"Migration check for cwd failed: {e}")
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_token_columns():
|
||||
"""Add cumulative token tracking columns to sessions table."""
|
||||
import sqlite3
|
||||
@@ -1306,6 +1472,29 @@ def _migrate_add_token_columns():
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_total_cost_usd():
|
||||
"""Add cumulative USD cost column to sessions table."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.execute("PRAGMA table_info(sessions)")
|
||||
columns = [row[1] for row in cursor.fetchall()]
|
||||
if "total_cost_usd" not in columns:
|
||||
conn.execute("ALTER TABLE sessions ADD COLUMN total_cost_usd REAL DEFAULT 0.0")
|
||||
conn.commit()
|
||||
logging.getLogger(__name__).info("Migrated: added total_cost_usd column to sessions")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"Migration check for total_cost_usd failed: {e}")
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_owner_to_table(table_name: str, index_name: str):
|
||||
"""Generic helper: add owner TEXT column + index to a table if missing."""
|
||||
import sqlite3
|
||||
@@ -1824,6 +2013,7 @@ class Note(TimestampMixin, Base):
|
||||
session_id = Column(String, nullable=True)
|
||||
sort_order = Column(Integer, default=0)
|
||||
image_url = Column(String, nullable=True) # uploaded image URL (relative path)
|
||||
gallery_id = Column(String, nullable=True, index=True) # stable Gallery image for drawings
|
||||
repeat = Column(String, default="none") # none, daily, weekly, monthly, yearly
|
||||
# Auto-AI fields — populated by /api/notes/{id}/classify. The classification
|
||||
# JSON shape is { kind, solvable, confidence, task_prompt, tools, items?: [...] }.
|
||||
@@ -2109,12 +2299,18 @@ def init_db():
|
||||
_migrate_add_model_endpoint_owner_column()
|
||||
_migrate_add_provider_auth_id_column()
|
||||
_migrate_add_supports_tools_column()
|
||||
_migrate_add_model_tool_modes_column()
|
||||
_migrate_add_task_run_model_column()
|
||||
_migrate_add_owner_column()
|
||||
_migrate_add_document_archived_column()
|
||||
_migrate_add_last_message_at_column()
|
||||
_migrate_add_memory_extraction_enabled_column()
|
||||
_migrate_add_skill_injection_enabled_column()
|
||||
_migrate_add_session_generation_settings_columns()
|
||||
_migrate_add_folder_column()
|
||||
_migrate_add_session_cwd_column()
|
||||
_migrate_add_token_columns()
|
||||
_migrate_add_total_cost_usd()
|
||||
_migrate_add_mode_column()
|
||||
_migrate_add_multiuser_owner_columns()
|
||||
_migrate_add_gallery_caption_column()
|
||||
@@ -2142,6 +2338,7 @@ def init_db():
|
||||
_migrate_add_calendar_account_id()
|
||||
_migrate_add_caldav_sync_columns()
|
||||
_migrate_add_calendar_recurrence_exdates()
|
||||
_migrate_add_note_gallery_id()
|
||||
_migrate_chat_messages_fts()
|
||||
_migrate_encrypt_email_passwords()
|
||||
_migrate_encrypt_signatures()
|
||||
@@ -2239,17 +2436,33 @@ def _migrate_chat_messages_fts():
|
||||
END;
|
||||
"""
|
||||
)
|
||||
conn.execute(
|
||||
f"""
|
||||
INSERT INTO chat_messages_fts(content, message_id, session_id, role)
|
||||
SELECT {fts_content_expr_cm}, cm.id, cm.session_id, cm.role
|
||||
FROM chat_messages cm
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM chat_messages_fts fts
|
||||
WHERE fts.message_id = cm.id
|
||||
# message_id is deliberately UNINDEXED in the FTS table. A correlated
|
||||
# NOT EXISTS against it therefore becomes quadratic once the transcript
|
||||
# grows large, even when there is nothing left to backfill. Build a
|
||||
# temporary indexed set only when the row counts show that reconciliation
|
||||
# is needed. Normal inserts/updates/deletes stay synchronized by the
|
||||
# triggers above.
|
||||
chat_count = conn.execute("SELECT COUNT(*) FROM chat_messages").fetchone()[0]
|
||||
fts_count = conn.execute("SELECT COUNT(*) FROM chat_messages_fts").fetchone()[0]
|
||||
if chat_count != fts_count:
|
||||
conn.execute(
|
||||
"CREATE TEMP TABLE IF NOT EXISTS _odysseus_fts_message_ids "
|
||||
"(message_id TEXT PRIMARY KEY) WITHOUT ROWID"
|
||||
)
|
||||
conn.execute("DELETE FROM temp._odysseus_fts_message_ids")
|
||||
conn.execute(
|
||||
"INSERT OR IGNORE INTO temp._odysseus_fts_message_ids(message_id) "
|
||||
"SELECT message_id FROM chat_messages_fts"
|
||||
)
|
||||
conn.execute(
|
||||
f"""
|
||||
INSERT INTO chat_messages_fts(content, message_id, session_id, role)
|
||||
SELECT {fts_content_expr_cm}, cm.id, cm.session_id, cm.role
|
||||
FROM chat_messages cm
|
||||
LEFT JOIN temp._odysseus_fts_message_ids known ON known.message_id = cm.id
|
||||
WHERE known.message_id IS NULL
|
||||
"""
|
||||
)
|
||||
"""
|
||||
)
|
||||
_scrub_legacy_chat_message_fts_media(conn)
|
||||
conn.commit()
|
||||
except Exception as e:
|
||||
@@ -2565,6 +2778,27 @@ def _migrate_add_calendar_recurrence_exdates():
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_note_gallery_id():
|
||||
"""Keep a drawn note linked to one Gallery image across edits."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
columns = [row[1] for row in conn.execute("PRAGMA table_info(notes)").fetchall()]
|
||||
if columns and "gallery_id" not in columns:
|
||||
conn.execute("ALTER TABLE notes ADD COLUMN gallery_id VARCHAR")
|
||||
conn.execute("CREATE INDEX IF NOT EXISTS ix_notes_gallery_id ON notes(gallery_id)")
|
||||
conn.commit()
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"notes gallery_id migration failed: {e}")
|
||||
finally:
|
||||
if conn is not None:
|
||||
conn.close()
|
||||
|
||||
|
||||
def get_db():
|
||||
"""
|
||||
Dependency to get a database session.
|
||||
|
||||
@@ -108,6 +108,12 @@ 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.headers is None:
|
||||
@@ -155,6 +161,24 @@ class Session:
|
||||
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
|
||||
|
||||
|
||||
+21
-3
@@ -150,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
|
||||
@@ -208,6 +214,12 @@ 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,
|
||||
)
|
||||
|
||||
# The rows just loaded are the whole transcript, so they — not the
|
||||
@@ -485,6 +497,7 @@ class SessionManager:
|
||||
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)
|
||||
@@ -545,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(
|
||||
@@ -556,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)
|
||||
)
|
||||
@@ -570,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
|
||||
|
||||
Reference in New Issue
Block a user