fix(providers): tighten local endpoint detection
Parse the endpoint host before disabling keepalive so public hostnames that merely contain private-network substrings keep the default connection pool behavior. Made-with: Cursor
This commit is contained in:
@@ -3,16 +3,18 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
|
||||||
import hashlib
|
import hashlib
|
||||||
import importlib.util
|
import importlib.util
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import secrets
|
import secrets
|
||||||
import string
|
import string
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
|
from ipaddress import ip_address
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import json_repair
|
import json_repair
|
||||||
@@ -174,23 +176,21 @@ def _is_local_endpoint(
|
|||||||
return True
|
return True
|
||||||
if not api_base:
|
if not api_base:
|
||||||
return False
|
return False
|
||||||
host = api_base.strip().lower().rstrip("/")
|
raw = api_base.strip().lower()
|
||||||
private_patterns = (
|
parsed = urlparse(raw if "://" in raw else f"//{raw}")
|
||||||
"localhost",
|
try:
|
||||||
"127.",
|
host = parsed.hostname
|
||||||
"192.168.",
|
except ValueError:
|
||||||
"10.",
|
return False
|
||||||
"host.docker.internal",
|
if host in {"localhost", "host.docker.internal"}:
|
||||||
"[::1]",
|
|
||||||
)
|
|
||||||
if any(p in host for p in private_patterns):
|
|
||||||
return True
|
return True
|
||||||
# 172.16.0.0 – 172.31.255.255
|
if not host:
|
||||||
import re
|
return False
|
||||||
m = re.search(r"172\.(\d+)\." , host)
|
try:
|
||||||
if m and 16 <= int(m.group(1)) <= 31:
|
addr = ip_address(host)
|
||||||
return True
|
except ValueError:
|
||||||
return False
|
return False
|
||||||
|
return addr.is_loopback or addr.is_private
|
||||||
|
|
||||||
|
|
||||||
def _is_direct_openai_base(api_base: str | None) -> bool:
|
def _is_direct_openai_base(api_base: str | None) -> bool:
|
||||||
|
|||||||
@@ -2,8 +2,6 @@
|
|||||||
|
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
from nanobot.providers.openai_compat_provider import (
|
from nanobot.providers.openai_compat_provider import (
|
||||||
OpenAICompatProvider,
|
OpenAICompatProvider,
|
||||||
_is_local_endpoint,
|
_is_local_endpoint,
|
||||||
@@ -74,6 +72,15 @@ class TestIsLocalEndpoint:
|
|||||||
def test_trailing_slash(self):
|
def test_trailing_slash(self):
|
||||||
assert _is_local_endpoint(None, "http://192.168.1.1:8080/v1/") is True
|
assert _is_local_endpoint(None, "http://192.168.1.1:8080/v1/") is True
|
||||||
|
|
||||||
|
def test_public_hostname_containing_localhost_is_not_local(self):
|
||||||
|
assert _is_local_endpoint(None, "https://notlocalhost.example/v1") is False
|
||||||
|
|
||||||
|
def test_public_hostname_containing_private_ip_prefix_is_not_local(self):
|
||||||
|
assert _is_local_endpoint(None, "https://api10.example.com/v1") is False
|
||||||
|
|
||||||
|
def test_url_without_scheme(self):
|
||||||
|
assert _is_local_endpoint(None, "192.168.1.1:8080/v1") is True
|
||||||
|
|
||||||
|
|
||||||
class TestLocalKeepaliveConfig:
|
class TestLocalKeepaliveConfig:
|
||||||
"""Verify that local endpoints get keepalive_expiry=0."""
|
"""Verify that local endpoints get keepalive_expiry=0."""
|
||||||
|
|||||||
Reference in New Issue
Block a user