Coverage for src/chat_limiter/types.py: 96%

80 statements  

« prev     ^ index     » next       coverage.py v7.9.2, created at 2025-09-12 11:48 +0100

1""" 

2Type definitions for chat completion requests and responses. 

3""" 

4 

5from dataclasses import dataclass 

6from enum import Enum 

7from typing import Any 

8 

9from pydantic import BaseModel 

10 

11 

12class MessageRole(str, Enum): 

13 """Message roles supported across providers.""" 

14 

15 USER = "user" 

16 ASSISTANT = "assistant" 

17 SYSTEM = "system" 

18 

19 

20@dataclass 

21class Message: 

22 """A chat message that works across all providers.""" 

23 

24 role: MessageRole 

25 content: str 

26 

27 

28class ChatCompletionRequest(BaseModel): 

29 """High-level chat completion request.""" 

30 

31 model: str 

32 messages: list[Message] 

33 max_tokens: int | None = None 

34 temperature: float | None = None 

35 top_p: float | None = None 

36 stop: str | list[str] | None = None 

37 stream: bool = False 

38 seed: int | None = None 

39 

40 # Provider-specific parameters (will be filtered per provider) 

41 frequency_penalty: float | None = None # OpenAI 

42 presence_penalty: float | None = None # OpenAI 

43 top_k: int | None = None # Anthropic 

44 reasoning_effort: str | None = None # OpenAI/OpenRouter reasoning models 

45 providers: list[str] | None = None # OpenRouter provider routing 

46 

47 

48@dataclass 

49class Usage: 

50 """Token usage information.""" 

51 

52 prompt_tokens: int 

53 completion_tokens: int 

54 total_tokens: int 

55 

56 

57@dataclass 

58class Choice: 

59 """A completion choice.""" 

60 

61 index: int 

62 message: Message 

63 finish_reason: str | None = None 

64 

65 

66@dataclass 

67class ChatCompletionResponse: 

68 """High-level chat completion response.""" 

69 

70 id: str 

71 model: str 

72 choices: list[Choice] 

73 usage: Usage | None = None 

74 created: int | None = None 

75 

76 # Error information 

77 success: bool = True 

78 error_message: str | None = None 

79 

80 # Provider-specific metadata 

81 provider: str | None = None 

82 raw_response: dict[str, Any] | None = None 

83 

84 

85# Model mappings for each provider 

86OPENAI_MODELS = { 

87 "gpt-4o", 

88 "gpt-4o-mini", 

89 "gpt-4-turbo", 

90 "gpt-4", 

91 "gpt-3.5-turbo", 

92 "gpt-3.5-turbo-16k", 

93} 

94 

95ANTHROPIC_MODELS = { 

96 "claude-3-5-sonnet-20241022", 

97 "claude-3-5-haiku-20241022", 

98 "claude-3-opus-20240229", 

99 "claude-3-sonnet-20240229", 

100 "claude-3-haiku-20240307", 

101} 

102 

103OPENROUTER_MODELS = { 

104 # OpenAI models via OpenRouter 

105 "openai/gpt-4o", 

106 "openai/gpt-4o-mini", 

107 "openai/gpt-4-turbo", 

108 "openai/gpt-3.5-turbo", 

109 

110 # Anthropic models via OpenRouter 

111 "anthropic/claude-3-5-sonnet", 

112 "anthropic/claude-3-opus", 

113 "anthropic/claude-3-sonnet", 

114 "anthropic/claude-3-haiku", 

115 

116 # Other providers via OpenRouter 

117 "meta-llama/llama-3.1-405b-instruct", 

118 "meta-llama/llama-3.1-70b-instruct", 

119 "google/gemini-pro", 

120 "cohere/command-r-plus", 

121} 

122 

123ALL_MODELS = OPENAI_MODELS | ANTHROPIC_MODELS | OPENROUTER_MODELS 

124 

125 

126def detect_provider_from_model(model: str, use_dynamic_discovery: bool = False, api_keys: dict[str, str] | None = None) -> str | None: 

127 """ 

128 Detect provider from model name. 

129 

130 Args: 

131 model: The model name to check 

132 use_dynamic_discovery: Whether to use live API queries for model discovery 

133 api_keys: Dictionary of API keys for dynamic discovery 

134 

135 Returns: 

136 Provider name or None if not found 

137 """ 

138 # Handle provider-prefixed models (e.g., "openai/o3", "anthropic/claude-3-sonnet") 

139 preferred_provider = None 

140 base_model = model 

141 

142 if "/" in model: 

143 parts = model.split("/", 1) 

144 if len(parts) == 2: 

145 provider_prefix, base_model = parts 

146 if provider_prefix == "openai": 

147 preferred_provider = "openai" 

148 elif provider_prefix == "anthropic": 

149 preferred_provider = "anthropic" 

150 

151 # If we have a preferred provider, check if the base model exists in hardcoded lists 

152 if preferred_provider: 

153 if preferred_provider == "openai" and base_model in OPENAI_MODELS: 

154 return "openai" 

155 elif preferred_provider == "anthropic" and base_model in ANTHROPIC_MODELS: 

156 return "anthropic" 

157 # If base model not found in preferred provider, fall back to checking if full model is in OpenRouter 

158 elif model in OPENROUTER_MODELS: 

159 return "openrouter" 

160 

161 # Check hardcoded lists for fast lookup (for models without provider prefix) 

162 if model in OPENAI_MODELS: 

163 return "openai" 

164 elif model in ANTHROPIC_MODELS: 

165 return "anthropic" 

166 elif model in OPENROUTER_MODELS: 

167 return "openrouter" 

168 

169 # If dynamic discovery is enabled and we have API keys, try that 

170 if use_dynamic_discovery and api_keys: 

171 from .models import detect_provider_from_model_sync 

172 result = detect_provider_from_model_sync(model, api_keys) 

173 return result.found_provider 

174 

175 return None