mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-09 08:22:19 +02:00
feat(provider): support multiple ChatGPT subscriptions with usage
This commit is contained in:
@@ -203,6 +203,12 @@ class Session(TimestampMixin, Base):
|
||||
# Organization
|
||||
folder = Column(String, nullable=True, default=None)
|
||||
cwd = Column(String, nullable=True, default=None)
|
||||
# Registered ModelEndpoint this session is bound to. endpoint_url alone
|
||||
# cannot distinguish two endpoints that share a provider URL but use
|
||||
# different credentials (e.g. two ChatGPT Subscription accounts), so the
|
||||
# exact endpoint id is remembered here. NULL = legacy session; the first
|
||||
# deterministic, owner-scoped resolution persists a binding.
|
||||
endpoint_id = Column(String, nullable=True, index=True)
|
||||
|
||||
# Headers stored as JSON
|
||||
headers = Column(JSON, default=dict)
|
||||
@@ -1472,6 +1478,19 @@ def _migrate_add_session_cwd_column():
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _migrate_add_session_endpoint_id_column():
|
||||
"""Add the nullable binding and index without rewriting existing sessions."""
|
||||
with engine.begin() as connection:
|
||||
schema = inspect(connection)
|
||||
if not schema.has_table("sessions"):
|
||||
return
|
||||
columns = {column["name"] for column in schema.get_columns("sessions")}
|
||||
if "endpoint_id" not in columns:
|
||||
connection.execute(text("ALTER TABLE sessions ADD COLUMN endpoint_id VARCHAR"))
|
||||
index = next(index for index in Session.__table__.indexes if index.name == "ix_sessions_endpoint_id")
|
||||
index.create(bind=connection, checkfirst=True)
|
||||
|
||||
|
||||
def _migrate_add_token_columns():
|
||||
"""Add cumulative token tracking columns to sessions table."""
|
||||
import sqlite3
|
||||
@@ -2378,6 +2397,7 @@ def init_db():
|
||||
_migrate_add_session_generation_settings_columns()
|
||||
_migrate_add_folder_column()
|
||||
_migrate_add_session_cwd_column()
|
||||
_migrate_add_session_endpoint_id_column()
|
||||
_migrate_add_token_columns()
|
||||
_migrate_add_total_cost_usd()
|
||||
_migrate_add_mode_column()
|
||||
|
||||
@@ -115,6 +115,10 @@ class Session:
|
||||
temperature_override: Optional[float] = None
|
||||
max_tokens_override: Optional[int] = None
|
||||
cwd: Optional[str] = None
|
||||
# Registered ModelEndpoint id this session is bound to (None = legacy /
|
||||
# URL-matched). Lets two endpoints that share a provider URL but not
|
||||
# credentials stay distinguishable.
|
||||
endpoint_id: Optional[str] = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.headers is None:
|
||||
|
||||
@@ -157,6 +157,7 @@ class SessionManager:
|
||||
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,
|
||||
endpoint_id=getattr(db_session, "endpoint_id", None) or None,
|
||||
)
|
||||
session.message_count = getattr(db_session, "message_count", 0) or 0
|
||||
return session
|
||||
@@ -222,6 +223,7 @@ class SessionManager:
|
||||
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,
|
||||
endpoint_id=getattr(db_session, "endpoint_id", None) or None,
|
||||
)
|
||||
|
||||
# The rows just loaded are the whole transcript, so they — not the
|
||||
@@ -493,6 +495,7 @@ class SessionManager:
|
||||
headers = {}
|
||||
session.name = db_session.name
|
||||
session.endpoint_url = db_session.endpoint_url or ""
|
||||
session.endpoint_id = getattr(db_session, "endpoint_id", None) or None
|
||||
session.model = db_session.model or ""
|
||||
session.headers = headers or {}
|
||||
session.rag = db_session.rag
|
||||
@@ -563,9 +566,12 @@ class SessionManager:
|
||||
owner: str = None,
|
||||
cwd: str = None,
|
||||
headers: Optional[Dict[str, str]] = None,
|
||||
endpoint_id: Optional[str] = None,
|
||||
) -> Session:
|
||||
"""Create a new session and save to database."""
|
||||
session_headers = dict(headers or {})
|
||||
from src.chatgpt_subscription import is_chatgpt_subscription_base
|
||||
session_headers = {} if is_chatgpt_subscription_base(endpoint_url) else dict(headers or {})
|
||||
endpoint_id = (endpoint_id or "").strip() or None
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db_session = DbSession(
|
||||
@@ -577,6 +583,7 @@ class SessionManager:
|
||||
headers=session_headers,
|
||||
owner=owner,
|
||||
cwd=cwd or None,
|
||||
endpoint_id=endpoint_id,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc)
|
||||
)
|
||||
@@ -592,6 +599,7 @@ class SessionManager:
|
||||
headers=session_headers,
|
||||
owner=owner,
|
||||
cwd=cwd or None,
|
||||
endpoint_id=endpoint_id,
|
||||
)
|
||||
|
||||
self.sessions[session_id] = session
|
||||
|
||||
Reference in New Issue
Block a user