This commit is contained in:
IvanShang 2026-08-26 15:35:32 +08:00 committed by GitHub
commit 3c9d101532
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 695 additions and 9 deletions

View file

@ -152,6 +152,7 @@ LLM_CONFIG_NAMES: Final = (
"InfinityRerankConfig",
"JinaAIRerankConfig",
"DeepinfraRerankConfig",
"GPUStackRerankConfig",
"HostedVLLMRerankConfig",
"NvidiaNimRerankConfig",
"NvidiaNimRankingConfig",
@ -282,6 +283,7 @@ LLM_CONFIG_NAMES: Final = (
"AzureSpeechAudioTranscriptionConfig",
"HostedVLLMChatConfig",
"HostedVLLMEmbeddingConfig",
"GPUStackEmbeddingConfig",
# Alias for backwards compatibility
"VolcEngineConfig", # Alias for VolcEngineChatConfig
"LlamafileChatConfig",
@ -675,6 +677,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.deepinfra.rerank.transformation",
"DeepinfraRerankConfig",
),
"GPUStackRerankConfig": (
".llms.gpustack.rerank.transformation",
"GPUStackRerankConfig",
),
"HostedVLLMRerankConfig": (
".llms.hosted_vllm.rerank.transformation",
"HostedVLLMRerankConfig",
@ -1094,6 +1100,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.hosted_vllm.embedding.transformation",
"HostedVLLMEmbeddingConfig",
),
"GPUStackEmbeddingConfig": (
".llms.gpustack.embedding.transformation",
"GPUStackEmbeddingConfig",
),
# Alias for backwards compatibility
"VolcEngineConfig": (
".llms.volcengine.chat.transformation",

View file

@ -0,0 +1,76 @@
from typing import Final
from urllib.parse import urlsplit, urlunsplit
from litellm.secret_managers.main import get_secret_str
def get_gpustack_api_base(api_base: str | None) -> str:
resolved_api_base: Final = api_base or get_secret_str("GPUSTACK_API_BASE")
if resolved_api_base is None:
raise ValueError("api_base is required for GPUStack. Set it in the call or via GPUSTACK_API_BASE.")
return resolved_api_base
def get_gpustack_endpoint(api_base: str | None, endpoint: str) -> str:
parsed_api_base: Final = urlsplit(get_gpustack_api_base(api_base))
normalized_path: Final = parsed_api_base.path.rstrip("/")
normalized_endpoint: Final = endpoint.strip("/")
endpoint_path: Final = (
normalized_path
if normalized_path.endswith(f"/{normalized_endpoint}")
else (
f"{normalized_path}/{normalized_endpoint}"
if normalized_path.endswith("/v1")
else f"{normalized_path}/v1/{normalized_endpoint}"
)
)
return urlunsplit(parsed_api_base._replace(path=endpoint_path))
def get_gpustack_api_key(api_key: str | None) -> str | None:
return api_key or get_secret_str("GPUSTACK_API_KEY")
# fmt: off
def get_gpustack_headers(
headers: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
api_key: str | None,
*,
include_accept: bool = False,
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
resolved_api_key: Final = get_gpustack_api_key(api_key)
deduplicated_headers_by_name: Final = { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
header_name.lower(): (header_name, header_value) for header_name, header_value in headers.items()
}
deduplicated_headers: Final = { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
header_name: header_value for header_name, header_value in deduplicated_headers_by_name.values()
}
header_names: Final = set( # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
deduplicated_headers_by_name
) # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
default_headers: Final = { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
**(
{"Content-Type": "application/json"} if "content-type" not in header_names else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
), # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
**(
{"Accept": "application/json"} if include_accept and "accept" not in header_names else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
), # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
**(
{ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"Authorization": f"Bearer {resolved_api_key}"
} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
if resolved_api_key is not None and "authorization" not in header_names
else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
),
}
return { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
**default_headers,
**deduplicated_headers,
} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
# fmt: on
def strip_gpustack_model_prefix(model: str) -> str:
if model.startswith("gpustack/"):
return model.replace("gpustack/", "", 1)
return model

View file

@ -0,0 +1,105 @@
from typing import Final
import httpx
from pydantic import TypeAdapter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.gpustack.common_utils import (
get_gpustack_endpoint,
get_gpustack_headers,
strip_gpustack_model_prefix,
)
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse
class GPUStackEmbeddingError(BaseLLMException):
pass
# fmt: off
class GPUStackEmbeddingConfig(BaseEmbeddingConfig):
def validate_environment(
self,
headers: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
model: str,
messages: list[AllMessageValues], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
optional_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
litellm_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
api_key: str | None = None,
api_base: str | None = None,
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return get_gpustack_headers(headers=headers, api_key=api_key)
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
litellm_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
stream: bool | None = None,
) -> str:
return get_gpustack_endpoint(api_base=api_base, endpoint="/embeddings")
def transform_embedding_request(
self,
model: str,
input: AllEmbeddingInputValues,
optional_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
headers: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
encoding_format: Final[object | None] = optional_params.get("encoding_format")
encoding_format_body: Final = (
{"encoding_format": encoding_format} if encoding_format not in (None, "") else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"model": strip_gpustack_model_prefix(model),
"input": input,
**encoding_format_body,
}
def transform_embedding_response(
self,
model: str,
raw_response: httpx.Response,
model_response: EmbeddingResponse,
logging_obj: object,
api_key: str | None,
request_data: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
optional_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
litellm_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> EmbeddingResponse:
return TypeAdapter(EmbeddingResponse).validate_json(raw_response.content)
def get_supported_openai_params(
self, model: str
) -> list[str]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return ["encoding_format", "timeout"] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
def map_openai_params(
self,
non_default_params: dict[ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
str, object
], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
optional_params: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
model: str,
drop_params: bool,
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
**optional_params,
**{ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
param: value for param, value in non_default_params.items() if param == "encoding_format"
}, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
}
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| httpx.Headers, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> BaseLLMException:
return GPUStackEmbeddingError(message=error_message, status_code=status_code, headers=headers)
# fmt: on

View file

@ -0,0 +1,185 @@
from typing import Final
import httpx
from pydantic import BaseModel, Field, TypeAdapter
from litellm._uuid import uuid
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.gpustack.common_utils import (
get_gpustack_endpoint,
get_gpustack_headers,
strip_gpustack_model_prefix,
)
from litellm.types.rerank import (
RerankBilledUnits,
RerankResponse,
RerankResponseDocument,
RerankResponseMeta,
RerankResponseResult,
RerankTokens,
)
class GPUStackRerankError(BaseLLMException):
pass
class GPUStackRerankDocumentPayload(BaseModel):
text: str | None = None
class GPUStackRerankResultPayload(BaseModel):
index: int
relevance_score: float
document: GPUStackRerankDocumentPayload | None = None
class GPUStackRerankUsagePayload(BaseModel):
total_tokens: int | None = None
# fmt: off
class GPUStackRerankResponsePayload(BaseModel):
id: str | None = None
results: list[ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
GPUStackRerankResultPayload
] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
usage: GPUStackRerankUsagePayload = Field(default_factory=GPUStackRerankUsagePayload)
class GPUStackRerankConfig(BaseRerankConfig):
def validate_environment(
self,
headers: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
model: str,
api_key: str | None = None,
optional_params: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return get_gpustack_headers(headers=headers, api_key=api_key, include_accept=True)
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> str:
return get_gpustack_endpoint(api_base=api_base, endpoint="/rerank")
def get_supported_cohere_rerank_params(
self, model: str
) -> list[str]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return [ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"query",
"documents",
"top_n",
"return_documents",
] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
def map_cohere_rerank_params(
self,
non_default_params: dict[ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
str, object
], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
model: str,
drop_params: bool,
query: str,
documents: list[ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
str | dict[str, object]
], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: list[str] | None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
instruction: str | None = None,
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
top_n_body: Final = (
{"top_n": top_n} if top_n is not None else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return_documents_body: Final = (
{"return_documents": return_documents} if return_documents is not None else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"query": query,
"documents": documents,
**top_n_body,
**return_documents_body,
}
def transform_rerank_request(
self,
model: str,
optional_rerank_params: dict[ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
str, object
], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
headers: dict[str, object], # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
litellm_params: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> dict[str, object]: # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
return { # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"model": strip_gpustack_model_prefix(model),
"query": optional_rerank_params["query"],
"documents": optional_rerank_params["documents"],
**(
{"top_n": optional_rerank_params["top_n"]} if optional_rerank_params.get("top_n") is not None else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
), # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
**(
{ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"return_documents": optional_rerank_params["return_documents"]
} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
if optional_rerank_params.get("return_documents") is not None
else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
),
}
def transform_rerank_response(
self,
model: str,
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: object,
api_key: str | None = None,
request_data: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
optional_params: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
litellm_params: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| None = None, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> RerankResponse:
response_json: Final = TypeAdapter(GPUStackRerankResponsePayload).validate_json(raw_response.content)
total_tokens: Final = response_json.usage.total_tokens or 0
return RerankResponse(
id=response_json.id or str(uuid.uuid4()),
results=[ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
RerankResponseResult(
index=result.index,
relevance_score=result.relevance_score,
**(
{ # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
"document": RerankResponseDocument(text=result.document.text)
} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
if result.document is not None and result.document.text is not None
else {} # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
),
)
for result in response_json.results
],
meta=RerankResponseMeta(
billed_units=RerankBilledUnits(total_tokens=total_tokens),
tokens=RerankTokens(input_tokens=total_tokens),
),
)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict[str, object] # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
| httpx.Headers, # mutable-ok: LiteLLM provider interfaces require mutable JSON containers
) -> BaseLLMException:
return GPUStackRerankError(message=error_message, status_code=status_code, headers=headers)
# fmt: on

View file

@ -6335,14 +6335,17 @@ def embedding(
client=client,
aembedding=aembedding,
)
elif custom_llm_provider == "hosted_vllm":
api_base = api_base or litellm.api_base or get_secret_str("HOSTED_VLLM_API_BASE")
# set API KEY
elif custom_llm_provider in ("hosted_vllm", "gpustack"):
provider_env_prefix: Final = custom_llm_provider.upper()
api_base = ( # rebind-ok: provider dispatch resolves explicit and environment configuration
api_base or litellm.api_base or get_secret_str(f"{provider_env_prefix}_API_BASE")
)
if api_key is None:
api_key = litellm.api_key or get_secret_str("HOSTED_VLLM_API_KEY")
api_key = ( # rebind-ok: provider dispatch resolves explicit and environment configuration
litellm.api_key or get_secret_str(f"{provider_env_prefix}_API_KEY")
)
response = base_llm_http_handler.embedding(
response = base_llm_http_handler.embedding( # rebind-ok: provider dispatch resolves explicit and environment configuration
model=model,
input=input,
custom_llm_provider=custom_llm_provider,
@ -6355,7 +6358,7 @@ def embedding(
client=client,
aembedding=aembedding,
litellm_params=litellm_params_dict,
headers=headers or {},
headers=headers or {}, # mutable-ok: HTTP handler requires mutable request headers
)
elif (
custom_llm_provider == "openai_like"

View file

@ -85,6 +85,7 @@ def rerank(
"azure_ai",
"infinity",
"litellm_proxy",
"gpustack",
"hosted_vllm",
"deepinfra",
"fireworks_ai",
@ -395,6 +396,32 @@ def rerank(
litellm_params=rerank_litellm_params,
)
# fmt: off
elif _custom_llm_provider == litellm.LlmProviders.GPUSTACK:
api_key = ( # rebind-ok: provider dispatch resolves explicit and environment configuration
dynamic_api_key or optional_params.api_key or litellm.api_key or get_secret_str("GPUSTACK_API_KEY")
)
api_base = ( # rebind-ok: provider dispatch resolves explicit and environment configuration
dynamic_api_base or optional_params.api_base or litellm.api_base or get_secret_str("GPUSTACK_API_BASE")
)
response = base_llm_http_handler.rerank( # rebind-ok: provider dispatch resolves explicit and environment configuration
model=model,
custom_llm_provider=_custom_llm_provider,
provider_config=rerank_provider_config,
optional_rerank_params=optional_rerank_params,
logging_obj=litellm_logging_obj,
timeout=optional_params.timeout,
api_key=api_key,
api_base=api_base,
_is_async=_is_async,
headers=headers or litellm.headers or {}, # mutable-ok: HTTP handler requires mutable request headers
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
# fmt: on
elif _custom_llm_provider == litellm.LlmProviders.DEEPINFRA:
api_key = dynamic_api_key or optional_params.api_key or get_secret_str("DEEPINFRA_API_KEY")

View file

@ -3747,6 +3747,7 @@ class LlmProviders(str, Enum):
DOCKER_MODEL_RUNNER = "docker_model_runner"
CUSTOM = "custom"
LITELLM_PROXY = "litellm_proxy"
GPUSTACK = "gpustack"
HOSTED_VLLM = "hosted_vllm"
TENCENT = "tencent"
LLAMAFILE = "llamafile"

View file

@ -8274,6 +8274,8 @@ class ProviderConfigManager:
return litellm.CometAPIEmbeddingConfig()
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
return litellm.GithubCopilotEmbeddingConfig()
elif litellm.LlmProviders.GPUSTACK == provider:
return litellm.GPUStackEmbeddingConfig()
elif litellm.LlmProviders.OPENROUTER == provider:
from litellm.llms.openrouter.embedding.transformation import (
OpenrouterEmbeddingConfig,
@ -8318,8 +8320,17 @@ class ProviderConfigManager:
return litellm.InfinityRerankConfig()
elif litellm.LlmProviders.JINA_AI == provider:
return litellm.JinaAIRerankConfig()
elif litellm.LlmProviders.HOSTED_VLLM == provider:
return litellm.HostedVLLMRerankConfig()
elif provider in (
litellm.LlmProviders.GPUSTACK,
litellm.LlmProviders.HOSTED_VLLM,
):
rerank_configs: Final = MappingProxyType(
{
litellm.LlmProviders.GPUSTACK: litellm.GPUStackRerankConfig,
litellm.LlmProviders.HOSTED_VLLM: litellm.HostedVLLMRerankConfig,
}
)
return rerank_configs[provider]()
elif litellm.LlmProviders.HUGGINGFACE == provider:
return litellm.HuggingFaceRerankConfig()
elif litellm.LlmProviders.DEEPINFRA == provider:

View file

@ -1243,6 +1243,22 @@
"interactions": true
}
},
"gpustack": {
"display_name": "GPUStack (`gpustack`)",
"url": "https://docs.litellm.ai/docs/providers/gpustack",
"endpoints": {
"chat_completions": false,
"messages": false,
"responses": false,
"embeddings": true,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": true
}
},
"heroku": {
"display_name": "Heroku (`heroku`)",
"url": "https://docs.litellm.ai/docs/providers/heroku",

View file

@ -0,0 +1,252 @@
import json
from typing import Final
import httpx
import pytest
import litellm
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.gpustack.common_utils import get_gpustack_endpoint, get_gpustack_headers
def test_gpustack_provider_resolution_preserves_owner_route_id() -> None:
model, provider, _, _ = litellm.get_llm_provider(model="gpustack/owner/bge-m3")
assert model == "owner/bge-m3"
assert provider == litellm.LlmProviders.GPUSTACK.value
@pytest.mark.parametrize(
("api_base", "endpoint", "expected"),
[
(
"https://gpustack.test/v1/embeddings?tenant=a#fragment",
"/embeddings",
"https://gpustack.test/v1/embeddings?tenant=a#fragment",
),
(
"https://embeddings/v1?tenant=a",
"/embeddings",
"https://embeddings/v1/embeddings?tenant=a",
),
(
"https://rerank/prefix?tenant=a",
"/rerank",
"https://rerank/prefix/v1/rerank?tenant=a",
),
(
"https://gpustack.test/v1?tenant=a/",
"/embeddings",
"https://gpustack.test/v1/embeddings?tenant=a/",
),
(
"https://gpustack.test/v1#section/",
"/rerank",
"https://gpustack.test/v1/rerank#section/",
),
],
)
def test_gpustack_endpoint_normalization_uses_url_path(
api_base: str,
endpoint: str,
expected: str,
) -> None:
assert get_gpustack_endpoint(api_base, endpoint) == expected
def test_gpustack_headers_deduplicate_caller_header_casing() -> None:
headers = get_gpustack_headers(
{
"Authorization": "Bearer first",
"authorization": "Bearer second",
"CONTENT-TYPE": "application/custom",
},
"generated-key",
include_accept=True,
)
assert [key for key in headers if key.lower() == "authorization"] == ["authorization"]
assert headers["authorization"] == "Bearer second"
assert [key for key in headers if key.lower() == "content-type"] == ["CONTENT-TYPE"]
def test_gpustack_embedding_posts_to_v1_embeddings_with_caller_authorization() -> None:
captured_requests: list[httpx.Request] = []
def handler(request: httpx.Request) -> httpx.Response:
captured_requests.append(request)
return httpx.Response(
200,
json={
"object": "list",
"model": "owner/bge-m3",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"usage": {"prompt_tokens": 3, "total_tokens": 3},
},
)
client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))
response = litellm.embedding(
model="gpustack/owner/bge-m3",
input=["hello"],
api_base="https://gpustack.test",
api_key="generated-key",
encoding_format="float",
headers={"authorization": "Bearer caller-key"},
client=client,
)
assert len(captured_requests) == 1
request: Final = captured_requests[0]
assert str(request.url) == "https://gpustack.test/v1/embeddings"
assert request.headers["Authorization"] == "Bearer caller-key"
assert request.headers.get_list("authorization") == ["Bearer caller-key"]
assert json.loads(request.content) == {
"model": "owner/bge-m3",
"input": ["hello"],
"encoding_format": "float",
}
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
assert response.usage.total_tokens == 3
@pytest.mark.asyncio()
async def test_gpustack_aembedding_uses_env_base_and_key_without_double_v1(monkeypatch: pytest.MonkeyPatch) -> None:
captured_requests: list[httpx.Request] = []
async def handler(request: httpx.Request) -> httpx.Response:
captured_requests.append(request)
return httpx.Response(
200,
json={
"object": "list",
"model": "embeddings",
"data": [{"object": "embedding", "index": 0, "embedding": [0.4, 0.5]}],
"usage": {"prompt_tokens": 2, "total_tokens": 2},
},
)
monkeypatch.setenv("GPUSTACK_API_BASE", "https://gpustack-env.test/v1")
monkeypatch.setenv("GPUSTACK_API_KEY", "env-key")
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
try:
response = await litellm.aembedding(
model="gpustack/embeddings",
input=["hello"],
client=client,
)
await GLOBAL_LOGGING_WORKER.clear_queue()
await GLOBAL_LOGGING_WORKER.stop()
finally:
await client.close()
assert len(captured_requests) == 1
request: Final = captured_requests[0]
assert str(request.url) == "https://gpustack-env.test/v1/embeddings"
assert request.headers["Authorization"] == "Bearer env-key"
assert json.loads(request.content) == {"model": "embeddings", "input": ["hello"]}
assert response.data[0]["embedding"] == [0.4, 0.5]
def test_gpustack_rerank_posts_supported_fields_and_preserves_usage() -> None:
captured_requests: list[httpx.Request] = []
def handler(request: httpx.Request) -> httpx.Response:
captured_requests.append(request)
return httpx.Response(
200,
json={
"model": "owner/reranker",
"results": [
{
"index": 1,
"document": {"text": "second"},
"relevance_score": 0.98,
}
],
"usage": {"total_tokens": 7},
},
)
client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))
response = litellm.rerank(
model="gpustack/owner/reranker",
query="best doc",
documents=["first", "second"],
top_n=1,
return_documents=True,
rank_fields=["ignored"],
max_tokens_per_doc=128,
api_base="https://gpustack.test/v1/rerank",
api_key="generated-key",
headers={"authorization": "Bearer caller-key"},
client=client,
)
assert len(captured_requests) == 1
request: Final = captured_requests[0]
assert str(request.url) == "https://gpustack.test/v1/rerank"
assert request.headers["Authorization"] == "Bearer caller-key"
assert request.headers.get_list("authorization") == ["Bearer caller-key"]
assert json.loads(request.content) == {
"model": "owner/reranker",
"query": "best doc",
"documents": ["first", "second"],
"top_n": 1,
"return_documents": True,
}
assert response.results[0]["document"]["text"] == "second"
assert response.results[0]["relevance_score"] == 0.98
assert response.meta["billed_units"]["total_tokens"] == 7
assert response.meta["tokens"]["input_tokens"] == 7
@pytest.mark.asyncio()
async def test_gpustack_arerank_uses_env_base_and_key(monkeypatch: pytest.MonkeyPatch) -> None:
captured_requests: list[httpx.Request] = []
async def handler(request: httpx.Request) -> httpx.Response:
captured_requests.append(request)
return httpx.Response(
200,
json={
"results": [{"index": 0, "document": {"text": "first"}, "relevance_score": 0.8}],
"usage": {"total_tokens": None},
},
)
monkeypatch.setenv("GPUSTACK_API_BASE", "https://gpustack.test/")
monkeypatch.setenv("GPUSTACK_API_KEY", "generated-key")
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
try:
response = await litellm.arerank(
model="gpustack/reranker",
query="best doc",
documents=["first", "second"],
return_documents=True,
client=client,
)
await GLOBAL_LOGGING_WORKER.clear_queue()
await GLOBAL_LOGGING_WORKER.stop()
finally:
await client.close()
assert len(captured_requests) == 1
request: Final = captured_requests[0]
assert str(request.url) == "https://gpustack.test/v1/rerank"
assert request.headers["Authorization"] == "Bearer generated-key"
assert json.loads(request.content) == {
"model": "reranker",
"query": "best doc",
"documents": ["first", "second"],
"return_documents": True,
}
assert response.results[0]["index"] == 0
assert response.meta["billed_units"]["total_tokens"] == 0