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:
Xubin Ren
2026-04-26 16:14:24 +08:00
committed by Xubin Ren
parent 5943ab386d
commit 1e11b35b45
2 changed files with 26 additions and 19 deletions
+17 -17
View File
@@ -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.",
"host.docker.internal",
"[::1]",
)
if any(p in host for p in private_patterns):
return True
# 172.16.0.0 172.31.255.255
import re
m = re.search(r"172\.(\d+)\." , host)
if m and 16 <= int(m.group(1)) <= 31:
return True
return False return False
if host in {"localhost", "host.docker.internal"}:
return True
if not host:
return False
try:
addr = ip_address(host)
except ValueError:
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."""