diff --git a/litellm/llms/chatgpt/db_authenticator.py b/litellm/llms/chatgpt/db_authenticator.py index 9d823120516..46f6630bb7c 100644 --- a/litellm/llms/chatgpt/db_authenticator.py +++ b/litellm/llms/chatgpt/db_authenticator.py @@ -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) :] ) diff --git a/litellm/llms/github_copilot/db_authenticator.py b/litellm/llms/github_copilot/db_authenticator.py index a2470f4320d..4d1ddf5e703 100644 --- a/litellm/llms/github_copilot/db_authenticator.py +++ b/litellm/llms/github_copilot/db_authenticator.py @@ -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,), diff --git a/litellm/proxy/chatgpt_oauth_endpoints/endpoints.py b/litellm/proxy/chatgpt_oauth_endpoints/endpoints.py index 144a1ef853c..a61aa48b661 100644 --- a/litellm/proxy/chatgpt_oauth_endpoints/endpoints.py +++ b/litellm/proxy/chatgpt_oauth_endpoints/endpoints.py @@ -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", diff --git a/litellm/proxy/copilot_oauth_endpoints/endpoints.py b/litellm/proxy/copilot_oauth_endpoints/endpoints.py index 061e2bea0f3..9f43376692b 100644 --- a/litellm/proxy/copilot_oauth_endpoints/endpoints.py +++ b/litellm/proxy/copilot_oauth_endpoints/endpoints.py @@ -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", diff --git a/tests/test_litellm/proxy/chatgpt_oauth_endpoints/test_chatgpt_oauth_endpoints.py b/tests/test_litellm/proxy/chatgpt_oauth_endpoints/test_chatgpt_oauth_endpoints.py index 761bfffab1d..6220160d70e 100644 --- a/tests/test_litellm/proxy/chatgpt_oauth_endpoints/test_chatgpt_oauth_endpoints.py +++ b/tests/test_litellm/proxy/chatgpt_oauth_endpoints/test_chatgpt_oauth_endpoints.py @@ -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"}, diff --git a/tests/test_litellm/proxy/copilot_oauth_endpoints/test_copilot_oauth_endpoints.py b/tests/test_litellm/proxy/copilot_oauth_endpoints/test_copilot_oauth_endpoints.py index 265c14d3b6f..bb3b1b7782c 100644 --- a/tests/test_litellm/proxy/copilot_oauth_endpoints/test_copilot_oauth_endpoints.py +++ b/tests/test_litellm/proxy/copilot_oauth_endpoints/test_copilot_oauth_endpoints.py @@ -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"},