Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/client/chat.py: 0%

81 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import json 

2from collections.abc import Iterator 

3from typing import Any, Final 

4 

5import requests 

6 

7from .exceptions import UnauthorizedError 

8 

9 

10class ChatClient: 

11 def __init__(self, base_url: str, api_key: str | None = None, timeout: int = 600): 

12 """ 

13 Initialize the ChatClient. 

14 

15 Args: 

16 base_url (str): The base URL of the LiteLLM proxy server (e.g., "http://localhost:8000") 

17 api_key (Optional[str]): API key for authentication. If provided, it will be sent as a Bearer token. 

18 timeout (int): Request timeout in seconds (default: 600, the OpenAI SDK default, since a completion 

19 can legitimately take minutes) 

20 """ 

21 self._base_url = base_url.rstrip("/") # Remove trailing slash if present 

22 self._api_key = api_key 

23 self._timeout = timeout 

24 

25 def _get_headers(self) -> dict[str, str]: 

26 """ 

27 Get the headers for API requests, including authorization if api_key is set. 

28 

29 Returns: 

30 Dict[str, str]: Headers to use for API requests 

31 """ 

32 headers: Final = {"Content-Type": "application/json"} 

33 if self._api_key: 

34 headers["Authorization"] = f"Bearer {self._api_key}" 

35 return headers 

36 

37 def completions( 

38 self, 

39 model: str, 

40 messages: list[dict[str, str]], 

41 temperature: float | None = None, 

42 top_p: float | None = None, 

43 n: int | None = None, 

44 max_tokens: int | None = None, 

45 presence_penalty: float | None = None, 

46 frequency_penalty: float | None = None, 

47 user: str | None = None, 

48 return_request: bool = False, 

49 ) -> dict[str, Any] | requests.Request: 

50 """ 

51 Create a chat completion. 

52 

53 Args: 

54 model (str): The model to use for completion 

55 messages (List[Dict[str, str]]): The messages to generate a completion for 

56 temperature (Optional[float]): Sampling temperature between 0 and 2 

57 top_p (Optional[float]): Nucleus sampling parameter between 0 and 1 

58 n (Optional[int]): Number of completions to generate 

59 max_tokens (Optional[int]): Maximum number of tokens to generate 

60 presence_penalty (Optional[float]): Presence penalty between -2.0 and 2.0 

61 frequency_penalty (Optional[float]): Frequency penalty between -2.0 and 2.0 

62 user (Optional[str]): Unique identifier for the end user 

63 return_request (bool): If True, returns the prepared request object instead of executing it 

64 

65 Returns: 

66 Union[Dict[str, Any], requests.Request]: Either the completion response from the server or 

67 a prepared request object if return_request is True 

68 

69 Raises: 

70 UnauthorizedError: If the request fails with a 401 status code 

71 requests.exceptions.RequestException: If the request fails with any other error 

72 """ 

73 url: Final = f"{self._base_url}/chat/completions" 

74 

75 # Build request data with required fields 

76 data: Final[dict[str, object]] = {"model": model, "messages": messages} 

77 

78 # Add optional parameters if provided 

79 if temperature is not None: 

80 data["temperature"] = temperature 

81 if top_p is not None: 

82 data["top_p"] = top_p 

83 if n is not None: 

84 data["n"] = n 

85 if max_tokens is not None: 

86 data["max_tokens"] = max_tokens 

87 if presence_penalty is not None: 

88 data["presence_penalty"] = presence_penalty 

89 if frequency_penalty is not None: 

90 data["frequency_penalty"] = frequency_penalty 

91 if user is not None: 

92 data["user"] = user 

93 

94 request: Final = requests.Request("POST", url, headers=self._get_headers(), json=data) 

95 

96 if return_request: 

97 return request 

98 

99 # Prepare and send the request 

100 session: Final = requests.Session() 

101 try: 

102 response: Final = session.send(request.prepare(), timeout=self._timeout) 

103 response.raise_for_status() 

104 return response.json() 

105 except requests.exceptions.HTTPError as e: 

106 if e.response.status_code == 401: 

107 raise UnauthorizedError(e) 

108 raise 

109 

110 def completions_stream( 

111 self, 

112 model: str, 

113 messages: list[dict[str, str]], 

114 temperature: float | None = None, 

115 top_p: float | None = None, 

116 n: int | None = None, 

117 max_tokens: int | None = None, 

118 presence_penalty: float | None = None, 

119 frequency_penalty: float | None = None, 

120 user: str | None = None, 

121 ) -> Iterator[dict[str, Any]]: 

122 """ 

123 Create a streaming chat completion. 

124 

125 Args: 

126 model (str): The model to use for completion 

127 messages (List[Dict[str, str]]): The messages to generate a completion for 

128 temperature (Optional[float]): Sampling temperature between 0 and 2 

129 top_p (Optional[float]): Nucleus sampling parameter between 0 and 1 

130 n (Optional[int]): Number of completions to generate 

131 max_tokens (Optional[int]): Maximum number of tokens to generate 

132 presence_penalty (Optional[float]): Presence penalty between -2.0 and 2.0 

133 frequency_penalty (Optional[float]): Frequency penalty between -2.0 and 2.0 

134 user (Optional[str]): Unique identifier for the end user 

135 

136 Yields: 

137 Dict[str, Any]: Streaming response chunks from the server 

138 

139 Raises: 

140 UnauthorizedError: If the request fails with a 401 status code 

141 requests.exceptions.RequestException: If the request fails with any other error 

142 """ 

143 url: Final = f"{self._base_url}/chat/completions" 

144 

145 # Build request data with required fields 

146 data: Final[dict[str, object]] = {"model": model, "messages": messages, "stream": True} 

147 

148 # Add optional parameters if provided 

149 if temperature is not None: 

150 data["temperature"] = temperature 

151 if top_p is not None: 

152 data["top_p"] = top_p 

153 if n is not None: 

154 data["n"] = n 

155 if max_tokens is not None: 

156 data["max_tokens"] = max_tokens 

157 if presence_penalty is not None: 

158 data["presence_penalty"] = presence_penalty 

159 if frequency_penalty is not None: 

160 data["frequency_penalty"] = frequency_penalty 

161 if user is not None: 

162 data["user"] = user 

163 

164 # Make streaming request 

165 session: Final = requests.Session() 

166 try: 

167 response: Final = session.post( 

168 url, headers=self._get_headers(), json=data, stream=True, timeout=self._timeout 

169 ) 

170 response.raise_for_status() 

171 

172 # Parse SSE stream 

173 for line in response.iter_lines(): 

174 if line: 

175 line = line.decode("utf-8") 

176 if line.startswith("data: "): 

177 data_str = line[6:] # Remove 'data: ' prefix 

178 if data_str.strip() == "[DONE]": 

179 break 

180 try: 

181 chunk = json.loads(data_str) 

182 yield chunk 

183 except json.JSONDecodeError: 

184 continue 

185 

186 except requests.exceptions.HTTPError as e: 

187 if e.response.status_code == 401: 

188 raise UnauthorizedError(e) 

189 raise