mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
(sap) clean up code
This commit is contained in:
parent
70e64c572d
commit
3736812c64
1 changed files with 63 additions and 58 deletions
|
|
@ -7,9 +7,10 @@ from dataclasses import dataclass
|
|||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import httpx
|
||||
|
||||
from litellm import sap_service_key
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client, HTTPHandler
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
AUTH_ENDPOINT_SUFFIX = "/oauth/token"
|
||||
|
|
@ -64,6 +65,10 @@ def _get_vcap_service(label: str) -> Optional[Dict[str, Any]]:
|
|||
return svc
|
||||
return None
|
||||
|
||||
@dataclass
|
||||
class Source:
|
||||
name: str
|
||||
get: Callable[[CredentialsValue], Optional[str]]
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CredentialsValue:
|
||||
|
|
@ -106,12 +111,6 @@ CREDENTIAL_VALUES: Final[List[CredentialsValue]] = [
|
|||
),
|
||||
]
|
||||
|
||||
@dataclass
|
||||
class Source:
|
||||
name: str
|
||||
get: Callable[[CredentialsValue], Optional[str]]
|
||||
|
||||
|
||||
def init_conf(profile: Optional[str] = None) -> Dict[str, Any]:
|
||||
"""
|
||||
Loads config JSON from:
|
||||
|
|
@ -154,7 +153,7 @@ def init_conf(profile: Optional[str] = None) -> Dict[str, Any]:
|
|||
def _env_name(name: str) -> str:
|
||||
return f"AICORE_{name.upper()}"
|
||||
|
||||
def extract_credentials(source: Source, exclude: List[str] = None) -> Dict[str, str]:
|
||||
def extract_credentials(source: Source, exclude: Optional[List[str]] = None) -> Dict[str, str]:
|
||||
"""Extract all credentials from a source."""
|
||||
exclude = exclude or []
|
||||
credentials = {}
|
||||
|
|
@ -171,6 +170,7 @@ def resolve_credentials(sources: List[Source]) -> Dict[str, str]:
|
|||
for source in sources:
|
||||
credentials = extract_credentials(source, exclude=['resource_group'])
|
||||
if credentials:
|
||||
verbose_logger.debug(f"Resolved SAP credentials from source {source.name}")
|
||||
return credentials
|
||||
raise ValueError("No credentials found in any source")
|
||||
|
||||
|
|
@ -180,12 +180,11 @@ def resolve_resource_group(sources: List[Source]) -> Optional[str]:
|
|||
for source in sources:
|
||||
value = source.get(rg_cred)
|
||||
if value:
|
||||
verbose_logger.debug(f"Resolved GEN AI Hub resource_group from source {source.name}")
|
||||
return value
|
||||
return rg_cred.default
|
||||
|
||||
def fetch_credentials(
|
||||
service_key: Optional[str] = None, profile: Optional[str] = None, **kwargs
|
||||
) -> Dict[str, str]:
|
||||
def fetch_credentials(service_key: Optional[str, dict] = None, profile: Optional[str] = None, **kwargs) -> Dict[str, str]:
|
||||
"""
|
||||
Resolution order per key:
|
||||
kwargs
|
||||
|
|
@ -201,7 +200,7 @@ def fetch_credentials(
|
|||
vcap_service = _get_vcap_service(VCAP_AICORE_SERVICE_NAME)
|
||||
|
||||
sources = [
|
||||
Source("service key", lambda cv: _get_nested(service_key, cv.vcap_key if cv.vcap_key else (cv.name,))),
|
||||
Source("service key", lambda cv: _get_nested(service_key, cv.vcap_key if cv.vcap_key else (cv.name,))), # type: ignore[arg-type]
|
||||
Source("kwargs",
|
||||
lambda cv: _str_or_none(kwargs.get(cv.name))),
|
||||
Source("environment variables",
|
||||
|
|
@ -209,7 +208,7 @@ def fetch_credentials(
|
|||
Source("config file",
|
||||
lambda cv: _str_or_none(config.get(f'AICORE_{cv.name.upper()}'))),
|
||||
Source("VCAP service",
|
||||
lambda cv: _get_nested(vcap_service, ("credentials",) + cv.vcap_key if cv.vcap_key else (cv.name,))),
|
||||
lambda cv: _get_nested(vcap_service, ("credentials",) + cv.vcap_key if cv.vcap_key else (cv.name,))), # type: ignore[arg-type]
|
||||
]
|
||||
|
||||
credentials = resolve_credentials(sources)
|
||||
|
|
@ -223,48 +222,15 @@ def fetch_credentials(
|
|||
credentials['auth_url'] = credentials.pop('cert_url')
|
||||
return credentials
|
||||
|
||||
|
||||
def get_token_creator(
|
||||
service_key: Optional[str] = None,
|
||||
profile: Optional[str] = None,
|
||||
*,
|
||||
timeout: float = 30.0,
|
||||
expiry_buffer_minutes: int = 60,
|
||||
**overrides,
|
||||
) -> Tuple[Callable[[], str], str, str]:
|
||||
"""
|
||||
Creates a callable that fetches and caches an OAuth2 bearer token
|
||||
using credentials from `fetch_credentials()`.
|
||||
|
||||
The callable:
|
||||
- Automatically loads credentials via fetch_credentials(profile, **overrides)
|
||||
- Fetches a new token only if expired or near expiry
|
||||
- Caches token thread-safely with a configurable refresh buffer
|
||||
|
||||
Args:
|
||||
profile: Optional AICore profile name
|
||||
timeout: HTTP request timeout in seconds (default 30s)
|
||||
expiry_buffer_minutes: Refresh the token this many minutes before expiry
|
||||
overrides: Any explicit credential overrides (client_id, client_secret, etc.)
|
||||
|
||||
Returns:
|
||||
Callable[[], str]: function returning a valid "Bearer <token>" string.
|
||||
"""
|
||||
|
||||
# Resolve credentials using your helper
|
||||
credentials: Dict[str, str] = fetch_credentials(
|
||||
service_key=service_key, profile=profile, **overrides
|
||||
)
|
||||
|
||||
auth_url = credentials.get("auth_url")
|
||||
client_id = credentials.get("client_id")
|
||||
client_secret = credentials.get("client_secret")
|
||||
cert_str = credentials.get("cert_str")
|
||||
key_str = credentials.get("key_str")
|
||||
cert_file_path = credentials.get("cert_file_path")
|
||||
key_file_path = credentials.get("key_file_path")
|
||||
|
||||
# Sanity check
|
||||
def validate_credentials(
|
||||
auth_url: Optional[str] = None,
|
||||
client_id: Optional[str] = None,
|
||||
client_secret: Optional[str] = None,
|
||||
cert_str: Optional[str] = None,
|
||||
key_str: Optional[str] = None,
|
||||
cert_file_path: Optional[str] = None,
|
||||
key_file_path: Optional[str] = None
|
||||
):
|
||||
if not auth_url or not client_id:
|
||||
raise ValueError(
|
||||
"SAP AI Core credentials not found."
|
||||
|
|
@ -284,6 +250,45 @@ def get_token_creator(
|
|||
"(cert_str & key_str), or (cert_file_path & key_file_path)."
|
||||
)
|
||||
|
||||
def get_token_creator(
|
||||
service_key: Optional[str] = None,
|
||||
profile: Optional[str] = None,
|
||||
*,
|
||||
expiry_buffer_minutes: int = 60,
|
||||
**overrides,
|
||||
) -> Tuple[Callable[[], str], str, str]:
|
||||
"""
|
||||
Creates a callable that fetches and caches an OAuth2 bearer token
|
||||
using credentials from `fetch_credentials()`.
|
||||
|
||||
The callable:
|
||||
- Automatically loads credentials via fetch_credentials(profile, **overrides)
|
||||
- Fetches a new token only if expired or near expiry
|
||||
- Caches token thread-safely with a configurable refresh buffer
|
||||
|
||||
Args:
|
||||
profile: Optional AICore profile name
|
||||
expiry_buffer_minutes: Refresh the token this many minutes before expiry
|
||||
overrides: Any explicit credential overrides (client_id, client_secret, etc.)
|
||||
|
||||
Returns:
|
||||
Callable[[], str]: function returning a valid "Bearer <token>" string.
|
||||
"""
|
||||
|
||||
# Resolve credentials using your helper
|
||||
credentials: Dict[str, str] = fetch_credentials(service_key=service_key, profile=profile, **overrides)
|
||||
|
||||
auth_url = credentials.get("auth_url")
|
||||
client_id = credentials.get("client_id")
|
||||
client_secret = credentials.get("client_secret")
|
||||
cert_str = credentials.get("cert_str")
|
||||
key_str = credentials.get("key_str")
|
||||
cert_file_path = credentials.get("cert_file_path")
|
||||
key_file_path = credentials.get("key_file_path")
|
||||
|
||||
# Sanity check
|
||||
validate_credentials(auth_url, client_id, client_secret, cert_str, key_str, cert_file_path, key_file_path)
|
||||
|
||||
lock = Lock()
|
||||
token: Optional[str] = None
|
||||
token_expiry: Optional[datetime] = None
|
||||
|
|
@ -293,10 +298,10 @@ def get_token_creator(
|
|||
if client_secret:
|
||||
data["client_secret"] = client_secret
|
||||
|
||||
client = _get_httpx_client()
|
||||
# with httpx.Client(cert=cert_pair, timeout=timeout) as client:
|
||||
resp = client.post(auth_url, data=data)
|
||||
client = HTTPHandler(client=httpx.Client(cert=cert_pair)) if cert_pair else _get_httpx_client()
|
||||
|
||||
try:
|
||||
resp = client.post(auth_url, data=data) # type: ignore[arg-type]
|
||||
resp.raise_for_status()
|
||||
payload = resp.json()
|
||||
access_token = payload["access_token"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue