mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Adding oci_key_file support for oci provider
This commit is contained in:
parent
e638c5490d
commit
87abfc3c11
3 changed files with 66 additions and 17 deletions
|
|
@ -44,7 +44,11 @@ response = completion(
|
|||
oci_user=<your_oci_user>,
|
||||
oci_fingerprint=<your_oci_fingerprint>,
|
||||
oci_tenancy=<your_oci_tenancy>,
|
||||
# Provide either the private key string OR the path to the key file:
|
||||
# Option 1: pass the private key as a string
|
||||
oci_key=<string_with_content_of_oci_key>,
|
||||
# Option 2: pass the private key file path
|
||||
# oci_key_file="<path/to/oci_key.pem>",
|
||||
oci_compartment_id=<oci_compartment_id>,
|
||||
)
|
||||
print(response)
|
||||
|
|
@ -67,7 +71,11 @@ response = completion(
|
|||
oci_user=<your_oci_user>,
|
||||
oci_fingerprint=<your_oci_fingerprint>,
|
||||
oci_tenancy=<your_oci_tenancy>,
|
||||
# Provide either the private key string OR the path to the key file:
|
||||
# Option 1: pass the private key as a string
|
||||
oci_key=<string_with_content_of_oci_key>,
|
||||
# Option 2: pass the private key file path
|
||||
# oci_key_file="<path/to/oci_key.pem>",
|
||||
oci_compartment_id=<oci_compartment_id>,
|
||||
)
|
||||
for chunk in response:
|
||||
|
|
|
|||
|
|
@ -92,6 +92,22 @@ def load_private_key_from_str(key_str: str):
|
|||
return key
|
||||
|
||||
|
||||
def load_private_key_from_file(file_path: str):
|
||||
"""Loads a private key from a file path"""
|
||||
try:
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
key_str = f.read().strip()
|
||||
except FileNotFoundError:
|
||||
raise FileNotFoundError(f"Private key file not found: {file_path}")
|
||||
except OSError as e:
|
||||
raise OSError(f"Failed to read private key file '{file_path}': {e}") from e
|
||||
|
||||
if not key_str:
|
||||
raise ValueError(f"Private key file is empty: {file_path}")
|
||||
|
||||
return load_private_key_from_str(key_str)
|
||||
|
||||
|
||||
def get_vendor_from_model(model: str) -> OCIVendors:
|
||||
"""
|
||||
Extracts the vendor from the model name.
|
||||
|
|
@ -237,10 +253,12 @@ class OCIChatConfig(BaseConfig):
|
|||
oci_fingerprint = optional_params.get("oci_fingerprint")
|
||||
oci_tenancy = optional_params.get("oci_tenancy")
|
||||
oci_key = optional_params.get("oci_key")
|
||||
oci_key_file = optional_params.get("oci_key_file")
|
||||
|
||||
if not oci_user or not oci_fingerprint or not oci_tenancy or not oci_key:
|
||||
if not oci_user or not oci_fingerprint or not oci_tenancy or not (oci_key or oci_key_file):
|
||||
raise Exception(
|
||||
"Missing one of the following parameters: oci_user, oci_fingerprint, oci_tenancy, oci_key"
|
||||
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, "
|
||||
"and at least one of oci_key or oci_key_file."
|
||||
)
|
||||
|
||||
method = str(optional_params.get("method", "POST")).upper()
|
||||
|
|
@ -283,7 +301,14 @@ class OCIChatConfig(BaseConfig):
|
|||
"Please install it with: pip install cryptography"
|
||||
) from e
|
||||
|
||||
private_key = load_private_key_from_str(oci_key)
|
||||
private_key = (
|
||||
load_private_key_from_str(oci_key)
|
||||
if oci_key
|
||||
else load_private_key_from_file(oci_key_file)
|
||||
if oci_key_file
|
||||
else None
|
||||
)
|
||||
|
||||
signature = private_key.sign(
|
||||
signing_string.encode("utf-8"),
|
||||
padding.PKCS1v15(),
|
||||
|
|
@ -334,17 +359,19 @@ class OCIChatConfig(BaseConfig):
|
|||
oci_fingerprint = optional_params.get("oci_fingerprint")
|
||||
oci_tenancy = optional_params.get("oci_tenancy")
|
||||
oci_key = optional_params.get("oci_key")
|
||||
oci_key_file = optional_params.get("oci_key_file")
|
||||
oci_compartment_id = optional_params.get("oci_compartment_id")
|
||||
|
||||
if (
|
||||
not oci_user
|
||||
or not oci_fingerprint
|
||||
or not oci_tenancy
|
||||
or not oci_key
|
||||
or not (oci_key or oci_key_file)
|
||||
or not oci_compartment_id
|
||||
):
|
||||
raise Exception(
|
||||
"Missing one of the following parameters: oci_user, oci_fingerprint, oci_tenancy, oci_key, oci_compartment_id"
|
||||
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, "
|
||||
"and at least one of oci_key or oci_key_file."
|
||||
)
|
||||
|
||||
if not api_base:
|
||||
|
|
|
|||
|
|
@ -17,18 +17,30 @@ TEST_MODEL_NAME = "xai.grok-4"
|
|||
TEST_MODEL = f"oci/{TEST_MODEL_NAME}"
|
||||
TEST_MESSAGES = [{"role": "user", "content": "Hello, how are you?"}]
|
||||
TEST_COMPARTMENT_ID = "ocid1.compartment.oc1..xxxxxx"
|
||||
TEST_OCI_PARAMS = {
|
||||
BASE_OCI_PARAMS = {
|
||||
"oci_region": "us-ashburn-1",
|
||||
"oci_user": "ocid1.user.oc1..xxxxxxEXAMPLExxxxxx",
|
||||
"oci_fingerprint": "4f:29:77:cc:b1:3e:55:ab:61:2a:de:47:f1:38:4c:90",
|
||||
"oci_tenancy": "ocid1.tenancy.oc1..xxxxxxEXAMPLExxxxxx",
|
||||
"oci_compartment_id": TEST_COMPARTMENT_ID,
|
||||
"oci_key": "<private_key.pem as string>"
|
||||
}
|
||||
|
||||
TEST_OCI_PARAMS_KEY = {
|
||||
**BASE_OCI_PARAMS,
|
||||
"oci_key": "<private_key.pem as string>",
|
||||
}
|
||||
|
||||
TEST_OCI_PARAMS_KEY_FILE = {
|
||||
**BASE_OCI_PARAMS,
|
||||
"oci_key_file": "<private_key.pem as a Path>",
|
||||
}
|
||||
|
||||
class TestOCIChatConfig:
|
||||
def test_validate_environment_with_oci_region(self):
|
||||
@pytest.mark.parametrize(
|
||||
"optional_params",
|
||||
[TEST_OCI_PARAMS_KEY, TEST_OCI_PARAMS_KEY_FILE],
|
||||
)
|
||||
def test_validate_environment_with_oci_region(self, optional_params):
|
||||
config = OCIChatConfig()
|
||||
|
||||
headers = {}
|
||||
|
|
@ -37,25 +49,25 @@ class TestOCIChatConfig:
|
|||
headers=headers,
|
||||
model=TEST_MODEL,
|
||||
messages=TEST_MESSAGES, # type: ignore
|
||||
optional_params=TEST_OCI_PARAMS,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert result["content-type"] == "application/json"
|
||||
assert result["user-agent"] == f"litellm/{version}"
|
||||
|
||||
def test_missing_oci_auth_parameters(self):
|
||||
optional_params = TEST_OCI_PARAMS.copy()
|
||||
optional_params.pop("oci_region")
|
||||
def test_missing_oci_auth_parameters(self, optional_params):
|
||||
# Start with a copy and remove `oci_region`
|
||||
params = optional_params.copy()
|
||||
params.pop("oci_region")
|
||||
|
||||
# Remove optional_params one by one and verify that an exception is raised
|
||||
for key in optional_params.keys():
|
||||
modified_params = optional_params.copy()
|
||||
# Remove each remaining param one by one
|
||||
for key in list(params.keys()):
|
||||
modified_params = params.copy()
|
||||
del modified_params[key]
|
||||
|
||||
with pytest.raises(Exception) as excinfo:
|
||||
config = OCIChatConfig()
|
||||
|
||||
headers = {}
|
||||
|
||||
config.validate_environment(
|
||||
|
|
@ -67,7 +79,9 @@ class TestOCIChatConfig:
|
|||
litellm_params={},
|
||||
)
|
||||
|
||||
assert f"Missing one of the following parameters: oci_user, oci_fingerprint, oci_tenancy, oci_key, oci_compartment_id" in str(excinfo.value)
|
||||
assert (
|
||||
"Missing required parameters"
|
||||
) in str(excinfo.value)
|
||||
|
||||
def test_transform_request_simple(self):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue