mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
4f798c0d23
commit
c3932e06d4
6 changed files with 173 additions and 74 deletions
|
|
@ -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) :]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,),
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue