diff --git a/litellm/llms/oci/common_utils.py b/litellm/llms/oci/common_utils.py index 50fcf6f4194..2d9c640fea4 100644 --- a/litellm/llms/oci/common_utils.py +++ b/litellm/llms/oci/common_utils.py @@ -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/``), 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): diff --git a/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py index c68af52e7c6..30f49bea344 100644 --- a/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py +++ b/tests/test_litellm/llms/oci/embed/test_oci_embed_transformation.py @@ -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 # ------------------------------------------------------------------ diff --git a/tests/test_litellm/llms/oci/test_oci_common_utils.py b/tests/test_litellm/llms/oci/test_oci_common_utils.py index 90fea6fe2aa..d306d7351dd 100644 --- a/tests/test_litellm/llms/oci/test_oci_common_utils.py +++ b/tests/test_litellm/llms/oci/test_oci_common_utils.py @@ -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" diff --git a/tests/test_litellm/llms/oci/test_oci_coverage_boost.py b/tests/test_litellm/llms/oci/test_oci_coverage_boost.py index 015cd423e4f..575122bf70e 100644 --- a/tests/test_litellm/llms/oci/test_oci_coverage_boost.py +++ b/tests/test_litellm/llms/oci/test_oci_coverage_boost.py @@ -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):