fix(oci): strip trailing action path in get_oci_base_url to avoid URL doubling

A fully-formed OCI endpoint URL (e.g. https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/chat) passed via api_base previously had the action path appended a second time by get_complete_url in both chat and embed configs, yielding a 404. get_oci_base_url now strips a trailing /20231130/actions/<name> so callers can always append the action path safely.
This commit is contained in:
mateo-berri 2026-05-19 09:11:18 +00:00
parent 878df5b9be
commit 212b30ee66
No known key found for this signature in database
4 changed files with 54 additions and 2 deletions

View file

@ -187,12 +187,18 @@ def resolve_oci_credentials(optional_params: dict) -> dict:
_OCI_REGION_RE = re.compile(r"^[a-z][a-z0-9-]{0,30}[a-z0-9]$")
_OCI_ACTION_PATH_RE = re.compile(rf"/{OCI_API_VERSION}/actions/[^/?#]+/?$")
def get_oci_base_url(optional_params: dict, api_base: Optional[str] = None) -> str:
"""Return the OCI inference base URL, respecting any explicit api_base override."""
"""Return the OCI inference base URL, respecting any explicit api_base override.
If ``api_base`` already ends with a fully-formed OCI action path
(``/{OCI_API_VERSION}/actions/<name>``), that suffix is stripped so callers
can append their own action path without producing a doubled URL.
"""
if api_base:
return api_base.rstrip("/")
return _OCI_ACTION_PATH_RE.sub("", api_base).rstrip("/")
creds = resolve_oci_credentials(optional_params)
region = creds["oci_region"]
if not isinstance(region, str) or not _OCI_REGION_RE.match(region):

View file

@ -96,6 +96,22 @@ class TestOCIEmbedConfig:
)
assert url == "https://custom.endpoint.example.com/20231130/actions/embedText"
def test_get_complete_url_full_url_is_not_doubled(self):
"""A fully-formed embedText URL must not have the action path appended twice."""
cfg = self._config()
full_url = (
"https://inference.generativeai.us-chicago-1.oci.oraclecloud.com"
"/20231130/actions/embedText"
)
url = cfg.get_complete_url(
api_base=full_url,
api_key=None,
model="cohere.embed-v3.0",
optional_params={},
litellm_params={},
)
assert url == full_url
# ------------------------------------------------------------------
# transform_embedding_request
# ------------------------------------------------------------------

View file

@ -152,6 +152,21 @@ def test_get_oci_base_url_explicit_api_base():
assert url == "https://custom.endpoint.com"
@pytest.mark.parametrize(
"api_base",
[
"https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/chat",
"https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/chat/",
"https://inference.generativeai.us-chicago-1.oci.oraclecloud.com/20231130/actions/embedText",
],
)
def test_get_oci_base_url_strips_trailing_action_path(api_base):
assert (
get_oci_base_url({}, api_base=api_base)
== "https://inference.generativeai.us-chicago-1.oci.oraclecloud.com"
)
def test_get_oci_base_url_from_region():
url = get_oci_base_url({"oci_region": "eu-frankfurt-1"})
assert url == "https://inference.generativeai.eu-frankfurt-1.oci.oraclecloud.com"

View file

@ -738,6 +738,21 @@ class TestOCIChatConfigGetCompleteUrl:
)
assert url == "https://custom.endpoint.com/20231130/actions/chat"
def test_full_chat_url_is_not_doubled(self):
config = OCIChatConfig()
full_url = (
"https://inference.generativeai.us-chicago-1.oci.oraclecloud.com"
"/20231130/actions/chat"
)
url = config.get_complete_url(
api_base=full_url,
api_key=None,
model=_GENERIC_MODEL,
optional_params={},
litellm_params={},
)
assert url == full_url
class TestOCIChatConfigGetErrorClass:
def test_returns_oci_error(self):