This commit is contained in:
Julian Wachter 2026-09-28 19:23:23 -04:00 • committed by GitHub
commit c01ce89df6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 108 additions and 8 deletions

View file

@ -125,7 +125,6 @@ func (c *Client) UpdateKey(key *Key) (*Key, error) {
// Create a new map with only the fields that can be updated
updateData := map[string]interface{}{
"key": key.Key,
"team_id": key.TeamID,
"key_alias": key.KeyAlias,
"aliases": key.Aliases,
"permissions": key.Permissions,
@ -144,6 +143,12 @@ func (c *Client) UpdateKey(key *Key) (*Key, error) {
updateData["model_tpm_limit"] = key.ModelTPMLimit
}
if key.TeamID != "" {
updateData["team_id"] = key.TeamID
} else if key.clearTeamID {
updateData["team_id"] = nil
}
// The proxy rejects an empty-string budget_duration with a 400, so only
// send it when set.
if key.BudgetDuration != "" {

View file

@ -322,6 +322,7 @@ func resourceKeyUpdate(ctx context.Context, d *schema.ResourceData, m interface{
key := &Key{Key: d.Id()}
mapResourceDataToKey(d, key)
key.clearTeamID = d.HasChange("team_id") && key.TeamID == ""
if !d.HasChange("duration") {
key.Duration = ""
}

View file

@ -321,6 +321,98 @@ func TestUpdateKeyOmitsEmptyBudgetDuration(t *testing.T) {
}
}
func TestUpdateKeyOmitsEmptyTeamID(t *testing.T) {
var captured map[string]interface{}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
json.Unmarshal(body, &captured)
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(`{"key": "sk-test"}`))
}))
defer srv.Close()
client := NewClient(srv.URL, "test-key", true)
if _, err := client.UpdateKey(&Key{Key: "sk-test"}); err != nil {
t.Fatalf("UpdateKey returned error: %v", err)
}
if _, present := captured["team_id"]; present {
t.Errorf("update payload contains empty team_id: %v", captured["team_id"])
}
if _, err := client.UpdateKey(&Key{Key: "sk-test", TeamID: "team-1"}); err != nil {
t.Fatalf("UpdateKey returned error: %v", err)
}
if captured["team_id"] != "team-1" {
t.Errorf("team_id = %v, want team-1", captured["team_id"])
}
}
func TestResourceKeyTeamChangesConverge(t *testing.T) {
for _, tc := range []struct {
name string
priorTeam string
configuredTeam interface{}
wantTeam string
wantPresent bool
wantNull bool
}{
{name: "teamless"},
{name: "remove", priorTeam: "team-1", wantPresent: true, wantNull: true},
{name: "blank", priorTeam: "team-1", configuredTeam: "", wantPresent: true, wantNull: true},
{name: "assign", configuredTeam: "team-1", wantTeam: "team-1", wantPresent: true},
{name: "move", priorTeam: "team-1", configuredTeam: "team-2", wantTeam: "team-2", wantPresent: true},
{name: "retain", priorTeam: "team-1", configuredTeam: "team-1", wantTeam: "team-1", wantPresent: true},
} {
t.Run(tc.name, func(t *testing.T) {
storedTeam := tc.priorTeam
updates := 0
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
w.Header().Set("Content-Type", "application/json")
if req.URL.Path == "/key/update" {
updates++
var payload map[string]interface{}
if err := json.NewDecoder(req.Body).Decode(&payload); err != nil {
t.Error(err)
w.WriteHeader(http.StatusBadRequest)
return
}
value, present := payload["team_id"]
if present != tc.wantPresent || (present && (value == nil) != tc.wantNull) {
t.Errorf("team_id = %#v, present = %v; want present = %v, null = %v", value, present, tc.wantPresent, tc.wantNull)
}
if present {
storedTeam, _ = value.(string)
}
}
json.NewEncoder(w).Encode(map[string]interface{}{
"key": "hash-1",
"info": map[string]interface{}{"team_id": storedTeam, "key_alias": "test"},
})
}))
defer srv.Close()
res := resourceKey()
priorData := newKeyResourceData(t, map[string]interface{}{"team_id": tc.priorTeam, "key_alias": "test", "max_budget": 10.0})
priorData.SetId("hash-1")
config := terraform.NewResourceConfigRaw(map[string]interface{}{"team_id": tc.configuredTeam, "key_alias": "test", "max_budget": 25.0})
diff, err := res.Diff(context.Background(), priorData.State(), config, nil)
if err != nil {
t.Fatal(err)
}
state, diags := res.Apply(context.Background(), priorData.State(), diff, NewClient(srv.URL, "test-key", true))
if diags.HasError() {
t.Fatalf("apply failed: %v", diags)
}
if updates != 1 || storedTeam != tc.wantTeam || state.Attributes["team_id"] != tc.wantTeam {
t.Fatalf("updates = %d, stored team = %q, state team = %q; want %q", updates, storedTeam, state.Attributes["team_id"], tc.wantTeam)
}
nextDiff, err := res.Diff(context.Background(), state, config, nil)
if err != nil || (nextDiff != nil && !nextDiff.Empty()) {
t.Fatalf("subsequent plan not clean: diff = %v, error = %v", nextDiff, err)
}
})
}
}
func TestResourceKeyUpdateFailureKeepsPriorState(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")

View file

@ -129,6 +129,7 @@ type ModelInfo struct {
// Key represents a LiteLLM API key.
type Key struct {
clearTeamID bool
Key string `json:"key,omitempty"`
TokenID string `json:"token_id,omitempty"`
Models []string `json:"models"`

View file

@ -3,12 +3,14 @@ OpenAPI compliance tests for Google Interactions API.
Validates that our SDK requests/responses match the OpenAPI spec at:
https://ai.google.dev/static/api/interactions.openapi.json
Schema names verified against that spec on 2026-09-28.
Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
"""
import json
import os
import re
from typing import Any, Dict
from unittest.mock import MagicMock, patch
@ -60,12 +62,11 @@ class TestRequestCompliance:
"""Tests that our request bodies match the OpenAPI spec."""
def test_create_model_interaction_request_schema(self, spec_dict):
"""Verify CreateModelInteractionParams schema fields."""
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
"""Verify ModelInteraction schema fields."""
schema = spec_dict["components"]["schemas"]["ModelInteraction"]
# Required fields per spec
assert "model" in schema["required"]
assert "input" in schema["required"]
assert "input" in schema["properties"]
# Check our supported optional fields exist in spec
our_optional_fields = [
@ -88,7 +89,7 @@ class TestRequestCompliance:
def test_input_types_match_spec(self, spec_dict):
"""Verify input field supports string, Content, Content[], Turn[]."""
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
schema = spec_dict["components"]["schemas"]["ModelInteraction"]
input_schema = schema["properties"]["input"]
# The input property may be inline oneOf or a $ref to InteractionsInput
@ -313,7 +314,7 @@ class TestEndpointCompliance:
get_path = None
for path, methods in paths.items():
if "{id}" in path and "interactions" in path and "get" in methods:
if re.fullmatch(r".*/interactions/\{[^/{}]+\}", path) and "get" in methods:
get_path = path
break
@ -326,7 +327,7 @@ class TestEndpointCompliance:
delete_path = None
for path, methods in paths.items():
if "{id}" in path and "interactions" in path and "delete" in methods:
if re.fullmatch(r".*/interactions/\{[^/{}]+\}", path) and "delete" in methods:
delete_path = path
break