fix(oauth): run the OAuth worker on the proxy's main asyncio loop

Reported error during device-code login:

  Tokens obtained but DB persist failed: <asyncio.locks.Event object at
  0x... [unset]> is bound to a different event loop

Root cause: the background worker was a ``threading.Thread`` that ran
``asyncio.run(persist_credential_to_db(item))``. That creates a fresh
event loop in the worker thread, but ``prisma_client``'s internal
asyncio primitives (locks, futures) are bound to the proxy's main loop.
Awaiting a prisma call from the worker loop hits the cross-loop error.

Fix:

- Refactor the background flow from a thread to an asyncio task
  (``asyncio.create_task``) scheduled on the proxy's main loop. Blocking
  IO in the flow (device-code poll, token exchange) runs in
  ``loop.run_in_executor``. The DB persist step (``await
  persist_credential_to_db(item)``) now naturally shares a loop with
  prisma_client. Same treatment for both ChatGPT and Copilot endpoints.

- For the DBAuthenticator refresh path (called from sync
  ``validate_environment`` during a request, same cross-loop risk), add
  ``_register_proxy_main_loop`` + ``_schedule_db_persist`` now prefers
  ``asyncio.run_coroutine_threadsafe`` onto the registered loop.
  ``/chatgpt/oauth/start`` and ``/copilot/oauth/start`` register the loop
  on invocation (guaranteed to be the proxy's main loop). Falls back to
  the old thread + ``asyncio.run`` path only when no loop is registered
  (CLI / tests, where prisma_client is ``None`` anyway).

Tests updated: ``TestBackgroundWorker`` now drives the async task
directly via ``await _run_device_code_flow_async(...)``. The
``test_creates_session_and_spawns_worker`` test stubs the task to a
no-op instead of patching ``threading.Thread``.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Jason Cook 2026-04-23 10:48:59 -04:00
parent 4f798c0d23
commit c3932e06d4
6 changed files with 173 additions and 74 deletions

View file

@ -108,6 +108,21 @@ def _unpack_auth_record(values: Dict[str, Any]) -> Dict[str, Any]:
# Fire-and-forget DB persistence
# ---------------------------------------------------------------------------
# The proxy's main asyncio event loop, captured by the OAuth start endpoint.
# ``prisma_client``'s internal primitives (locks, futures) are bound to this
# loop; awaiting them from a new loop spun up by ``asyncio.run`` in a worker
# thread raises "bound to a different event loop". We schedule the persist
# coroutine back onto this loop via ``run_coroutine_threadsafe`` from the
# refresh path (which runs in a FastAPI threadpool thread, not on any loop).
_proxy_main_loop: Optional["asyncio.AbstractEventLoop"] = None
def _register_proxy_main_loop(loop: "asyncio.AbstractEventLoop") -> None:
"""Called from the OAuth ``/start`` endpoint (guaranteed to be on the
proxy's main loop) so the refresh-path persist can find the loop."""
global _proxy_main_loop
_proxy_main_loop = loop
def _schedule_db_persist(item: CredentialItem) -> None:
"""
@ -119,6 +134,14 @@ def _schedule_db_persist(item: CredentialItem) -> None:
lags or fails. Failures are logged; the cache will be reconciled on the
next successful write.
"""
# If the proxy's main loop is known and running, schedule directly on it
# (cross-loop-safe for prisma_client's bound primitives). Otherwise fall
# back to spawning a thread + asyncio.run for non-proxy contexts (CLI,
# tests) where no prisma_client is set anyway.
loop = _proxy_main_loop
if loop is not None and loop.is_running():
asyncio.run_coroutine_threadsafe(persist_credential_to_db(item), loop)
return
thread = threading.Thread(
target=_persist_item_sync,
args=(item,),
@ -203,9 +226,7 @@ def resolve_authenticator(
if isinstance(litellm_params, dict)
else getattr(litellm_params, "api_key", None)
)
if isinstance(api_key, str) and api_key.startswith(
OAUTH_CREDENTIAL_API_KEY_PREFIX
):
if isinstance(api_key, str) and api_key.startswith(OAUTH_CREDENTIAL_API_KEY_PREFIX):
return DBAuthenticator(
credential_name=api_key[len(OAUTH_CREDENTIAL_API_KEY_PREFIX) :]
)

View file

@ -136,8 +136,23 @@ class DBAuthenticator(Authenticator):
# Fire-and-forget DB persistence
# ---------------------------------------------------------------------------
# See chatgpt/db_authenticator.py for the full rationale: prisma_client's
# internal asyncio primitives are bound to the proxy's main loop, so we
# must schedule persist coroutines back onto that loop rather than spin
# up a fresh one in a worker thread.
_proxy_main_loop: Optional["asyncio.AbstractEventLoop"] = None
def _register_proxy_main_loop(loop: "asyncio.AbstractEventLoop") -> None:
global _proxy_main_loop
_proxy_main_loop = loop
def _schedule_db_persist(item: CredentialItem) -> None:
loop = _proxy_main_loop
if loop is not None and loop.is_running():
asyncio.run_coroutine_threadsafe(persist_credential_to_db(item), loop)
return
thread = threading.Thread(
target=_persist_item_sync,
args=(item,),

View file

@ -37,6 +37,7 @@ from litellm.llms.chatgpt.common_utils import (
from litellm.llms.chatgpt.db_authenticator import (
CREDENTIAL_TYPE,
DBAuthenticator,
_register_proxy_main_loop,
persist_credential_to_db,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
@ -141,9 +142,18 @@ async def start_oauth(
"cancelled": False,
}
# Capture the proxy's main event loop so the DBAuthenticator refresh
# path can schedule its fire-and-forget persist safely across threads
# (prevents the cross-event-loop errors you hit if you `asyncio.run`
# an async prisma call from a worker thread).
loop = asyncio.get_running_loop()
_register_proxy_main_loop(loop)
authenticator = Authenticator()
try:
device_code = authenticator._request_device_code()
device_code = await loop.run_in_executor(
None, authenticator._request_device_code
)
except ChatGPTAuthError as exc:
with _sessions_lock:
_sessions.pop(session_id, None)
@ -157,13 +167,15 @@ async def start_oauth(
entry["status"] = "pending"
entry["device_code"] = device_code
thread = threading.Thread(
target=_run_device_code_flow,
args=(session_id, body.credential_name, device_code, authenticator),
daemon=True,
name=f"chatgpt-oauth-{session_id[:8]}",
# Run the remaining poll → exchange → persist flow as a task on the main
# loop (was previously a thread doing ``asyncio.run``, which built a new
# loop and collided with prisma_client's main-loop-bound primitives).
# Blocking IO inside the task runs through ``loop.run_in_executor``.
asyncio.create_task(
_run_device_code_flow_async(
session_id, body.credential_name, device_code, authenticator
)
)
thread.start()
return StartResponse(
session_id=session_id,
@ -251,22 +263,29 @@ async def oauth_refresh(
)
def _run_device_code_flow(
async def _run_device_code_flow_async(
session_id: str,
credential_name: str,
device_code: Dict[str, str],
authenticator: Authenticator,
) -> None:
"""
Background worker: polls for the authorization code, exchanges it for
tokens, then upserts the credential into the in-memory cache and DB.
Background async task: polls for the authorization code, exchanges it
for tokens, then upserts the credential into the in-memory cache and
DB. Runs on the proxy's main event loop so the DB persist step shares
a loop with ``prisma_client``.
"""
loop = asyncio.get_running_loop()
try:
auth_code = authenticator._poll_for_authorization_code(device_code)
auth_code = await loop.run_in_executor(
None, authenticator._poll_for_authorization_code, device_code
)
session_snapshot = _get_session(session_id)
if session_snapshot is None or session_snapshot.get("cancelled"):
return
tokens = authenticator._exchange_code_for_tokens(auth_code)
tokens = await loop.run_in_executor(
None, authenticator._exchange_code_for_tokens, auth_code
)
except ChatGPTAuthError as exc:
_update_session(session_id, status="error", message=exc.message)
return
@ -288,7 +307,7 @@ def _run_device_code_flow(
CredentialAccessor.upsert_credentials([item])
try:
asyncio.run(persist_credential_to_db(item))
await persist_credential_to_db(item)
except Exception as exc:
verbose_proxy_logger.error(
"Failed to persist ChatGPT OAuth credential %s: %s",

View file

@ -15,6 +15,7 @@ Endpoints:
POST /copilot/oauth/refresh — force a Copilot API key refresh for a stored credential
"""
import asyncio
import threading
import time
import uuid
@ -24,11 +25,18 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.llms.github_copilot.authenticator import Authenticator
from litellm.llms.github_copilot.common_utils import GithubCopilotError
from litellm.llms.github_copilot.db_authenticator import DBAuthenticator
from litellm.llms.github_copilot.db_authenticator import (
CREDENTIAL_TYPE,
DBAuthenticator,
_register_proxy_main_loop,
persist_credential_to_db,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.utils import CredentialItem
router = APIRouter(prefix="/copilot/oauth", tags=["copilot oauth"])
@ -128,9 +136,16 @@ async def start_oauth(
"cancelled": False,
}
# Capture the proxy's main event loop so the DBAuthenticator refresh
# path can schedule persist work cross-thread via run_coroutine_threadsafe.
loop = asyncio.get_running_loop()
_register_proxy_main_loop(loop)
authenticator = Authenticator()
try:
device_code_info = authenticator._get_device_code()
device_code_info = await loop.run_in_executor(
None, authenticator._get_device_code
)
except GithubCopilotError as exc:
with _sessions_lock:
_sessions.pop(session_id, None)
@ -143,13 +158,13 @@ async def start_oauth(
entry["status"] = "pending"
entry["device_code_info"] = device_code_info
thread = threading.Thread(
target=_run_device_code_flow,
args=(session_id, body.credential_name, device_code_info, authenticator),
daemon=True,
name=f"copilot-oauth-{session_id[:8]}",
# Run the remaining flow as a task on the main loop so the DB persist
# step shares an event loop with prisma_client.
asyncio.create_task(
_run_device_code_flow_async(
session_id, body.credential_name, device_code_info, authenticator
)
)
thread.start()
return StartResponse(
session_id=session_id,
@ -228,19 +243,23 @@ async def oauth_refresh(
)
def _run_device_code_flow(
async def _run_device_code_flow_async(
session_id: str,
credential_name: str,
device_code_info: Dict[str, str],
authenticator: Authenticator,
) -> None:
"""
Background worker: polls GitHub for the access token, then persists it
via :class:`DBAuthenticator.store_access_token`.
Background async task: polls GitHub for the access token, then upserts
the credential into the in-memory cache and awaits the DB persist on
the same event loop as ``prisma_client``.
"""
loop = asyncio.get_running_loop()
try:
access_token = authenticator._poll_for_access_token(
device_code_info["device_code"]
access_token = await loop.run_in_executor(
None,
authenticator._poll_for_access_token,
device_code_info["device_code"],
)
session_snapshot = _get_session(session_id)
if session_snapshot is None or session_snapshot.get("cancelled"):
@ -253,10 +272,20 @@ def _run_device_code_flow(
_update_session(session_id, status="error", message=str(exc))
return
item = CredentialItem(
credential_name=credential_name,
credential_values={"access_token": access_token},
credential_info={
"type": CREDENTIAL_TYPE,
"custom_llm_provider": "github_copilot",
},
)
CredentialAccessor.upsert_credentials([item])
# Invalidate any cached Copilot API key tied to an old access token.
DBAuthenticator._api_key_cache.pop(credential_name, None)
try:
DBAuthenticator(credential_name=credential_name).store_access_token(
access_token
)
await persist_credential_to_db(item)
except Exception as exc:
verbose_proxy_logger.error(
"Failed to persist Copilot OAuth credential %s: %s",

View file

@ -117,12 +117,10 @@ class TestStartOAuth:
"interval": "5",
}
class _FakeThread:
def __init__(self, *args, **kwargs):
self.started = False
def start(self):
self.started = True
# Stub the background async task so start_oauth doesn't actually run
# the poll → exchange → persist flow during the test.
async def _noop(*args, **kwargs):
return None
with (
patch.object(
@ -131,8 +129,8 @@ class TestStartOAuth:
return_value=fake_device_code,
),
patch(
"litellm.proxy.chatgpt_oauth_endpoints.endpoints.threading.Thread",
_FakeThread,
"litellm.proxy.chatgpt_oauth_endpoints.endpoints._run_device_code_flow_async",
_noop,
),
):
response = await start_oauth(
@ -321,9 +319,10 @@ class TestRefreshEndpoint:
class TestBackgroundWorker:
def test_worker_marks_success_and_persists(self, monkeypatch):
@pytest.mark.asyncio
async def test_worker_marks_success_and_persists(self, monkeypatch):
from litellm.proxy.chatgpt_oauth_endpoints.endpoints import (
_run_device_code_flow,
_run_device_code_flow_async,
)
session_id = "s1"
@ -365,7 +364,7 @@ class TestBackgroundWorker:
)
monkeypatch.setattr(litellm, "credential_list", [])
_run_device_code_flow(
await _run_device_code_flow_async(
session_id=session_id,
credential_name="my-creds",
device_code={"interval": "5"},
@ -377,17 +376,17 @@ class TestBackgroundWorker:
persist_mock.assert_called_once()
persisted_item = persist_mock.call_args.args[0]
assert persisted_item.credential_name == "my-creds"
# Verify the item is in the in-memory cache too
assert any(c.credential_name == "my-creds" for c in litellm.credential_list)
def test_worker_marks_error_on_db_persist_failure(self, monkeypatch):
@pytest.mark.asyncio
async def test_worker_marks_error_on_db_persist_failure(self, monkeypatch):
"""
If tokens are obtained but the DB write fails, the session should
flip to ``error`` with an informative message (the in-memory cache
was already updated; next retry via UI will retry the DB write).
"""
from litellm.proxy.chatgpt_oauth_endpoints.endpoints import (
_run_device_code_flow,
_run_device_code_flow_async,
)
session_id = "s1"
@ -427,7 +426,7 @@ class TestBackgroundWorker:
)
monkeypatch.setattr(litellm, "credential_list", [])
_run_device_code_flow(
await _run_device_code_flow_async(
session_id=session_id,
credential_name="my-creds",
device_code={"interval": "5"},
@ -438,10 +437,11 @@ class TestBackgroundWorker:
assert _sessions[session_id]["status"] == "error"
assert "DB persist failed" in _sessions[session_id]["message"]
def test_worker_marks_error_on_auth_failure(self, monkeypatch):
@pytest.mark.asyncio
async def test_worker_marks_error_on_auth_failure(self, monkeypatch):
from litellm.llms.chatgpt.common_utils import GetAccessTokenError
from litellm.proxy.chatgpt_oauth_endpoints.endpoints import (
_run_device_code_flow,
_run_device_code_flow_async,
)
session_id = "s1"
@ -458,7 +458,7 @@ class TestBackgroundWorker:
status_code=408, message="timed out"
)
_run_device_code_flow(
await _run_device_code_flow_async(
session_id=session_id,
credential_name="my-creds",
device_code={"interval": "5"},

View file

@ -109,12 +109,8 @@ class TestStartOAuth:
"interval": 5,
}
class _FakeThread:
def __init__(self, *args, **kwargs):
pass
def start(self):
pass
async def _noop(*args, **kwargs):
return None
with (
patch.object(
@ -123,8 +119,8 @@ class TestStartOAuth:
return_value=fake_device_code,
),
patch(
"litellm.proxy.copilot_oauth_endpoints.endpoints.threading.Thread",
_FakeThread,
"litellm.proxy.copilot_oauth_endpoints.endpoints._run_device_code_flow_async",
_noop,
),
):
response = await start_oauth(
@ -227,9 +223,11 @@ class TestRefreshEndpoint:
class TestBackgroundWorker:
def test_worker_persists_access_token_on_success(self, monkeypatch):
@pytest.mark.asyncio
async def test_worker_persists_access_token_on_success(self, monkeypatch):
import litellm
from litellm.proxy.copilot_oauth_endpoints.endpoints import (
_run_device_code_flow,
_run_device_code_flow_async,
)
session_id = "s1"
@ -244,25 +242,37 @@ class TestBackgroundWorker:
auth = MagicMock()
auth._poll_for_access_token.return_value = "gho_fresh_token"
store_mock = MagicMock()
monkeypatch.setattr(
DBAuthenticator, "store_access_token", lambda self, tok: store_mock(tok)
)
persist_mock = MagicMock()
_run_device_code_flow(
async def _fake_persist(item):
persist_mock(item)
monkeypatch.setattr(
"litellm.proxy.copilot_oauth_endpoints.endpoints.persist_credential_to_db",
_fake_persist,
)
monkeypatch.setattr(litellm, "credential_list", [])
await _run_device_code_flow_async(
session_id=session_id,
credential_name="c",
device_code_info={"device_code": "dc", "user_code": "UC"},
authenticator=auth,
)
store_mock.assert_called_once_with("gho_fresh_token")
persist_mock.assert_called_once()
assert persist_mock.call_args.args[0].credential_name == "c"
assert persist_mock.call_args.args[0].credential_values["access_token"] == (
"gho_fresh_token"
)
with _sessions_lock:
assert _sessions[session_id]["status"] == "success"
def test_worker_marks_error_on_db_persist_failure(self, monkeypatch):
@pytest.mark.asyncio
async def test_worker_marks_error_on_db_persist_failure(self, monkeypatch):
import litellm
from litellm.proxy.copilot_oauth_endpoints.endpoints import (
_run_device_code_flow,
_run_device_code_flow_async,
)
session_id = "s1"
@ -277,12 +287,16 @@ class TestBackgroundWorker:
auth = MagicMock()
auth._poll_for_access_token.return_value = "gho_fresh"
def _boom(self, tok):
async def _boom(item):
raise RuntimeError("prisma disconnected")
monkeypatch.setattr(DBAuthenticator, "store_access_token", _boom)
monkeypatch.setattr(
"litellm.proxy.copilot_oauth_endpoints.endpoints.persist_credential_to_db",
_boom,
)
monkeypatch.setattr(litellm, "credential_list", [])
_run_device_code_flow(
await _run_device_code_flow_async(
session_id=session_id,
credential_name="c",
device_code_info={"device_code": "dc", "user_code": "UC"},
@ -293,9 +307,10 @@ class TestBackgroundWorker:
assert _sessions[session_id]["status"] == "error"
assert "DB persist failed" in _sessions[session_id]["message"]
def test_worker_marks_error_on_poll_failure(self):
@pytest.mark.asyncio
async def test_worker_marks_error_on_poll_failure(self):
from litellm.proxy.copilot_oauth_endpoints.endpoints import (
_run_device_code_flow,
_run_device_code_flow_async,
)
session_id = "s1"
@ -312,7 +327,7 @@ class TestBackgroundWorker:
status_code=408, message="timed out"
)
_run_device_code_flow(
await _run_device_code_flow_async(
session_id=session_id,
credential_name="c",
device_code_info={"device_code": "dc", "user_code": "UC"},