feat(provider): support multiple ChatGPT subscriptions with usage

This commit is contained in:
Alexandre Teixeira
2026-09-22 13:12:19 +01:00
parent 45330097b8
commit ed7ccfd584
31 changed files with 3281 additions and 215 deletions
+20
View File
@@ -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()
+4
View File
@@ -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:
+9 -1
View File
@@ -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