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
« 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
5import requests
7from .exceptions import UnauthorizedError
10class ChatClient:
11 def __init__(self, base_url: str, api_key: str | None = None, timeout: int = 600):
12 """
13 Initialize the ChatClient.
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
25 def _get_headers(self) -> dict[str, str]:
26 """
27 Get the headers for API requests, including authorization if api_key is set.
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
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.
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
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
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"
75 # Build request data with required fields
76 data: Final[dict[str, object]] = {"model": model, "messages": messages}
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
94 request: Final = requests.Request("POST", url, headers=self._get_headers(), json=data)
96 if return_request:
97 return request
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
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.
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
136 Yields:
137 Dict[str, Any]: Streaming response chunks from the server
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"
145 # Build request data with required fields
146 data: Final[dict[str, object]] = {"model": model, "messages": messages, "stream": True}
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
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()
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
186 except requests.exceptions.HTTPError as e:
187 if e.response.status_code == 401:
188 raise UnauthorizedError(e)
189 raise