mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_bedrock_guardrail_stream_audit
This commit is contained in:
commit
e486e43d0f
20 changed files with 1671 additions and 25 deletions
|
|
@ -114,4 +114,4 @@ jobs:
|
|||
|
||||
- name: Audit provider endpoints against the schema
|
||||
working-directory: terraform/provider
|
||||
run: go run ./tools/endpointaudit -provider-dir ./litellm -spec "${RUNNER_TEMP}/openapi.json"
|
||||
run: go run ./tools/endpointaudit -provider-dir ./litellm -spec "${RUNNER_TEMP}/openapi.json" -coverage-allowlist ./tools/endpointaudit/coverage_allowlist.txt
|
||||
|
|
|
|||
|
|
@ -342,7 +342,7 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
self.start_time = datetime.now()
|
||||
self.completion_start_time: datetime | None = None
|
||||
|
||||
async def _handle_streaming_logging(self, collected_chunks: list[bytes]):
|
||||
async def _handle_streaming_logging(self, collected_chunks: list[bytes], *, stream_teardown: bool = False):
|
||||
"""Handle the logging after all chunks have been collected."""
|
||||
from litellm.proxy.pass_through_endpoints.streaming_handler import (
|
||||
PassThroughStreamingHandler,
|
||||
|
|
@ -354,21 +354,26 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
if self.completion_start_time is not None:
|
||||
self.litellm_logging_obj.completion_start_time = self.completion_start_time
|
||||
self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time
|
||||
logging_coroutine: Final = PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
litellm_logging_obj=self.litellm_logging_obj,
|
||||
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
|
||||
url_route="/v1/messages",
|
||||
request_body=self.request_body or {},
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=self.start_time,
|
||||
raw_bytes=collected_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
deferred_dispatch_armed: Final = (
|
||||
getattr(self.litellm_logging_obj, "_on_deferred_stream_complete", None) is not None
|
||||
)
|
||||
if deferred_dispatch_armed and not stream_teardown:
|
||||
self.litellm_logging_obj._deferred_stream_complete_args = (logging_coroutine,)
|
||||
return
|
||||
# Enqueue on the rooted logging worker rather than asyncio.create_task:
|
||||
# this also runs during generator teardown after a client disconnect,
|
||||
# where an unrooted task could be garbage-collected before it bills.
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
litellm_logging_obj=self.litellm_logging_obj,
|
||||
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
|
||||
url_route="/v1/messages",
|
||||
request_body=self.request_body or {},
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=self.start_time,
|
||||
raw_bytes=collected_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine)
|
||||
|
||||
def get_async_streaming_response_iterator(
|
||||
self,
|
||||
|
|
@ -433,7 +438,7 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
# post-loop logging below never runs and the tokens already streamed
|
||||
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
|
||||
if collected_chunks:
|
||||
await self._handle_streaming_logging(collected_chunks)
|
||||
await self._handle_streaming_logging(collected_chunks, stream_teardown=True)
|
||||
raise
|
||||
|
||||
if not saw_terminal_event:
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import json
|
|||
import logging
|
||||
import math
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
|
|
@ -2372,6 +2372,21 @@ class ProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
|
||||
elif (
|
||||
_post_call_guardrails_active
|
||||
and route_type == "anthropic_messages"
|
||||
and self._is_streaming_response(response)
|
||||
):
|
||||
from litellm.litellm_core_utils.logging_worker import (
|
||||
GLOBAL_LOGGING_WORKER,
|
||||
)
|
||||
|
||||
async def _on_deferred_native_stream_complete(
|
||||
logging_coroutine: Coroutine[object, object, object],
|
||||
) -> None:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _on_deferred_native_stream_complete
|
||||
|
||||
if route_type == "allm_passthrough_route":
|
||||
# Check if response is an async generator
|
||||
|
|
|
|||
|
|
@ -3167,8 +3167,14 @@ class ProxyLogging:
|
|||
# through each of them adds N pass-through trampolines per chunk for
|
||||
# zero behavior change. Skip the chain entirely and stream through.
|
||||
if not caps.iterator_overrides:
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
try:
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
except (GeneratorExit, asyncio.CancelledError):
|
||||
raise
|
||||
except Exception:
|
||||
ProxyLogging._fire_deferred_stream_logging(request_data)
|
||||
raise
|
||||
ProxyLogging._fire_deferred_stream_logging(request_data)
|
||||
return
|
||||
|
||||
|
|
@ -3221,9 +3227,14 @@ class ProxyLogging:
|
|||
),
|
||||
)
|
||||
|
||||
# Actually iterate through the chained async generator and yield chunks
|
||||
async for chunk in current_response:
|
||||
yield chunk
|
||||
try:
|
||||
async for chunk in current_response:
|
||||
yield chunk
|
||||
except (GeneratorExit, asyncio.CancelledError):
|
||||
raise
|
||||
except Exception:
|
||||
ProxyLogging._fire_deferred_stream_logging(request_data)
|
||||
raise
|
||||
|
||||
# Fire deferred logging AFTER all guardrail end-of-stream blocks
|
||||
# completed. unified_guardrail writes guardrail_information during
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ longer signal it.
|
|||
|
||||
### Added
|
||||
|
||||
- **jwt_key_mapping**: New `litellm_jwt_key_mapping` resource for the proxy's JWT to virtual key mappings, so JWT clients identified by a claim (`client_id`, `azp`, `sub`) map to virtual keys and inherit their models, budgets and rate limits. Supports `description` and `is_active`, rotating the mapped key in place, and forces replacement when the claim name or value changes
|
||||
- **team**: `soft_budget`, `tags`, and `soft_budget_alerting_emails` attributes on `litellm_team`, matching what `/team/new` and `/team/update` already accept; `soft_budget_alerting_emails` is sent under `metadata`, where the proxy reads it
|
||||
- **user**: New `litellm_user` resource and `litellm_user` / `litellm_users` data sources for managing internal users
|
||||
- **budget**: New `litellm_budget` resource and `litellm_budget` / `litellm_budgets` data sources for reusable budget objects
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ This Terraform provider allows you to manage LiteLLM resources through Infrastru
|
|||
|
||||
## Source of truth
|
||||
|
||||
This directory (`terraform/provider/` in [BerriAI/litellm](https://github.com/BerriAI/litellm)) is the source of truth for the provider. [BerriAI/terraform-provider-litellm](https://github.com/BerriAI/terraform-provider-litellm) is a thin release mirror that the public Terraform Registry ingests from; do not open PRs there. Changes land here, where CI builds the provider, runs its tests, and statically audits every endpoint the provider calls against the proxy's generated OpenAPI schema (`tools/endpointaudit/`), so the provider cannot drift from the LiteLLM API silently. Releases are published by mirroring this directory into the split repo and tagging it, which triggers the goreleaser workflow there (see `RELEASING.md`)
|
||||
This directory (`terraform/provider/` in [BerriAI/litellm](https://github.com/BerriAI/litellm)) is the source of truth for the provider. [BerriAI/terraform-provider-litellm](https://github.com/BerriAI/terraform-provider-litellm) is a thin release mirror that the public Terraform Registry ingests from; do not open PRs there. Changes land here, where CI builds the provider, runs its tests, and statically audits every endpoint the provider calls against the proxy's generated OpenAPI schema (`tools/endpointaudit/`), so the provider cannot drift from the LiteLLM API silently. The same audit runs in reverse as a coverage gate: every management endpoint in the schema must be covered by a resource or data source, or carry a documented entry in `tools/endpointaudit/coverage_allowlist.txt`, and stale allowlist entries fail CI. Releases are published by mirroring this directory into the split repo and tagging it, which triggers the goreleaser workflow there (see `RELEASING.md`)
|
||||
|
||||
## Versioning
|
||||
|
||||
|
|
@ -151,6 +151,7 @@ For full details on the <code>litellm_key</code> resource, see the [key resource
|
|||
- <code>litellm_mcp_server</code>: Manage MCP (Model Context Protocol) servers. [Documentation](docs/resources/mcp_server.md)
|
||||
- <code>litellm_credential</code>: Manage credentials for secure authentication. [Documentation](docs/resources/credential.md)
|
||||
- <code>litellm_vector_store</code>: Manage vector stores for embeddings and RAG. [Documentation](docs/resources/vector_store.md)
|
||||
- <code>litellm_jwt_key_mapping</code>: Map JWT claim values to virtual keys for per-client budgets and limits. [Documentation](docs/resources/jwt_key_mapping.md)
|
||||
|
||||
### Available Data Sources
|
||||
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@ The LiteLLM provider supports the following resources:
|
|||
* [`litellm_mcp_server`](./resources/mcp_server) - Manage MCP (Model Context Protocol) servers
|
||||
* [`litellm_credential`](./resources/credential) - Manage credentials for various providers
|
||||
* [`litellm_vector_store`](./resources/vector_store) - Manage vector stores
|
||||
* [`litellm_jwt_key_mapping`](./resources/jwt_key_mapping) - Map JWT claim values to virtual keys
|
||||
|
||||
## Available Data Sources
|
||||
|
||||
|
|
|
|||
94
terraform/provider/docs/resources/jwt_key_mapping.md
Normal file
94
terraform/provider/docs/resources/jwt_key_mapping.md
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
# litellm_jwt_key_mapping
|
||||
|
||||
Maps a JWT claim value to a LiteLLM virtual key. Every JWT client identified by a claim, typically `client_id`, `azp` or `sub`, then gets the model restrictions, budgets, rate limits, guardrails and spend tracking of the virtual key it maps to, without that key ever being handed to the client.
|
||||
|
||||
The mappings only take effect once JWT auth is enabled on the proxy, which is configuration rather than API state:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
enable_jwt_auth: True
|
||||
litellm_jwtauth:
|
||||
virtual_key_claim_field: "client_id"
|
||||
unregistered_jwt_client_behavior: "fallback_team_mapping"
|
||||
```
|
||||
|
||||
See [JWT to virtual key mapping](https://docs.litellm.ai/docs/proxy/jwt_key_mapping) for the proxy side of the feature
|
||||
|
||||
## Example Usage
|
||||
|
||||
The mapped virtual key has to exist already and its value has to be known to Terraform, so it comes from a variable or a secret manager rather than from a `litellm_key` resource. `litellm_key` deliberately made its generated `key` write-only, to avoid storing raw API keys in state, so referencing it here does not merely read back null: Terraform's write-only enforcement turns `key = litellm_key.foo.key` into a static `Missing required argument` error at `terraform plan`, before any API call, in every apply ordering, including a first apply where both resources are created together:
|
||||
|
||||
```hcl
|
||||
variable "alice_key" {
|
||||
type = string
|
||||
sensitive = true
|
||||
}
|
||||
|
||||
resource "litellm_jwt_key_mapping" "alice" {
|
||||
jwt_claim_name = "client_id"
|
||||
jwt_claim_value = "dev-alice"
|
||||
key = var.alice_key
|
||||
}
|
||||
```
|
||||
|
||||
Per-client limits live on the virtual key, so one mapping per client is how each JWT client gets its own budget and quota:
|
||||
|
||||
```hcl
|
||||
resource "litellm_jwt_key_mapping" "billing_service" {
|
||||
jwt_claim_name = "client_id"
|
||||
jwt_claim_value = "billing-service"
|
||||
key = var.billing_service_key
|
||||
description = "Billing service JWT client"
|
||||
is_active = true
|
||||
}
|
||||
```
|
||||
|
||||
Several clients at once, with the key values coming from a map of secrets:
|
||||
|
||||
```hcl
|
||||
variable "jwt_client_keys" {
|
||||
type = map(string)
|
||||
sensitive = true
|
||||
}
|
||||
|
||||
resource "litellm_jwt_key_mapping" "developer" {
|
||||
for_each = var.jwt_client_keys
|
||||
|
||||
jwt_claim_name = "client_id"
|
||||
jwt_claim_value = each.key
|
||||
key = each.value
|
||||
description = "Developer JWT client ${each.key}"
|
||||
}
|
||||
```
|
||||
|
||||
## Argument Reference
|
||||
|
||||
- `jwt_claim_name` - (Required, ForceNew) Name of the JWT claim to match on, for example `client_id`, `azp` or `sub`. Must match `virtual_key_claim_field` in the proxy JWT config
|
||||
- `jwt_claim_value` - (Required, ForceNew) Value of the claim identifying the JWT client. Unique together with `jwt_claim_name`, so a second mapping for the same pair fails with a 409
|
||||
- `key` - (Required, Sensitive) The virtual key this claim value maps to. It has to exist already, otherwise the proxy rejects the mapping with `The provided key does not match an existing virtual key`
|
||||
- `description` - (Optional) Description of the mapping
|
||||
- `is_active` - (Optional) Whether the mapping is active. Inactive mappings are ignored during JWT auth. Defaults to `true`
|
||||
|
||||
## Attribute Reference
|
||||
|
||||
- `id` - The mapping ID assigned by LiteLLM
|
||||
- `created_at` - Timestamp when the mapping was created
|
||||
- `updated_at` - Timestamp when the mapping was last updated
|
||||
- `created_by` - User who created the mapping
|
||||
- `updated_by` - User who last updated the mapping
|
||||
|
||||
## Notes
|
||||
|
||||
The proxy stores only a hash of `key` and never returns it, so drift on that attribute cannot be detected and Terraform tracks the value from your configuration. Changing `key` rotates the mapping onto the new virtual key in place, with no replacement. Like the other secrets this provider accepts, such as `credential_values` and `model_api_key`, the configured value is kept in state, so treat the state as sensitive
|
||||
|
||||
Only proxy admins can create, update or delete mappings, so the provider `api_key` has to be a master key or an admin key
|
||||
|
||||
## Import
|
||||
|
||||
Mappings are imported by their mapping ID:
|
||||
|
||||
```shell
|
||||
terraform import litellm_jwt_key_mapping.alice 297a5536-1aeb-4cf1-b666-b3809c2750a8
|
||||
```
|
||||
|
||||
Because the API does not return the mapped key, `key` is empty in state right after an import, so the first plan shows an in-place update that pushes the configured key back to the proxy. That update is harmless, the proxy just rehashes the same value when the key has not actually changed
|
||||
|
|
@ -19,6 +19,7 @@ func Provider() *schema.Provider {
|
|||
"litellm_mcp_server": resourceLiteLLMMCPServer(),
|
||||
"litellm_credential": resourceLiteLLMCredential(),
|
||||
"litellm_vector_store": resourceLiteLLMVectorStore(),
|
||||
"litellm_jwt_key_mapping": resourceLiteLLMJWTKeyMapping(),
|
||||
"litellm_fallback": resourceLiteLLMFallback(),
|
||||
"litellm_key_block": resourceLiteLLMKeyBlock(),
|
||||
"litellm_team_block": resourceLiteLLMTeamBlock(),
|
||||
|
|
|
|||
70
terraform/provider/litellm/resource_jwt_key_mapping.go
Normal file
70
terraform/provider/litellm/resource_jwt_key_mapping.go
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
package litellm
|
||||
|
||||
import (
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
)
|
||||
|
||||
func resourceLiteLLMJWTKeyMapping() *schema.Resource {
|
||||
return &schema.Resource{
|
||||
Create: resourceLiteLLMJWTKeyMappingCreate,
|
||||
Read: resourceLiteLLMJWTKeyMappingRead,
|
||||
Update: resourceLiteLLMJWTKeyMappingUpdate,
|
||||
Delete: resourceLiteLLMJWTKeyMappingDelete,
|
||||
|
||||
Importer: &schema.ResourceImporter{
|
||||
StateContext: schema.ImportStatePassthroughContext,
|
||||
},
|
||||
|
||||
Schema: map[string]*schema.Schema{
|
||||
"jwt_claim_name": {
|
||||
Type: schema.TypeString,
|
||||
Required: true,
|
||||
ForceNew: true,
|
||||
Description: "Name of the JWT claim to match on, for example client_id, azp or sub. Must match virtual_key_claim_field in the proxy JWT config",
|
||||
},
|
||||
"jwt_claim_value": {
|
||||
Type: schema.TypeString,
|
||||
Required: true,
|
||||
ForceNew: true,
|
||||
Description: "Value of the claim identifying the JWT client. Unique together with jwt_claim_name",
|
||||
},
|
||||
"key": {
|
||||
Type: schema.TypeString,
|
||||
Required: true,
|
||||
Sensitive: true,
|
||||
Description: "The virtual key this claim value maps to. The proxy stores only a hash of it and never returns it, so drift on this attribute cannot be detected and Terraform tracks the configured value",
|
||||
},
|
||||
"description": {
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
Description: "Description of the mapping",
|
||||
},
|
||||
"is_active": {
|
||||
Type: schema.TypeBool,
|
||||
Optional: true,
|
||||
Default: true,
|
||||
Description: "Whether the mapping is active. Inactive mappings are ignored during JWT auth",
|
||||
},
|
||||
"created_at": {
|
||||
Type: schema.TypeString,
|
||||
Computed: true,
|
||||
Description: "Timestamp when the mapping was created",
|
||||
},
|
||||
"updated_at": {
|
||||
Type: schema.TypeString,
|
||||
Computed: true,
|
||||
Description: "Timestamp when the mapping was last updated",
|
||||
},
|
||||
"created_by": {
|
||||
Type: schema.TypeString,
|
||||
Computed: true,
|
||||
Description: "User who created the mapping",
|
||||
},
|
||||
"updated_by": {
|
||||
Type: schema.TypeString,
|
||||
Computed: true,
|
||||
Description: "User who last updated the mapping",
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
186
terraform/provider/litellm/resource_jwt_key_mapping_crud.go
Normal file
186
terraform/provider/litellm/resource_jwt_key_mapping_crud.go
Normal file
|
|
@ -0,0 +1,186 @@
|
|||
package litellm
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
)
|
||||
|
||||
const jwtKeyMappingNotFound = "jwt_key_mapping_not_found"
|
||||
|
||||
func resourceLiteLLMJWTKeyMappingCreate(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
|
||||
createRequest := JWTKeyMappingRequest{
|
||||
JWTClaimName: d.Get("jwt_claim_name").(string),
|
||||
JWTClaimValue: d.Get("jwt_claim_value").(string),
|
||||
Key: d.Get("key").(string),
|
||||
Description: d.Get("description").(string),
|
||||
}
|
||||
|
||||
resp, err := MakeRequest(client, "POST", "/jwt/key/mapping/new", createRequest)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create JWT key mapping: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
var mapping JWTKeyMappingResponse
|
||||
if err := handleJWTKeyMappingAPIResponse(resp, &mapping, client); err != nil {
|
||||
return fmt.Errorf("failed to create JWT key mapping: %w", err)
|
||||
}
|
||||
|
||||
if mapping.ID == "" {
|
||||
return fmt.Errorf("failed to create JWT key mapping: the proxy returned no mapping id")
|
||||
}
|
||||
|
||||
d.SetId(mapping.ID)
|
||||
|
||||
// The create endpoint has no is_active field and always activates the
|
||||
// mapping, so a JWT client matching this claim can authenticate during
|
||||
// the gap before the deactivation call below runs. If deactivation
|
||||
// itself fails, delete the mapping rather than leaving it active and
|
||||
// unmanaged indefinitely.
|
||||
if !d.Get("is_active").(bool) {
|
||||
if err := updateJWTKeyMapping(d, client); err != nil {
|
||||
if deleteErr := deleteJWTKeyMapping(mapping.ID, client); deleteErr != nil {
|
||||
return fmt.Errorf(
|
||||
"JWT key mapping %s was created active and could not be deactivated (%v); it also could not be deleted and remains active on the proxy, remove it manually via POST /jwt/key/mapping/delete: %v",
|
||||
mapping.ID, err, deleteErr,
|
||||
)
|
||||
}
|
||||
d.SetId("")
|
||||
return fmt.Errorf("JWT key mapping was created active but could not be deactivated, so it was deleted instead: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return resourceLiteLLMJWTKeyMappingRead(d, m)
|
||||
}
|
||||
|
||||
func resourceLiteLLMJWTKeyMappingRead(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
|
||||
resp, err := MakeRequest(client, "GET", fmt.Sprintf("/jwt/key/mapping/info?id=%s", url.QueryEscape(d.Id())), nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read JWT key mapping: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
var mapping JWTKeyMappingResponse
|
||||
if err := handleJWTKeyMappingAPIResponse(resp, &mapping, client); err != nil {
|
||||
if err.Error() == jwtKeyMappingNotFound {
|
||||
d.SetId("")
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("failed to read JWT key mapping: %w", err)
|
||||
}
|
||||
|
||||
d.SetId(mapping.ID)
|
||||
d.Set("jwt_claim_name", mapping.JWTClaimName)
|
||||
d.Set("jwt_claim_value", mapping.JWTClaimValue)
|
||||
d.Set("description", mapping.Description)
|
||||
d.Set("is_active", mapping.IsActive)
|
||||
d.Set("created_at", mapping.CreatedAt)
|
||||
d.Set("updated_at", mapping.UpdatedAt)
|
||||
d.Set("created_by", mapping.CreatedBy)
|
||||
d.Set("updated_by", mapping.UpdatedBy)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func resourceLiteLLMJWTKeyMappingUpdate(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
|
||||
oldKey, _ := d.GetChange("key")
|
||||
oldDescription, _ := d.GetChange("description")
|
||||
oldIsActive, _ := d.GetChange("is_active")
|
||||
|
||||
if err := updateJWTKeyMapping(d, client); err != nil {
|
||||
// The update is a single atomic API call: on failure nothing changed
|
||||
// server-side. Revert every field the update could have changed before
|
||||
// attempting to resync, so a failed refresh can't leave the rejected
|
||||
// values persisted into state.
|
||||
d.Set("key", oldKey)
|
||||
d.Set("description", oldDescription)
|
||||
d.Set("is_active", oldIsActive)
|
||||
if readErr := resourceLiteLLMJWTKeyMappingRead(d, m); readErr != nil {
|
||||
return fmt.Errorf("failed to update JWT key mapping: %w (and failed to refresh state afterward: %v)", err, readErr)
|
||||
}
|
||||
return fmt.Errorf("failed to update JWT key mapping: %w", err)
|
||||
}
|
||||
|
||||
return resourceLiteLLMJWTKeyMappingRead(d, m)
|
||||
}
|
||||
|
||||
func resourceLiteLLMJWTKeyMappingDelete(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
|
||||
if err := deleteJWTKeyMapping(d.Id(), client); err != nil {
|
||||
return fmt.Errorf("failed to delete JWT key mapping: %w", err)
|
||||
}
|
||||
|
||||
d.SetId("")
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteJWTKeyMapping(id string, client *Client) error {
|
||||
resp, err := MakeRequest(client, "POST", "/jwt/key/mapping/delete", JWTKeyMappingDeleteRequest{ID: id})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if err := handleJWTKeyMappingAPIResponse(resp, nil, client); err != nil {
|
||||
if err.Error() != jwtKeyMappingNotFound {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateJWTKeyMapping(d *schema.ResourceData, client *Client) error {
|
||||
updateRequest := JWTKeyMappingUpdateRequest{
|
||||
ID: d.Id(),
|
||||
Key: d.Get("key").(string),
|
||||
Description: d.Get("description").(string),
|
||||
IsActive: d.Get("is_active").(bool),
|
||||
}
|
||||
|
||||
resp, err := MakeRequest(client, "POST", "/jwt/key/mapping/update", updateRequest)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
return handleJWTKeyMappingAPIResponse(resp, nil, client)
|
||||
}
|
||||
|
||||
func handleJWTKeyMappingAPIResponse(resp *http.Response, result interface{}, client *Client) error {
|
||||
bodyBytes, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read response body: %v", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return fmt.Errorf(jwtKeyMappingNotFound)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
|
||||
return fmt.Errorf("API request failed: Status: %s, Response: %s",
|
||||
resp.Status, client.redactSensitiveData(string(bodyBytes)))
|
||||
}
|
||||
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(bodyBytes, result); err != nil {
|
||||
return fmt.Errorf("failed to parse response: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
630
terraform/provider/litellm/resource_jwt_key_mapping_crud_test.go
Normal file
630
terraform/provider/litellm/resource_jwt_key_mapping_crud_test.go
Normal file
|
|
@ -0,0 +1,630 @@
|
|||
package litellm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/terraform"
|
||||
)
|
||||
|
||||
// resourceDataWithChange builds a ResourceData carrying a real diff between
|
||||
// prior state and new config, so d.GetChange reflects true old/new values.
|
||||
// schema.TestResourceDataRaw diffs against a nil prior state, which collapses
|
||||
// GetChange's old side to the zero value and can't exercise this.
|
||||
func resourceDataWithChange(t *testing.T, oldAttrs map[string]string, newRaw map[string]interface{}) *schema.ResourceData {
|
||||
t.Helper()
|
||||
|
||||
sm := schema.InternalMap(resourceLiteLLMJWTKeyMapping().Schema)
|
||||
state := &terraform.InstanceState{ID: oldAttrs["id"], Attributes: oldAttrs}
|
||||
config := terraform.NewResourceConfigRaw(newRaw)
|
||||
|
||||
diff, err := sm.Diff(context.Background(), state, config, nil, nil, true)
|
||||
if err != nil {
|
||||
t.Fatalf("diff: %v", err)
|
||||
}
|
||||
d, err := sm.Data(state, diff)
|
||||
if err != nil {
|
||||
t.Fatalf("data: %v", err)
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
type jwtKeyMappingCall struct {
|
||||
Method string
|
||||
Path string
|
||||
Query string
|
||||
Body map[string]interface{}
|
||||
}
|
||||
|
||||
func jwtKeyMappingTestServer(t *testing.T, mapping JWTKeyMappingResponse) (*httptest.Server, *[]jwtKeyMappingCall) {
|
||||
t.Helper()
|
||||
|
||||
calls := make([]jwtKeyMappingCall, 0)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
body := map[string]interface{}{}
|
||||
if r.Body != nil {
|
||||
_ = json.NewDecoder(r.Body).Decode(&body)
|
||||
}
|
||||
calls = append(calls, jwtKeyMappingCall{Method: r.Method, Path: r.URL.Path, Query: r.URL.RawQuery, Body: body})
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/jwt/key/mapping/delete":
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"status": "success"})
|
||||
default:
|
||||
_ = json.NewEncoder(w).Encode(mapping)
|
||||
}
|
||||
}))
|
||||
|
||||
return srv, &calls
|
||||
}
|
||||
|
||||
func jwtKeyMappingFixture() JWTKeyMappingResponse {
|
||||
return JWTKeyMappingResponse{
|
||||
ID: "map-abc-123",
|
||||
JWTClaimName: "client_id",
|
||||
JWTClaimValue: "dev-alice",
|
||||
Description: "dev-alice",
|
||||
IsActive: true,
|
||||
CreatedAt: "2026-08-06T10:00:00Z",
|
||||
UpdatedAt: "2026-08-06T11:00:00Z",
|
||||
CreatedBy: "admin",
|
||||
UpdatedBy: "admin",
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateSendsClaimAndKey(t *testing.T) {
|
||||
srv, calls := jwtKeyMappingTestServer(t, jwtKeyMappingFixture())
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
"description": "dev-alice",
|
||||
"is_active": true,
|
||||
})
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingCreate(d, client); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
if d.Id() != "map-abc-123" {
|
||||
t.Fatalf("expected id from the API response, got %q", d.Id())
|
||||
}
|
||||
|
||||
create := (*calls)[0]
|
||||
if create.Method != "POST" || create.Path != "/jwt/key/mapping/new" {
|
||||
t.Fatalf("expected POST /jwt/key/mapping/new, got %s %s", create.Method, create.Path)
|
||||
}
|
||||
if create.Body["jwt_claim_name"] != "client_id" || create.Body["jwt_claim_value"] != "dev-alice" {
|
||||
t.Fatalf("claim fields not sent: %v", create.Body)
|
||||
}
|
||||
if create.Body["key"] != "sk-abc123" {
|
||||
t.Fatalf("virtual key not sent: %v", create.Body["key"])
|
||||
}
|
||||
if create.Body["description"] != "dev-alice" {
|
||||
t.Fatalf("description not sent: %v", create.Body["description"])
|
||||
}
|
||||
if _, sent := create.Body["is_active"]; sent {
|
||||
t.Fatalf("is_active is not accepted by /jwt/key/mapping/new but was sent: %v", create.Body)
|
||||
}
|
||||
|
||||
for _, call := range (*calls)[1:] {
|
||||
if call.Path == "/jwt/key/mapping/update" {
|
||||
t.Fatalf("an active mapping must not trigger a follow-up update")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateOmitsEmptyDescription(t *testing.T) {
|
||||
srv, calls := jwtKeyMappingTestServer(t, jwtKeyMappingFixture())
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
"is_active": true,
|
||||
})
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingCreate(d, client); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
if _, sent := (*calls)[0].Body["description"]; sent {
|
||||
t.Fatalf("unset description should be omitted: %v", (*calls)[0].Body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateDeactivatesWhenNotActive(t *testing.T) {
|
||||
mapping := jwtKeyMappingFixture()
|
||||
mapping.IsActive = false
|
||||
srv, calls := jwtKeyMappingTestServer(t, mapping)
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
"is_active": false,
|
||||
})
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingCreate(d, client); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
var update *jwtKeyMappingCall
|
||||
for i := range *calls {
|
||||
if (*calls)[i].Path == "/jwt/key/mapping/update" {
|
||||
update = &(*calls)[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if update == nil {
|
||||
t.Fatal("expected a follow-up update, since the create endpoint always starts a mapping active")
|
||||
}
|
||||
if update.Body["id"] != "map-abc-123" {
|
||||
t.Fatalf("update must target the new mapping, got %v", update.Body["id"])
|
||||
}
|
||||
if update.Body["is_active"] != false {
|
||||
t.Fatalf("expected is_active false in the follow-up update, got %v", update.Body["is_active"])
|
||||
}
|
||||
if d.Get("is_active").(bool) {
|
||||
t.Fatal("state should reflect the inactive mapping after create")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateDeletesMappingWhenDeactivationFails(t *testing.T) {
|
||||
// Regression test: the create endpoint has no is_active field and always
|
||||
// activates the mapping, so a failed deactivation used to leave that
|
||||
// mapping active and unmanaged indefinitely. It must be deleted instead.
|
||||
calls := make([]jwtKeyMappingCall, 0)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
body := map[string]interface{}{}
|
||||
if r.Body != nil {
|
||||
_ = json.NewDecoder(r.Body).Decode(&body)
|
||||
}
|
||||
calls = append(calls, jwtKeyMappingCall{Method: r.Method, Path: r.URL.Path, Query: r.URL.RawQuery, Body: body})
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/jwt/key/mapping/new":
|
||||
_ = json.NewEncoder(w).Encode(jwtKeyMappingFixture())
|
||||
case "/jwt/key/mapping/update":
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"detail": "proxy unavailable"})
|
||||
case "/jwt/key/mapping/delete":
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"status": "success"})
|
||||
default:
|
||||
t.Fatalf("unexpected request to %s", r.URL.Path)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
"is_active": false,
|
||||
})
|
||||
|
||||
err := resourceLiteLLMJWTKeyMappingCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected the failed deactivation to surface as an error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "deleted instead") {
|
||||
t.Fatalf("expected the error to explain the mapping was deleted, got %v", err)
|
||||
}
|
||||
|
||||
deleteCalls := 0
|
||||
for _, c := range calls {
|
||||
if c.Path == "/jwt/key/mapping/delete" {
|
||||
deleteCalls++
|
||||
if c.Body["id"] != "map-abc-123" {
|
||||
t.Fatalf("delete must target the mapping that could not be deactivated, got %v", c.Body["id"])
|
||||
}
|
||||
}
|
||||
}
|
||||
if deleteCalls != 1 {
|
||||
t.Fatalf("expected exactly one cleanup delete call, got %d", deleteCalls)
|
||||
}
|
||||
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("a successfully deleted mapping must not remain in state, got id %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateReportsWhenDeactivationAndDeleteBothFail(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/jwt/key/mapping/new":
|
||||
_ = json.NewEncoder(w).Encode(jwtKeyMappingFixture())
|
||||
default:
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"detail": "proxy unavailable"})
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
"is_active": false,
|
||||
})
|
||||
|
||||
err := resourceLiteLLMJWTKeyMappingCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected an error when both deactivation and the cleanup delete fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "remove it manually") {
|
||||
t.Fatalf("expected the error to demand manual cleanup, got %v", err)
|
||||
}
|
||||
|
||||
// The mapping is still active on the proxy since neither call succeeded, so
|
||||
// the id must stay in state: the next apply taints and retries the delete,
|
||||
// rather than Terraform losing track of a live, active mapping entirely.
|
||||
if d.Id() != "map-abc-123" {
|
||||
t.Fatalf("expected the id to remain in state so a retry can find it, got %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingUpdateRevertsDescriptionAndIsActiveWhenTheRecoveryReadAlsoFails(t *testing.T) {
|
||||
// Regression test: on a failed update, only `key` was being reverted
|
||||
// before Read ran. If Read itself then failed too (network blip, proxy
|
||||
// hiccup), description/is_active kept the rejected, never-applied values,
|
||||
// and Terraform could persist them as if the update had succeeded.
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/jwt/key/mapping/update":
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"detail": "rejected"})
|
||||
case "/jwt/key/mapping/info":
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"detail": "proxy unavailable"})
|
||||
default:
|
||||
t.Fatalf("unexpected request to %s", r.URL.Path)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
d := resourceDataWithChange(t,
|
||||
map[string]string{
|
||||
"id": "map-abc-123",
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-old-key-0000000000",
|
||||
"description": "old description",
|
||||
"is_active": "true",
|
||||
},
|
||||
map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-old-key-0000000000",
|
||||
"description": "attempted new description",
|
||||
"is_active": false,
|
||||
},
|
||||
)
|
||||
d.SetId("map-abc-123")
|
||||
|
||||
err := resourceLiteLLMJWTKeyMappingUpdate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected the update failure to surface as an error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed to refresh state afterward") {
|
||||
t.Fatalf("expected the error to mention the failed recovery read, got %v", err)
|
||||
}
|
||||
|
||||
if d.Get("description").(string) != "old description" {
|
||||
t.Fatalf("a rejected description must not survive when the recovery read also fails, got %q", d.Get("description").(string))
|
||||
}
|
||||
if d.Get("is_active").(bool) != true {
|
||||
t.Fatalf("a rejected is_active must not survive when the recovery read also fails, got %v", d.Get("is_active").(bool))
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingReadPopulatesStateAndKeepsKey(t *testing.T) {
|
||||
srv, calls := jwtKeyMappingTestServer(t, jwtKeyMappingFixture())
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-configured-value",
|
||||
})
|
||||
d.SetId("map-abc-123")
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingRead(d, client); err != nil {
|
||||
t.Fatalf("read failed: %v", err)
|
||||
}
|
||||
|
||||
read := (*calls)[0]
|
||||
if read.Method != "GET" || read.Path != "/jwt/key/mapping/info" {
|
||||
t.Fatalf("expected GET /jwt/key/mapping/info, got %s %s", read.Method, read.Path)
|
||||
}
|
||||
if read.Query != "id=map-abc-123" {
|
||||
t.Fatalf("expected the mapping id in the query, got %q", read.Query)
|
||||
}
|
||||
|
||||
if d.Get("jwt_claim_value").(string) != "dev-alice" {
|
||||
t.Fatalf("claim value not populated: %q", d.Get("jwt_claim_value").(string))
|
||||
}
|
||||
if d.Get("description").(string) != "dev-alice" {
|
||||
t.Fatalf("description not populated: %q", d.Get("description").(string))
|
||||
}
|
||||
if !d.Get("is_active").(bool) {
|
||||
t.Fatal("is_active not populated")
|
||||
}
|
||||
if d.Get("created_at").(string) != "2026-08-06T10:00:00Z" || d.Get("created_by").(string) != "admin" {
|
||||
t.Fatalf("computed audit fields not populated: %v", d.State().Attributes)
|
||||
}
|
||||
if d.Get("key").(string) != "sk-configured-value" {
|
||||
t.Fatalf("the API never returns the key, so the configured value must survive a read, got %q", d.Get("key").(string))
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingReadClearsIDWhenMappingIsGone(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"detail": "Mapping not found"})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
})
|
||||
d.SetId("map-gone")
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingRead(d, client); err != nil {
|
||||
t.Fatalf("a deleted mapping must not fail the read: %v", err)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("expected the id to be cleared so Terraform plans a recreate, got %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingUpdateClearsDescriptionAndSendsKey(t *testing.T) {
|
||||
mapping := jwtKeyMappingFixture()
|
||||
mapping.Description = ""
|
||||
srv, calls := jwtKeyMappingTestServer(t, mapping)
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-rotated",
|
||||
"is_active": true,
|
||||
})
|
||||
d.SetId("map-abc-123")
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingUpdate(d, client); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
|
||||
update := (*calls)[0]
|
||||
if update.Method != "POST" || update.Path != "/jwt/key/mapping/update" {
|
||||
t.Fatalf("expected POST /jwt/key/mapping/update, got %s %s", update.Method, update.Path)
|
||||
}
|
||||
if update.Body["id"] != "map-abc-123" {
|
||||
t.Fatalf("update must carry the mapping id, got %v", update.Body["id"])
|
||||
}
|
||||
if update.Body["key"] != "sk-rotated" {
|
||||
t.Fatalf("rotated key not sent: %v", update.Body["key"])
|
||||
}
|
||||
description, sent := update.Body["description"]
|
||||
if !sent || description != "" {
|
||||
t.Fatalf("a dropped description must be sent as an empty string, since the proxy ignores absent fields: %v", update.Body)
|
||||
}
|
||||
if d.Get("description").(string) != "" {
|
||||
t.Fatalf("description should be cleared in state, got %q", d.Get("description").(string))
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingUpdateRevertsKeyOnFailureAndResyncsRest(t *testing.T) {
|
||||
// Regression test for a live-verified bug: Terraform's classic SDKv2 CRUD
|
||||
// model persists ResourceData's diff-applied (attempted) values to state
|
||||
// even when the callback returns an error, unless the provider reverts
|
||||
// them explicitly. Confirmed live: a rejected key rotation left the new,
|
||||
// never-applied key in `terraform state pull` while the proxy kept the
|
||||
// old one, so the next plan falsely reported convergence.
|
||||
calls := make([]jwtKeyMappingCall, 0)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
body := map[string]interface{}{}
|
||||
if r.Body != nil {
|
||||
_ = json.NewDecoder(r.Body).Decode(&body)
|
||||
}
|
||||
calls = append(calls, jwtKeyMappingCall{Method: r.Method, Path: r.URL.Path, Query: r.URL.RawQuery, Body: body})
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/jwt/key/mapping/update":
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"detail": "The provided key does not match an existing virtual key.",
|
||||
})
|
||||
case "/jwt/key/mapping/info":
|
||||
// Server truth: unchanged, since the rejected update above never applied.
|
||||
_ = json.NewEncoder(w).Encode(jwtKeyMappingFixture())
|
||||
default:
|
||||
t.Fatalf("unexpected request to %s", r.URL.Path)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
d := resourceDataWithChange(t,
|
||||
map[string]string{
|
||||
"id": "map-abc-123",
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-old-key-0000000000",
|
||||
"description": "dev-alice",
|
||||
"is_active": "true",
|
||||
},
|
||||
map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-rejected-new-key-00",
|
||||
"description": "attempted new description",
|
||||
"is_active": false,
|
||||
},
|
||||
)
|
||||
d.SetId("map-abc-123")
|
||||
|
||||
err := resourceLiteLLMJWTKeyMappingUpdate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected the rejected key to fail the update")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "does not match an existing virtual key") {
|
||||
t.Fatalf("expected the proxy's rejection reason in the error, got %v", err)
|
||||
}
|
||||
|
||||
if d.Get("key").(string) != "sk-old-key-0000000000" {
|
||||
t.Fatalf("a failed update must not persist the rejected key into state, got %q", d.Get("key").(string))
|
||||
}
|
||||
if d.Get("description").(string) != "dev-alice" {
|
||||
t.Fatalf("a failed update must resync description from the server, got %q", d.Get("description").(string))
|
||||
}
|
||||
if d.Get("is_active").(bool) != true {
|
||||
t.Fatalf("a failed update must resync is_active from the server, got %v", d.Get("is_active").(bool))
|
||||
}
|
||||
|
||||
readCalls := 0
|
||||
for _, c := range calls {
|
||||
if c.Path == "/jwt/key/mapping/info" {
|
||||
readCalls++
|
||||
}
|
||||
}
|
||||
if readCalls != 1 {
|
||||
t.Fatalf("expected exactly one read to resync state after the failed update, got %d", readCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingUpdateOmitsMissingKeyRatherThanBlankingIt(t *testing.T) {
|
||||
srv, calls := jwtKeyMappingTestServer(t, jwtKeyMappingFixture())
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"description": "dev-alice",
|
||||
"is_active": true,
|
||||
})
|
||||
d.SetId("map-abc-123")
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingUpdate(d, client); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
|
||||
if _, sent := (*calls)[0].Body["key"]; sent {
|
||||
t.Fatalf("a missing key must be omitted rather than blanking the mapping token: %v", (*calls)[0].Body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingDeleteToleratesMissingMapping(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"detail": "Mapping not found"})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
})
|
||||
d.SetId("map-already-gone")
|
||||
|
||||
if err := resourceLiteLLMJWTKeyMappingDelete(d, client); err != nil {
|
||||
t.Fatalf("deleting an already deleted mapping must succeed: %v", err)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("expected the id to be cleared after delete, got %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateSurfacesDuplicateClaimError(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusConflict)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"detail": "A mapping for claim 'client_id' = 'dev-alice' already exists.",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-abc123",
|
||||
"is_active": true,
|
||||
})
|
||||
|
||||
err := resourceLiteLLMJWTKeyMappingCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected a duplicate claim pair to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "already exists") {
|
||||
t.Fatalf("the proxy explanation must reach the user, got %v", err)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("no id should be recorded for a failed create, got %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTKeyMappingCreateDoesNotLeakKeyInErrors(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"key": "sk-super-secret",
|
||||
"detail": "The provided key does not match an existing virtual key.",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMJWTKeyMapping().Schema, map[string]interface{}{
|
||||
"jwt_claim_name": "client_id",
|
||||
"jwt_claim_value": "dev-alice",
|
||||
"key": "sk-super-secret",
|
||||
"is_active": true,
|
||||
})
|
||||
|
||||
err := resourceLiteLLMJWTKeyMappingCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected an unknown virtual key to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "does not match an existing virtual key") {
|
||||
t.Fatalf("the proxy explanation must reach the user, got %v", err)
|
||||
}
|
||||
if strings.Contains(err.Error(), "sk-super-secret") {
|
||||
t.Fatalf("the virtual key must be redacted in errors, got %v", err)
|
||||
}
|
||||
}
|
||||
|
|
@ -272,3 +272,33 @@ type VectorStoreDeleteRequest struct {
|
|||
type VectorStoreInfoRequest struct {
|
||||
VectorStoreID string `json:"vector_store_id"`
|
||||
}
|
||||
|
||||
type JWTKeyMappingRequest struct {
|
||||
JWTClaimName string `json:"jwt_claim_name"`
|
||||
JWTClaimValue string `json:"jwt_claim_value"`
|
||||
Key string `json:"key"`
|
||||
Description string `json:"description,omitempty"`
|
||||
}
|
||||
|
||||
type JWTKeyMappingUpdateRequest struct {
|
||||
ID string `json:"id"`
|
||||
Key string `json:"key,omitempty"`
|
||||
Description string `json:"description"`
|
||||
IsActive bool `json:"is_active"`
|
||||
}
|
||||
|
||||
type JWTKeyMappingDeleteRequest struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
|
||||
type JWTKeyMappingResponse struct {
|
||||
ID string `json:"id"`
|
||||
JWTClaimName string `json:"jwt_claim_name"`
|
||||
JWTClaimValue string `json:"jwt_claim_value"`
|
||||
Description string `json:"description,omitempty"`
|
||||
IsActive bool `json:"is_active"`
|
||||
CreatedAt string `json:"created_at,omitempty"`
|
||||
UpdatedAt string `json:"updated_at,omitempty"`
|
||||
CreatedBy string `json:"created_by,omitempty"`
|
||||
UpdatedBy string `json:"updated_by,omitempty"`
|
||||
}
|
||||
|
|
|
|||
106
terraform/provider/tools/endpointaudit/coverage.go
Normal file
106
terraform/provider/tools/endpointaudit/coverage.go
Normal file
|
|
@ -0,0 +1,106 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var managementPrefixes = map[string]bool{
|
||||
"access_group": true,
|
||||
"agent": true,
|
||||
"budget": true,
|
||||
"cache": true,
|
||||
"config": true,
|
||||
"coordination_redis": true,
|
||||
"credentials": true,
|
||||
"customer": true,
|
||||
"fallback": true,
|
||||
"guardrails": true,
|
||||
"jwt": true,
|
||||
"key": true,
|
||||
"model": true,
|
||||
"organization": true,
|
||||
"project": true,
|
||||
"prompts": true,
|
||||
"router": true,
|
||||
"search_tools": true,
|
||||
"tag": true,
|
||||
"team": true,
|
||||
"user": true,
|
||||
"vector_store": true,
|
||||
}
|
||||
|
||||
func isManagementPath(path string) bool {
|
||||
segments := strings.SplitN(strings.TrimPrefix(path, "/"), "/", 2)
|
||||
return len(segments) > 0 && managementPrefixes[segments[0]]
|
||||
}
|
||||
|
||||
func parseAllowlist(path string) (map[string]bool, error) {
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer file.Close()
|
||||
entries := make(map[string]bool)
|
||||
scanner := bufio.NewScanner(file)
|
||||
line := 0
|
||||
for scanner.Scan() {
|
||||
line++
|
||||
text := strings.TrimSpace(scanner.Text())
|
||||
if text == "" || strings.HasPrefix(text, "#") {
|
||||
continue
|
||||
}
|
||||
if idx := strings.Index(text, "#"); idx >= 0 {
|
||||
text = strings.TrimSpace(text[:idx])
|
||||
}
|
||||
fields := strings.Fields(text)
|
||||
if len(fields) != 2 || !strings.HasPrefix(fields[1], "/") {
|
||||
return nil, fmt.Errorf("%s:%d: allowlist entries must be \"METHOD /path\", got %q", path, line, text)
|
||||
}
|
||||
entries[strings.ToUpper(fields[0])+" "+fields[1]] = true
|
||||
}
|
||||
return entries, scanner.Err()
|
||||
}
|
||||
|
||||
func specCallCovered(calls []endpointCall, specMethod, specPath string) bool {
|
||||
for _, call := range calls {
|
||||
if strings.EqualFold(call.Method, specMethod) && pathMatches(call.Path, specPath) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func auditCoverage(calls []endpointCall, specPaths map[string]map[string]json.RawMessage, allowlist map[string]bool) []string {
|
||||
var violations []string
|
||||
seen := make(map[string]bool)
|
||||
for specPath, operations := range specPaths {
|
||||
if !isManagementPath(specPath) {
|
||||
continue
|
||||
}
|
||||
for method := range operations {
|
||||
entry := strings.ToUpper(method) + " " + specPath
|
||||
covered := specCallCovered(calls, method, specPath)
|
||||
switch {
|
||||
case allowlist[entry]:
|
||||
seen[entry] = true
|
||||
if covered {
|
||||
violations = append(violations, fmt.Sprintf("stale allowlist entry: %s is covered by the provider; remove it from the allowlist", entry))
|
||||
}
|
||||
case !covered:
|
||||
violations = append(violations, fmt.Sprintf("uncovered management endpoint: %s has no provider resource or data source; add coverage or allowlist it with a reason", entry))
|
||||
}
|
||||
}
|
||||
}
|
||||
for entry := range allowlist {
|
||||
if !seen[entry] {
|
||||
violations = append(violations, fmt.Sprintf("stale allowlist entry: %s is not a management endpoint in the proxy schema; remove it from the allowlist", entry))
|
||||
}
|
||||
}
|
||||
sort.Strings(violations)
|
||||
return violations
|
||||
}
|
||||
128
terraform/provider/tools/endpointaudit/coverage_allowlist.txt
Normal file
128
terraform/provider/tools/endpointaudit/coverage_allowlist.txt
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
# Management endpoints deliberately not covered by a Terraform resource or data source.
|
||||
#
|
||||
# Format: one "METHOD /path" per line, matching the proxy OpenAPI schema exactly;
|
||||
# "#" starts a comment. The coverage gate (endpointaudit -coverage-allowlist) fails
|
||||
# when a management endpoint is neither covered nor listed here, and also when an
|
||||
# entry goes stale (the provider now covers it, or the endpoint left the schema),
|
||||
# so this file can only shrink relative to the schema over time.
|
||||
#
|
||||
# Every entry needs a reason. Endpoints that are analytics, UI helpers, or
|
||||
# imperative one-shot operations never get a resource. Entries marked "known gap"
|
||||
# are real coverage gaps awaiting a resource; remove them when the resource lands.
|
||||
|
||||
# Read-only analytics and spend reporting; observability, not Terraform-managed state
|
||||
GET /agent/daily/activity
|
||||
GET /customer/daily/activity
|
||||
GET /guardrails/usage/detail/{guardrail_id}
|
||||
GET /guardrails/usage/logs
|
||||
GET /guardrails/usage/overview
|
||||
GET /key/spend/report
|
||||
GET /organization/daily/activity
|
||||
GET /organization/spend/report
|
||||
GET /tag/daily/activity
|
||||
GET /tag/dau
|
||||
GET /tag/distinct
|
||||
GET /tag/mau
|
||||
GET /tag/summary
|
||||
GET /tag/user-agent/per-user-analytics
|
||||
GET /tag/wau
|
||||
GET /team/daily/activity
|
||||
GET /team/daily/activity/aggregated
|
||||
GET /team/spend/report
|
||||
GET /user/daily/activity
|
||||
GET /user/daily/activity/aggregated
|
||||
GET /user/spend/report
|
||||
|
||||
# Admin UI helper endpoints; serve UI forms and caller-scoped views, not desired state
|
||||
GET /budget/settings
|
||||
GET /router/fields
|
||||
GET /guardrails/ui/add_guardrail_settings
|
||||
GET /guardrails/ui/category_yaml/{category_name}
|
||||
GET /guardrails/ui/major_airlines
|
||||
GET /guardrails/ui/provider_specific_params
|
||||
GET /key/aliases
|
||||
GET /model/deprecations
|
||||
GET /search_tools/ui/available_providers
|
||||
GET /team/available
|
||||
GET /team/metadata_schema
|
||||
GET /team/{team_id}/members/me
|
||||
GET /user/available_users
|
||||
|
||||
# Imperative one-shot operations: bulk edits, rotation, health probes, test hooks,
|
||||
# migrations, and approval workflows; procedural, not declarative state
|
||||
GET /cache/ping
|
||||
GET /cache/redis/info
|
||||
GET /credentials/migrate-encryption/check
|
||||
POST /cache/delete
|
||||
POST /cache/flushall
|
||||
POST /cache/settings/test
|
||||
POST /coordination_redis/settings/test
|
||||
GET /guardrails/submissions
|
||||
GET /guardrails/submissions/{guardrail_id}
|
||||
POST /credentials/migrate-encryption
|
||||
POST /customer/block
|
||||
POST /customer/unblock
|
||||
POST /guardrails/apply_guardrail
|
||||
POST /guardrails/register
|
||||
POST /guardrails/submissions/{guardrail_id}/approve
|
||||
POST /guardrails/submissions/{guardrail_id}/reject
|
||||
POST /guardrails/test_custom_code
|
||||
POST /guardrails/validate_blocked_words_file
|
||||
POST /key/bulk_update
|
||||
POST /key/health
|
||||
POST /key/regenerate
|
||||
POST /key/service-account/generate
|
||||
POST /key/{key}/regenerate
|
||||
POST /key/{key}/reset_spend
|
||||
POST /model/block
|
||||
POST /model/unblock
|
||||
POST /prompts/test
|
||||
POST /search_tools/test_connection
|
||||
POST /team/bulk_member_add
|
||||
POST /team/{team_id}/member/{user_id}/reset_spend
|
||||
POST /team/key/bulk_update
|
||||
POST /team/permissions_bulk_update
|
||||
POST /team/{team_id}/disable_logging
|
||||
POST /user/bulk_update
|
||||
|
||||
# Alternate method or path for functionality the provider already manages elsewhere
|
||||
GET /credentials/by_model/{model_id}
|
||||
GET /guardrails/{guardrail_id}
|
||||
GET /prompts/{prompt_id}
|
||||
GET /prompts/{prompt_id}/versions
|
||||
PATCH /guardrails/{guardrail_id}
|
||||
PATCH /model/{model_id}/update
|
||||
PATCH /prompts/{prompt_id}
|
||||
PATCH /team/{team_id}
|
||||
POST /team/model/add
|
||||
POST /team/model/delete
|
||||
|
||||
# Known gaps awaiting a resource or data source; remove the entry when it lands
|
||||
GET /credentials # known gap: plural credentials data source
|
||||
GET /cache/settings # known gap: cache settings resource
|
||||
POST /cache/settings # known gap: cache settings resource
|
||||
GET /coordination_redis/settings # known gap: coordination redis settings resource
|
||||
POST /coordination_redis/settings # known gap: coordination redis settings resource
|
||||
GET /router/settings # known gap: router settings data source
|
||||
GET /router/fields # known gap: router settings data source
|
||||
GET /config/block_requests_for_models_without_pricing # known gap: proxy config resource
|
||||
PATCH /config/block_requests_for_models_without_pricing # known gap: proxy config resource
|
||||
GET /config/cost_discount_config # known gap: proxy config resource
|
||||
PATCH /config/cost_discount_config # known gap: proxy config resource
|
||||
GET /config/cost_margin_config # known gap: proxy config resource
|
||||
PATCH /config/cost_margin_config # known gap: proxy config resource
|
||||
GET /config/pass_through_endpoint # known gap: pass-through endpoint resource
|
||||
POST /config/pass_through_endpoint # known gap: pass-through endpoint resource
|
||||
DELETE /config/pass_through_endpoint # known gap: pass-through endpoint resource
|
||||
POST /config/pass_through_endpoint/{endpoint_id} # known gap: pass-through endpoint resource
|
||||
GET /config/pass_through_endpoint/team/{team_id} # known gap: pass-through endpoint resource
|
||||
GET /vector_store/list # known gap: plural vector stores data source
|
||||
GET /jwt/key/mapping/list # known gap: plural jwt key mappings data source
|
||||
GET /customer/info # known gap: litellm_customer resource
|
||||
GET /customer/list # known gap: litellm_customer resource
|
||||
POST /customer/new # known gap: litellm_customer resource
|
||||
POST /customer/update # known gap: litellm_customer resource
|
||||
POST /customer/delete # known gap: litellm_customer resource
|
||||
GET /team/{team_id}/callback # known gap: team callback resource
|
||||
POST /team/{team_id}/callback # known gap: team callback resource
|
||||
DELETE /team/{team_id}/callback/{callback_name} # known gap: team callback resource
|
||||
136
terraform/provider/tools/endpointaudit/coverage_test.go
Normal file
136
terraform/provider/tools/endpointaudit/coverage_test.go
Normal file
|
|
@ -0,0 +1,136 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func coverageSpecFixture(paths map[string][]string) map[string]map[string]json.RawMessage {
|
||||
spec := make(map[string]map[string]json.RawMessage)
|
||||
for path, methods := range paths {
|
||||
operations := make(map[string]json.RawMessage)
|
||||
for _, method := range methods {
|
||||
operations[method] = json.RawMessage(`{}`)
|
||||
}
|
||||
spec[path] = operations
|
||||
}
|
||||
return spec
|
||||
}
|
||||
|
||||
func writeAllowlist(t *testing.T, body string) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(t.TempDir(), "allowlist.txt")
|
||||
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
func TestParseAllowlist(t *testing.T) {
|
||||
path := writeAllowlist(t, `# comment
|
||||
GET /team/spend/report
|
||||
|
||||
post /key/regenerate # inline reason
|
||||
`)
|
||||
entries, err := parseAllowlist(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(entries) != 2 || !entries["GET /team/spend/report"] || !entries["POST /key/regenerate"] {
|
||||
t.Fatalf("unexpected entries: %v", entries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseAllowlistRejectsMalformedLines(t *testing.T) {
|
||||
path := writeAllowlist(t, "GET\n")
|
||||
if _, err := parseAllowlist(path); err == nil {
|
||||
t.Fatal("expected error for malformed line")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditCoverageFailsOnUncoveredManagementEndpoint(t *testing.T) {
|
||||
spec := coverageSpecFixture(map[string][]string{
|
||||
"/team/new": {"post"},
|
||||
"/team/spend/report": {"get"},
|
||||
"/chat/completions": {"post"},
|
||||
"/health/liveliness": {"get"},
|
||||
"/v1/chat/completions": {"post"},
|
||||
})
|
||||
calls := []endpointCall{{Method: "POST", Path: "/team/new"}}
|
||||
violations := auditCoverage(calls, spec, nil)
|
||||
if len(violations) != 1 || !strings.Contains(violations[0], "GET /team/spend/report") {
|
||||
t.Fatalf("unexpected violations: %v", violations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditCoverageAllowlistSuppressesUncovered(t *testing.T) {
|
||||
spec := coverageSpecFixture(map[string][]string{"/team/spend/report": {"get"}})
|
||||
violations := auditCoverage(nil, spec, map[string]bool{"GET /team/spend/report": true})
|
||||
if len(violations) != 0 {
|
||||
t.Fatalf("unexpected violations: %v", violations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditCoverageFailsOnStaleCoveredEntry(t *testing.T) {
|
||||
spec := coverageSpecFixture(map[string][]string{"/team/new": {"post"}})
|
||||
calls := []endpointCall{{Method: "POST", Path: "/team/new"}}
|
||||
violations := auditCoverage(calls, spec, map[string]bool{"POST /team/new": true})
|
||||
if len(violations) != 1 || !strings.Contains(violations[0], "stale allowlist entry: POST /team/new is covered") {
|
||||
t.Fatalf("unexpected violations: %v", violations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditCoverageFailsOnEntryMissingFromSchema(t *testing.T) {
|
||||
spec := coverageSpecFixture(map[string][]string{"/team/new": {"post"}})
|
||||
calls := []endpointCall{{Method: "POST", Path: "/team/new"}}
|
||||
violations := auditCoverage(calls, spec, map[string]bool{"POST /team/removed": true})
|
||||
if len(violations) != 1 || !strings.Contains(violations[0], "POST /team/removed is not a management endpoint") {
|
||||
t.Fatalf("unexpected violations: %v", violations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditCoverageMatchesPathParams(t *testing.T) {
|
||||
spec := coverageSpecFixture(map[string][]string{"/team/{team_id}/callback": {"get"}})
|
||||
calls := []endpointCall{{Method: "GET", Path: "/team/{param}/callback"}}
|
||||
violations := auditCoverage(calls, spec, nil)
|
||||
if len(violations) != 0 {
|
||||
t.Fatalf("unexpected violations: %v", violations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountedDeclarativeAPIsAreManagementPaths(t *testing.T) {
|
||||
for _, path := range []string{
|
||||
"/cache/settings",
|
||||
"/config/cost_discount_config",
|
||||
"/coordination_redis/settings",
|
||||
"/router/settings",
|
||||
} {
|
||||
if !isManagementPath(path) {
|
||||
t.Fatalf("%s should be classified as a management path", path)
|
||||
}
|
||||
}
|
||||
for _, path := range []string{"/chat/completions", "/health/liveliness"} {
|
||||
if isManagementPath(path) {
|
||||
t.Fatalf("%s should not be classified as a management path", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBundledAllowlistEntriesAreManagementPaths(t *testing.T) {
|
||||
entries, err := parseAllowlist("coverage_allowlist.txt")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(entries) == 0 {
|
||||
t.Fatal("bundled allowlist parsed to zero entries")
|
||||
}
|
||||
for entry := range entries {
|
||||
fields := strings.Fields(entry)
|
||||
if !isManagementPath(fields[1]) {
|
||||
t.Fatalf("allowlist entry %q is not under a management prefix", entry)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -306,7 +306,7 @@ func auditCalls(calls []endpointCall, specPaths map[string]map[string]json.RawMe
|
|||
return violations
|
||||
}
|
||||
|
||||
func run(providerDir, specPath string) error {
|
||||
func run(providerDir, specPath, coverageAllowlistPath string) error {
|
||||
extracted, err := extractProviderCalls(providerDir)
|
||||
if err != nil {
|
||||
return err
|
||||
|
|
@ -326,6 +326,16 @@ func run(providerDir, specPath string) error {
|
|||
sort.Strings(violations)
|
||||
return fmt.Errorf("provider/proxy endpoint drift:\n %s", strings.Join(violations, "\n "))
|
||||
}
|
||||
if coverageAllowlistPath != "" {
|
||||
allowlist, err := parseAllowlist(coverageAllowlistPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
coverageViolations := auditCoverage(extracted.Calls, specPaths, allowlist)
|
||||
if len(coverageViolations) > 0 {
|
||||
return fmt.Errorf("provider coverage gaps:\n %s", strings.Join(coverageViolations, "\n "))
|
||||
}
|
||||
}
|
||||
fmt.Printf("OK: %d request call sites verified against %d proxy OpenAPI paths\n", len(extracted.Calls), len(specPaths))
|
||||
return nil
|
||||
}
|
||||
|
|
@ -333,12 +343,13 @@ func run(providerDir, specPath string) error {
|
|||
func main() {
|
||||
providerDir := flag.String("provider-dir", "./litellm", "directory containing the provider Go source")
|
||||
specPath := flag.String("spec", "", "path to the proxy OpenAPI schema JSON")
|
||||
coverageAllowlist := flag.String("coverage-allowlist", "", "path to the coverage allowlist; when set, also fail on management endpoints with no provider coverage")
|
||||
flag.Parse()
|
||||
if *specPath == "" {
|
||||
fmt.Fprintln(os.Stderr, "error: -spec is required")
|
||||
os.Exit(2)
|
||||
}
|
||||
if err := run(*providerDir, *specPath); err != nil {
|
||||
if err := run(*providerDir, *specPath, *coverageAllowlist); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "error: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,6 +6,9 @@ import pytest
|
|||
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages import (
|
||||
streaming_iterator as streaming_iterator_module,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE,
|
||||
AnthropicMessagesStreamHiddenParams,
|
||||
|
|
@ -26,7 +29,7 @@ class _RecordingLoggingIterator(BaseAnthropicMessagesStreamingIterator):
|
|||
self.logged_chunks: list = []
|
||||
self.logging_call_count: int = 0
|
||||
|
||||
async def _handle_streaming_logging(self, collected_chunks):
|
||||
async def _handle_streaming_logging(self, collected_chunks, *, stream_teardown=False):
|
||||
self.logged_chunks = list(collected_chunks)
|
||||
self.logging_call_count += 1
|
||||
|
||||
|
|
@ -543,3 +546,86 @@ def test_anthropic_messages_response_as_sse_events_no_content_blocks():
|
|||
response = {"id": "msg_4", "content": [], "stop_reason": "end_turn"}
|
||||
decoded = _decode_sse_events(anthropic_messages_response_as_sse_events(response))
|
||||
assert [event_type for event_type, _ in decoded] == ["message_start", "message_delta", "message_stop"]
|
||||
|
||||
|
||||
class _RecordingLoggingWorker:
|
||||
def __init__(self):
|
||||
self.enqueued = []
|
||||
|
||||
def ensure_initialized_and_enqueue(self, async_coroutine):
|
||||
self.enqueued.append(async_coroutine)
|
||||
|
||||
def close_enqueued(self):
|
||||
for coroutine in self.enqueued:
|
||||
coroutine.close()
|
||||
|
||||
|
||||
async def _noop_deferred_dispatch(logging_coroutine):
|
||||
logging_coroutine.close()
|
||||
|
||||
|
||||
async def _stream_of(events):
|
||||
for event in events:
|
||||
yield event
|
||||
|
||||
|
||||
COMPLETE_STREAM_EVENTS = TRUNCATED_TOOL_USE_EVENTS + ({"type": "message_stop"},)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normal_end_with_deferred_dispatch_armed_parks_logging_coroutine(monkeypatch):
|
||||
"""
|
||||
Regression test for LIT-6409: with post_call guardrails active the proxy
|
||||
arms logging_obj._on_deferred_stream_complete, and the native /v1/messages
|
||||
iterator must park its logging coroutine instead of enqueueing it at
|
||||
upstream exhaustion, otherwise the spend log is built before the
|
||||
guardrail end-of-stream scan writes its post_call entry.
|
||||
"""
|
||||
worker = _RecordingLoggingWorker()
|
||||
monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker)
|
||||
iterator = _make_iterator("test_deferred_parks_logging_coroutine")
|
||||
iterator.litellm_logging_obj._on_deferred_stream_complete = _noop_deferred_dispatch
|
||||
|
||||
await _collect(iterator, _stream_of(COMPLETE_STREAM_EVENTS))
|
||||
|
||||
parked = getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None)
|
||||
assert worker.enqueued == []
|
||||
assert parked is not None
|
||||
assert len(parked) == 1
|
||||
assert asyncio.iscoroutine(parked[0])
|
||||
parked[0].close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_disconnect_enqueues_immediately_even_when_deferred_dispatch_armed(monkeypatch):
|
||||
"""
|
||||
On client disconnect the guardrail end-of-stream scan never runs, so
|
||||
deferral would strand the spend log; the teardown path must keep
|
||||
enqueueing immediately (LIT-5839) even when the deferred callback is armed.
|
||||
"""
|
||||
worker = _RecordingLoggingWorker()
|
||||
monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker)
|
||||
iterator = _make_iterator("test_disconnect_enqueues_when_armed")
|
||||
iterator.litellm_logging_obj._on_deferred_stream_complete = _noop_deferred_dispatch
|
||||
|
||||
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
|
||||
for _ in range(len(TRUNCATED_TOOL_USE_EVENTS)):
|
||||
await wrapped.__anext__()
|
||||
await wrapped.aclose()
|
||||
|
||||
assert len(worker.enqueued) == 1
|
||||
assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None
|
||||
worker.close_enqueued()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normal_end_without_deferred_dispatch_enqueues_immediately(monkeypatch):
|
||||
worker = _RecordingLoggingWorker()
|
||||
monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker)
|
||||
iterator = _make_iterator("test_unarmed_enqueues_at_stream_end")
|
||||
|
||||
await _collect(iterator, _stream_of(COMPLETE_STREAM_EVENTS))
|
||||
|
||||
assert len(worker.enqueued) == 1
|
||||
assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None
|
||||
worker.close_enqueued()
|
||||
|
|
|
|||
|
|
@ -2385,6 +2385,7 @@ async def test_start_shadow_eval_seeds_a_zero_funnel_row_per_leg(monkeypatch: py
|
|||
separates 'nothing was skipped' from a job predating the funnel."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
_configure_anthropic_sdk_judge(monkeypatch)
|
||||
prisma = _shadow_prisma(legs=[])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ Covers ``_wrap_streaming_iterator_with_enrichment``,
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
|
|
@ -18,6 +19,10 @@ from fastapi import HTTPException
|
|||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
BaseAnthropicMessagesStreamingIterator,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
|
||||
|
|
@ -346,6 +351,134 @@ async def test_async_post_call_streaming_iterator_hook_upstream_error_raises(pro
|
|||
pass
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# deferred native /v1/messages stream logging (LIT-6409)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
_NATIVE_MESSAGES_STREAM_EVENTS = (
|
||||
{"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 3, "output_tokens": 1}}},
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}},
|
||||
{"type": "message_stop"},
|
||||
)
|
||||
|
||||
|
||||
def _armed_native_messages_stream(test_name: str, request_data: Dict[str, Any], events: List[Any]):
|
||||
"""The proxy-side setup for a native /v1/messages stream with post_call
|
||||
guardrails active: a real BaseAnthropicMessagesStreamingIterator whose
|
||||
logging_obj carries the deferred-dispatch callback the proxy arms in
|
||||
common_request_processing. The callback records what the guardrail
|
||||
metadata contained at the moment the deferred logging was dispatched."""
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="bedrock/invoke/anthropic.claude-sonnet-4-20250514-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id=test_name,
|
||||
function_id=test_name,
|
||||
)
|
||||
|
||||
async def _dispatch_deferred_logging(logging_coroutine):
|
||||
events.append(
|
||||
(
|
||||
"logging_dispatched",
|
||||
"post_call_entry_visible",
|
||||
bool(request_data.get("metadata", {}).get("standard_logging_guardrail_information")),
|
||||
)
|
||||
)
|
||||
logging_coroutine.close()
|
||||
|
||||
logging_obj._on_deferred_stream_complete = _dispatch_deferred_logging
|
||||
request_data["litellm_logging_obj"] = logging_obj
|
||||
|
||||
iterator = BaseAnthropicMessagesStreamingIterator(litellm_logging_obj=logging_obj, request_body={})
|
||||
|
||||
async def _upstream():
|
||||
for event in _NATIVE_MESSAGES_STREAM_EVENTS:
|
||||
yield event
|
||||
|
||||
return logging_obj, iterator.async_sse_wrapper(_upstream())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_messages_stream_logging_fires_after_guardrail_end_of_stream_scan(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""
|
||||
Regression test for LIT-6409: on native /v1/messages streams the
|
||||
end-of-stream guardrail scan writes its post_call entry AFTER the
|
||||
upstream iterator is exhausted, so success logging dispatched at
|
||||
upstream exhaustion never sees it. The deferred dispatch must fire
|
||||
only after the guardrail chain fully drains.
|
||||
"""
|
||||
events: List[Any] = []
|
||||
request_data: Dict[str, Any] = {"metadata": {}}
|
||||
_, native_stream = _armed_native_messages_stream(
|
||||
"test_native_stream_deferred_ordering", request_data, events
|
||||
)
|
||||
|
||||
class _EndOfStreamScanGuardrail(CustomLogger):
|
||||
async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data):
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
request_data.setdefault("metadata", {})["standard_logging_guardrail_information"] = [
|
||||
{"guardrail_mode": "post_call", "guardrail_status": "success"}
|
||||
]
|
||||
events.append("scan_appended")
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_EndOfStreamScanGuardrail()])
|
||||
|
||||
async for _ in proxy_logging.async_post_call_streaming_iterator_hook(
|
||||
response=native_stream,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert events == ["scan_appended", ("logging_dispatched", "post_call_entry_visible", True)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_messages_stream_logging_fires_when_guardrail_blocks_after_stream_end(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""
|
||||
A guardrail block raised after upstream exhaustion (unified_guardrail
|
||||
re-raises HTTPException for blocked content) must still flush the
|
||||
parked deferred logging, or the blocked stream loses its spend log.
|
||||
"""
|
||||
events: List[Any] = []
|
||||
request_data: Dict[str, Any] = {"metadata": {}}
|
||||
logging_obj, native_stream = _armed_native_messages_stream(
|
||||
"test_native_stream_deferred_block", request_data, events
|
||||
)
|
||||
|
||||
class _BlockingGuardrail(CustomLogger):
|
||||
async def async_post_call_streaming_iterator_hook(self, user_api_key_dict, response, request_data):
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
raise HTTPException(status_code=400, detail={"error": "Violated guardrail policy"})
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_BlockingGuardrail()])
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
async for _ in proxy_logging.async_post_call_streaming_iterator_hook(
|
||||
response=native_stream,
|
||||
user_api_key_dict=make_user_api_key_auth(),
|
||||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert [event[0] for event in events] == ["logging_dispatched"]
|
||||
assert logging_obj._deferred_stream_complete_args is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _fire_deferred_stream_logging
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue