File size: 11,135 Bytes
1def50b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
# ============================================================================
#  API KEY POOL MANAGER
# ============================================================================

import json
import random
import hashlib
import asyncio
from datetime import datetime, timedelta
from typing import Dict, Any, List, Optional
from pathlib import Path

from app.config import Config
from app.utils.logger import logger


class KeyStats:
    """Statistics for a single API key"""
    def __init__(self):
        self.rpm_used = 0
        self.rph_used = 0
        self.rpd_used = 0
        self.tpm_used = 0
        self.tpd_used = 0
        self.total_requests = 0
        self.total_tokens = 0
        self.last_used = None
        self.rpm_reset_at = datetime.now() + timedelta(minutes=1)
        self.rph_reset_at = datetime.now() + timedelta(hours=1)
        self.rpd_reset_at = datetime.now() + timedelta(days=1)
        self.tpm_reset_at = datetime.now() + timedelta(minutes=1)
        self.tpd_reset_at = datetime.now() + timedelta(days=1)


class APIKeyPool:
    """Manages multiple Groq API keys with shuffle rotation and usage tracking"""
    
    USAGE_FILE = Path(Config.OUTPUT_DIR) / "api_usage.json"
    
    def __init__(self, keys: List[str]):
        if not keys:
            raise ValueError("At least one API key is required")
        
        self.keys = keys
        self.key_aliases = {self._hash_key(key): f"key_{i+1}" for i, key in enumerate(keys)}
        self.key_stats: Dict[str, KeyStats] = {key: KeyStats() for key in keys}
        self._lock = asyncio.Lock()
        
        # Load existing usage data
        self._load_usage_data()
        
        logger.info(f"πŸ” API Key Pool initialized with {len(keys)} key(s)")
    
    def _hash_key(self, key: str) -> str:
        """Create SHA256 hash of API key for secure storage"""
        return hashlib.sha256(key.encode()).hexdigest()[:16]
    
    def _load_usage_data(self):
        """Load usage data from JSON file"""
        if self.USAGE_FILE.exists():
            try:
                with open(self.USAGE_FILE, 'r') as f:
                    data = json.load(f)
                
                for key in self.keys:
                    key_hash = self._hash_key(key)
                    if key_hash in data.get("keys", {}):
                        stored = data["keys"][key_hash]
                        stats = self.key_stats[key]
                        
                        # Load counters
                        stats.total_requests = stored.get("total_requests", 0)
                        stats.total_tokens = stored.get("total_tokens", 0)
                        stats.last_used = stored.get("last_used")
                        
                        # Load rate limit usage
                        limits = stored.get("limits", {})
                        stats.rpm_used = limits.get("rpm", {}).get("used", 0)
                        stats.rph_used = limits.get("rph", {}).get("used", 0)
                        stats.rpd_used = limits.get("rpd", {}).get("used", 0)
                        stats.tpm_used = limits.get("tpm", {}).get("used", 0)
                        stats.tpd_used = limits.get("tpd", {}).get("used", 0)
                
                logger.info(f"πŸ“Š Loaded usage data from {self.USAGE_FILE}")
            except Exception as e:
                logger.warning(f"Failed to load usage data: {e}")
    
    def _save_usage_data(self):
        """Save usage data to JSON file"""
        try:
            data = {
                "last_updated": datetime.now().isoformat(),
                "keys": {}
            }
            
            for key, stats in self.key_stats.items():
                key_hash = self._hash_key(key)
                alias = self.key_aliases[key_hash]
                
                data["keys"][key_hash] = {
                    "alias": alias,
                    "limits": {
                        "rpm": {"used": stats.rpm_used, "limit": Config.GROQ_RPM_LIMIT, "reset_at": stats.rpm_reset_at.isoformat()},
                        "rph": {"used": stats.rph_used, "limit": Config.GROQ_RPH_LIMIT, "reset_at": stats.rph_reset_at.isoformat()},
                        "rpd": {"used": stats.rpd_used, "limit": Config.GROQ_RPD_LIMIT, "reset_at": stats.rpd_reset_at.isoformat()},
                        "tpm": {"used": stats.tpm_used, "limit": Config.GROQ_TPM_LIMIT, "reset_at": stats.tpm_reset_at.isoformat()},
                        "tpd": {"used": stats.tpd_used, "limit": Config.GROQ_TPD_LIMIT, "reset_at": stats.tpd_reset_at.isoformat()}
                    },
                    "stats": {
                        "total_requests": stats.total_requests,
                        "total_tokens": stats.total_tokens,
                        "last_used": stats.last_used,
                        "is_available": self._is_key_available(key)
                    }
                }
            
            # Ensure directory exists
            self.USAGE_FILE.parent.mkdir(parents=True, exist_ok=True)
            
            with open(self.USAGE_FILE, 'w') as f:
                json.dump(data, f, indent=2, default=str)
            
        except Exception as e:
            logger.error(f"Failed to save usage data: {e}")
    
    def _is_key_available(self, key: str) -> bool:
        """Check if key has remaining quota for all limit types"""
        stats = self.key_stats[key]
        now = datetime.now()
        
        # Reset counters if window expired
        if now >= stats.rpm_reset_at:
            stats.rpm_used = 0
            stats.rpm_reset_at = now + timedelta(minutes=1)
        
        if now >= stats.rph_reset_at:
            stats.rph_used = 0
            stats.rph_reset_at = now + timedelta(hours=1)
        
        if now >= stats.rpd_reset_at:
            stats.rpd_used = 0
            stats.rpd_reset_at = now + timedelta(days=1)
        
        if now >= stats.tpm_reset_at:
            stats.tpm_used = 0
            stats.tpm_reset_at = now + timedelta(minutes=1)
        
        if now >= stats.tpd_reset_at:
            stats.tpd_used = 0
            stats.tpd_reset_at = now + timedelta(days=1)
        
        # Check all limits
        if stats.rpm_used >= Config.GROQ_RPM_LIMIT:
            return False
        if stats.rph_used >= Config.GROQ_RPH_LIMIT:
            return False
        if stats.rpd_used >= Config.GROQ_RPD_LIMIT:
            return False
        if stats.tpm_used >= Config.GROQ_TPM_LIMIT:
            return False
        if stats.tpd_used >= Config.GROQ_TPD_LIMIT:
            return False
        
        return True
    
    async def get_available_key(self, estimated_tokens: int = 2500) -> str:
        """Get an available API key using shuffle (random selection)
        
        Args:
            estimated_tokens: Estimated tokens needed for the request
            
        Returns:
            Available API key
            
        Raises:
            RuntimeError: If all keys are at rate limit
        """
        async with self._lock:
            # Filter available keys
            available_keys = [
                key for key in self.keys
                if self._is_key_available(key) and 
                   (self.key_stats[key].tpm_used + estimated_tokens) <= Config.GROQ_TPM_LIMIT and
                   (self.key_stats[key].tpd_used + estimated_tokens) <= Config.GROQ_TPD_LIMIT
            ]
            
            if not available_keys:
                # Find earliest reset time
                earliest_reset = min(
                    stats.tpd_reset_at for stats in self.key_stats.values()
                )
                wait_seconds = (earliest_reset - datetime.now()).total_seconds()
                
                logger.error(f"🚫 All API keys at rate limit. Next reset in {wait_seconds:.0f}s")
                raise RuntimeError(f"All API keys at rate limit. Try again in {wait_seconds:.0f} seconds.")
            
            # SHUFFLE - Random selection from available keys
            selected_key = random.choice(available_keys)
            alias = self.key_aliases[self._hash_key(selected_key)]
            
            logger.info(f"πŸ”€ Selected key: {alias} (from {len(available_keys)} available)")
            
            return selected_key
    
    async def record_usage(self, key: str, tokens_used: int, success: bool = True):
        """Record usage after API call
        
        Args:
            key: The API key used
            tokens_used: Total tokens consumed
            success: Whether the API call succeeded
        """
        async with self._lock:
            stats = self.key_stats[key]
            now = datetime.now()
            
            # Increment request counters
            stats.rpm_used += 1
            stats.rph_used += 1
            stats.rpd_used += 1
            
            # Increment token counters
            stats.tpm_used += tokens_used
            stats.tpd_used += tokens_used
            
            # Update totals
            stats.total_requests += 1
            stats.total_tokens += tokens_used
            stats.last_used = now.isoformat()
            
            # Save to JSON
            self._save_usage_data()
            
            alias = self.key_aliases[self._hash_key(key)]
            logger.info(f"πŸ“Š Recorded usage for {alias}: {tokens_used} tokens | RPM: {stats.rpm_used}/{Config.GROQ_RPM_LIMIT} | TPM: {stats.tpm_used}/{Config.GROQ_TPM_LIMIT}")
    
    def get_usage_summary(self) -> Dict[str, Any]:
        """Get usage summary for all keys"""
        summary = {
            "total_keys": len(self.keys),
            "available_keys": sum(1 for k in self.keys if self._is_key_available(k)),
            "total_requests_today": sum(s.rpd_used for s in self.key_stats.values()),
            "total_tokens_today": sum(s.tpd_used for s in self.key_stats.values()),
            "keys": {}
        }
        
        for key, stats in self.key_stats.items():
            key_hash = self._hash_key(key)
            alias = self.key_aliases[key_hash]
            
            summary["keys"][alias] = {
                "is_available": self._is_key_available(key),
                "rpm": {"used": stats.rpm_used, "limit": Config.GROQ_RPM_LIMIT},
                "rph": {"used": stats.rph_used, "limit": Config.GROQ_RPH_LIMIT},
                "rpd": {"used": stats.rpd_used, "limit": Config.GROQ_RPD_LIMIT},
                "tpm": {"used": stats.tpm_used, "limit": Config.GROQ_TPM_LIMIT},
                "tpd": {"used": stats.tpd_used, "limit": Config.GROQ_TPD_LIMIT},
                "total_requests": stats.total_requests,
                "total_tokens": stats.total_tokens,
                "last_used": stats.last_used
            }
        
        return summary
    
    def get_key_alias(self, key: str) -> str:
        """Get alias for a key (for logging/response headers)"""
        key_hash = self._hash_key(key)
        return self.key_aliases.get(key_hash, "unknown")


# Global instance
api_key_pool = APIKeyPool(Config.GROQ_API_KEYS) if Config.GROQ_API_KEYS else None