Skip to content

Commit 7115d1f

Browse files
committed
fix: address review feedback
1 parent fcf1fa8 commit 7115d1f

2 files changed

Lines changed: 18 additions & 16 deletions

File tree

api/config.py

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -127,23 +127,23 @@ def load_generator_config():
127127
generator_config = load_json_config("generator.json")
128128

129129
# Add client classes to each provider
130+
default_map = {
131+
"google": GoogleGenAIClient,
132+
"openai": OpenAIClient,
133+
"openrouter": OpenRouterClient,
134+
"ollama": OllamaClient,
135+
"bedrock": BedrockClient,
136+
"azure": AzureAIClient,
137+
"dashscope": DashscopeClient,
138+
"litellm": LiteLLMClient,
139+
}
130140
if "providers" in generator_config:
131141
for provider_id, provider_config in generator_config["providers"].items():
132142
# Try to set client class from client_class
133143
if provider_config.get("client_class") in CLIENT_CLASSES:
134144
provider_config["model_client"] = CLIENT_CLASSES[provider_config["client_class"]]
135145
# Fall back to default mapping based on provider_id
136-
elif provider_id in ["google", "openai", "openrouter", "ollama", "bedrock", "azure", "dashscope", "litellm"]:
137-
default_map = {
138-
"google": GoogleGenAIClient,
139-
"openai": OpenAIClient,
140-
"openrouter": OpenRouterClient,
141-
"ollama": OllamaClient,
142-
"bedrock": BedrockClient,
143-
"azure": AzureAIClient,
144-
"dashscope": DashscopeClient,
145-
"litellm": LiteLLMClient,
146-
}
146+
elif provider_id in default_map:
147147
provider_config["model_client"] = default_map[provider_id]
148148
else:
149149
logger.warning(f"Unknown provider or client class: {provider_id}")

api/litellm_client.py

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -167,14 +167,18 @@ def convert_inputs_to_api_kwargs(
167167
return final_model_kwargs
168168

169169
def parse_chat_completion(self, completion) -> GeneratorOutput:
170+
import types
171+
170172
try:
171-
data = self.chat_completion_parser(completion)
173+
is_stream = isinstance(completion, types.GeneratorType) or type(completion).__name__ == "CustomStreamWrapper"
174+
parser = handle_streaming_response if is_stream else self.chat_completion_parser
175+
data = parser(completion)
172176
except Exception as e:
173177
log.error(f"Error parsing the completion: {e}")
174178
return GeneratorOutput(data=None, error=str(e), raw_response=completion)
175179

176180
try:
177-
usage = self.track_completion_usage(completion)
181+
usage = None if is_stream else self.track_completion_usage(completion)
178182
return GeneratorOutput(
179183
data=None, error=None, raw_response=data, usage=usage
180184
)
@@ -207,7 +211,7 @@ def call(self, api_kwargs: Optional[Dict] = None, model_type: ModelType = ModelT
207211
import litellm
208212

209213
api_kwargs = api_kwargs or {}
210-
log.info(f"api_kwargs: {api_kwargs}")
214+
log.debug(f"api_kwargs: {api_kwargs}")
211215

212216
extra: Dict[str, Any] = {}
213217
if self._api_key:
@@ -218,8 +222,6 @@ def call(self, api_kwargs: Optional[Dict] = None, model_type: ModelType = ModelT
218222
if model_type == ModelType.EMBEDDER:
219223
return litellm.embedding(drop_params=True, **api_kwargs, **extra)
220224
elif model_type == ModelType.LLM:
221-
if api_kwargs.get("stream", False):
222-
self.chat_completion_parser = handle_streaming_response
223225
return litellm.completion(drop_params=True, **api_kwargs, **extra)
224226
else:
225227
raise ValueError(f"model_type {model_type} is not supported")

0 commit comments

Comments
 (0)