File size: 4,274 Bytes
3a6c182
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
"""

Shared LLM + Zep API key slot (same index for both) with rate-limit handling:

try primary → on limit try secondary → on limit sleep 60s → retry from primary once more.

"""

from __future__ import annotations

import contextvars
import time
from typing import Any, Callable, TypeVar

from ..config import Config
from ..utils.logger import get_logger

logger = get_logger("mirofish.api_key_runtime")

T = TypeVar("T")

_key_slot: contextvars.ContextVar[int] = contextvars.ContextVar("miro_api_key_slot", default=0)


def reset_pipeline_api_keys() -> None:
    """Call at the start of each full pipeline run (worker thread)."""
    _key_slot.set(0)


def current_key_slot() -> int:
    return _key_slot.get()


def _max_key_slots() -> int:
    lk = Config.get_llm_api_keys()
    zk = Config.get_zep_api_keys()
    n = max(len(lk or []), len(zk or []), 1)
    return n


def set_key_slot(i: int) -> None:
    n = _max_key_slots()
    _key_slot.set(max(0, min(i, n - 1)))


def get_llm_key() -> str | None:
    keys = Config.get_llm_api_keys()
    if not keys:
        return None
    i = min(_key_slot.get(), len(keys) - 1)
    return keys[i]


def get_zep_key() -> str | None:
    keys = Config.get_zep_api_keys()
    if not keys:
        return None
    i = min(_key_slot.get(), len(keys) - 1)
    return keys[i]


def is_rate_limit_error(exc: BaseException) -> bool:
    try:
        import openai

        if isinstance(exc, getattr(openai, "RateLimitError", ())):
            return True
        if isinstance(exc, getattr(openai, "APIStatusError", ())):
            if getattr(exc, "status_code", None) == 429:
                return True
    except Exception:
        pass

    st = getattr(exc, "status_code", None)
    if st == 429:
        return True

    msg = str(exc).lower()
    for phrase in (
        "rate limit",
        "429",
        "too many requests",
        "quota",
        "resource_exhausted",
        "resource exhausted",
        "throttl",
        "over capacity",
        "capacity",
    ):
        if phrase in msg:
            return True
    return False


def is_auth_or_invalid_key_error(exc: BaseException) -> bool:
    """401 / invalid API key — rotate to next paired slot when LLM_2/ZEP_2 exist."""
    try:
        import openai

        if isinstance(exc, getattr(openai, "AuthenticationError", ())):
            return True
        if isinstance(exc, getattr(openai, "APIStatusError", ())):
            if getattr(exc, "status_code", None) == 401:
                return True
    except Exception:
        pass

    if getattr(exc, "status_code", None) == 401:
        return True

    msg = str(exc).lower()
    if "invalid api key" in msg or "invalid_api_key" in msg:
        return True
    if "401" in msg and ("unauthorized" in msg or "authentication" in msg):
        return True
    return False


def should_rotate_api_key(exc: BaseException) -> bool:
    return is_rate_limit_error(exc) or is_auth_or_invalid_key_error(exc)


def call_with_limit_rotation(op: Callable[[], T]) -> T:
    """

    Run op() using keys at the current slot. On rate-limit:

    advance to next key slot if available; else sleep 60s, reset to slot 0, try again;

    if still rate-limited after that full cycle, re-raise.

    """
    cooldown_used = False
    last_exc: BaseException | None = None

    while True:
        try:
            return op()
        except BaseException as e:
            if not should_rotate_api_key(e):
                raise
            last_exc = e
            idx = current_key_slot()
            n = _max_key_slots()
            if idx + 1 < n:
                set_key_slot(idx + 1)
                logger.warning(
                    "Rate limit hit; switching API key slot to %s (LLM+Zep)", current_key_slot()
                )
                continue
            if not cooldown_used:
                logger.warning("Rate limit on all API key slots; sleeping 60s then retrying from slot 0")
                time.sleep(60)
                cooldown_used = True
                set_key_slot(0)
                continue
            raise last_exc from None