fix(llms,embeddings): repair HTTP proxy support (httpx>=0.28) and preserve proxies in LlmFactory (#5447)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -2,9 +2,8 @@ import os
|
||||
from abc import ABC
|
||||
from typing import Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from mem0.configs.base import AzureConfig
|
||||
from mem0.utils.http import build_http_client
|
||||
|
||||
|
||||
class BaseEmbedderConfig(ABC):
|
||||
@@ -81,7 +80,8 @@ class BaseEmbedderConfig(ABC):
|
||||
self.embedding_dims = embedding_dims
|
||||
|
||||
# AzureOpenAI specific
|
||||
self.http_client = httpx.Client(proxies=http_client_proxies) if http_client_proxies else None
|
||||
self.http_client_proxies = http_client_proxies
|
||||
self.http_client = build_http_client(http_client_proxies)
|
||||
|
||||
# Ollama specific
|
||||
self.ollama_base_url = ollama_base_url
|
||||
@@ -109,4 +109,3 @@ class BaseEmbedderConfig(ABC):
|
||||
self.aws_secret_access_key = aws_secret_access_key
|
||||
self.aws_session_token = aws_session_token
|
||||
self.aws_region = aws_region or os.environ.get("AWS_REGION") or "us-west-2"
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from abc import ABC
|
||||
from typing import Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
from mem0.utils.http import build_http_client
|
||||
|
||||
|
||||
class BaseLlmConfig(ABC):
|
||||
@@ -74,4 +74,5 @@ class BaseLlmConfig(ABC):
|
||||
self.vision_details = vision_details
|
||||
self.reasoning_effort = reasoning_effort
|
||||
self.is_reasoning_model = is_reasoning_model
|
||||
self.http_client = httpx.Client(proxies=http_client_proxies) if http_client_proxies else None
|
||||
self.http_client_proxies = http_client_proxies
|
||||
self.http_client = build_http_client(http_client_proxies)
|
||||
|
||||
@@ -29,7 +29,7 @@ class AnthropicLLM(LLMBase):
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
@@ -32,7 +32,7 @@ class AzureOpenAILLM(LLMBase):
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
reasoning_effort=getattr(config, 'reasoning_effort', None),
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
is_reasoning_model=getattr(config, 'is_reasoning_model', None),
|
||||
)
|
||||
|
||||
|
||||
@@ -28,7 +28,7 @@ class DeepSeekLLM(LLMBase):
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
@@ -27,7 +27,7 @@ class LMStudioLLM(LLMBase):
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
@@ -28,7 +28,7 @@ class MiniMaxLLM(LLMBase):
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
+1
-1
@@ -30,7 +30,7 @@ class OllamaLLM(LLMBase):
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
+1
-1
@@ -30,7 +30,7 @@ class OpenAILLM(LLMBase):
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
reasoning_effort=getattr(config, 'reasoning_effort', None),
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
is_reasoning_model=getattr(config, 'is_reasoning_model', None),
|
||||
)
|
||||
|
||||
|
||||
+1
-1
@@ -28,7 +28,7 @@ class VllmLLM(LLMBase):
|
||||
top_k=config.top_k,
|
||||
enable_vision=config.enable_vision,
|
||||
vision_details=config.vision_details,
|
||||
http_client_proxies=config.http_client,
|
||||
http_client_proxies=config.http_client_proxies,
|
||||
)
|
||||
|
||||
super().__init__(config)
|
||||
|
||||
@@ -100,7 +100,7 @@ class LlmFactory:
|
||||
"top_k": config.top_k,
|
||||
"enable_vision": config.enable_vision,
|
||||
"vision_details": config.vision_details,
|
||||
"http_client_proxies": config.http_client,
|
||||
"http_client_proxies": config.http_client_proxies,
|
||||
}
|
||||
config_dict.update(kwargs)
|
||||
config = config_class(**config_dict)
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
from typing import Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
def build_http_client(http_client_proxies: Optional[Union[Dict, str]]) -> Optional[httpx.Client]:
|
||||
if not http_client_proxies:
|
||||
return None
|
||||
if isinstance(http_client_proxies, dict):
|
||||
return httpx.Client(
|
||||
mounts={scheme: httpx.HTTPTransport(proxy=url) for scheme, url in http_client_proxies.items()}
|
||||
)
|
||||
return httpx.Client(proxy=http_client_proxies)
|
||||
@@ -17,6 +17,7 @@ dependencies = [
|
||||
"qdrant-client>=1.12.0",
|
||||
"pydantic>=2.7.3",
|
||||
"openai>=1.90.0",
|
||||
"httpx>=0.28.0",
|
||||
"posthog>=7.14.0",
|
||||
"pytz>=2024.1",
|
||||
"sqlalchemy>=2.0.31",
|
||||
|
||||
@@ -289,7 +289,7 @@ def test_generate_with_http_proxies(default_headers):
|
||||
api_version=None,
|
||||
default_headers=default_headers,
|
||||
)
|
||||
mock_http_client.assert_called_once_with(proxies="http://testproxy.mem0.net:8000")
|
||||
mock_http_client.assert_called_once_with(proxy="http://testproxy.mem0.net:8000")
|
||||
|
||||
|
||||
def test_init_with_api_key(monkeypatch):
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
import os
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.configs.llms.openai import OpenAIConfig
|
||||
from mem0.llms.openai import OpenAILLM
|
||||
|
||||
@@ -451,3 +453,14 @@ def test_callback_with_tools(mock_openai_client):
|
||||
mock_callback.assert_called_once()
|
||||
# Check that tool_calls exists in the message
|
||||
assert hasattr(mock_callback.call_args[0][1].choices[0].message, 'tool_calls')
|
||||
|
||||
|
||||
def test_openai_llm_preserves_proxies_from_base_config(mock_openai_client):
|
||||
config = BaseLlmConfig(
|
||||
model="gpt-4.1-nano-2025-04-14",
|
||||
api_key="api_key",
|
||||
http_client_proxies="http://proxy.local:8080",
|
||||
)
|
||||
llm = OpenAILLM(config)
|
||||
assert llm.config.http_client_proxies == "http://proxy.local:8080"
|
||||
assert isinstance(llm.config.http_client, httpx.Client)
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from mem0.configs.embeddings.base import BaseEmbedderConfig
|
||||
from mem0.configs.llms.base import BaseLlmConfig
|
||||
from mem0.utils.factory import LlmFactory
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_cls", [BaseLlmConfig, BaseEmbedderConfig])
|
||||
def test_config_with_string_proxy_builds_client(config_cls):
|
||||
config = config_cls(http_client_proxies="http://proxy.local:8080")
|
||||
assert isinstance(config.http_client, httpx.Client)
|
||||
assert config.http_client_proxies == "http://proxy.local:8080"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_cls", [BaseLlmConfig, BaseEmbedderConfig])
|
||||
def test_config_with_dict_proxy_builds_client(config_cls):
|
||||
proxies = {"http://": "http://p:8080", "https://": "http://p:8080"}
|
||||
config = config_cls(http_client_proxies=proxies)
|
||||
assert isinstance(config.http_client, httpx.Client)
|
||||
assert config.http_client_proxies == proxies
|
||||
|
||||
|
||||
@pytest.mark.parametrize("config_cls", [BaseLlmConfig, BaseEmbedderConfig])
|
||||
def test_config_without_proxy_has_no_client(config_cls):
|
||||
config = config_cls()
|
||||
assert config.http_client is None
|
||||
assert config.http_client_proxies is None
|
||||
|
||||
|
||||
def test_llm_factory_preserves_http_client_proxies():
|
||||
base = BaseLlmConfig(
|
||||
model="gpt-4o-mini",
|
||||
api_key="sk-test",
|
||||
http_client_proxies="http://proxy.local:8080",
|
||||
)
|
||||
llm = LlmFactory.create("openai", base)
|
||||
assert llm.config.http_client_proxies == "http://proxy.local:8080"
|
||||
assert isinstance(llm.config.http_client, httpx.Client)
|
||||
Reference in New Issue
Block a user