diff --git a/README.md b/README.md index ab93c2ae2..49e52c195 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,12 @@ Topic pages: - [Purchase Safety](docs/cli/purchase-safety.md) - dry-run, audit log, idempotency window, and guardrails - [Cloud Setup](docs/cli/cloud-setup.md) - `configure-azure` and `configure-gcp` self-hosted credential bootstrap +## MCP Server + +CUDly also ships an MCP server (`cudly-mcp`) that lets Claude and other MCP clients search recommendations and drive RI, Savings Plan, and CUD purchases across AWS, Azure, and GCP, with the same dry-run-by-default safety as the CLI. + +Setup and usage: [mcp/README.md](mcp/README.md) + ## Key Features - **Multi-Cloud Support** - Unified interface for AWS (production), Azure (experimental), and GCP (experimental) diff --git a/cmd/cudly-mcp/main.go b/cmd/cudly-mcp/main.go new file mode 100644 index 000000000..fde741f11 --- /dev/null +++ b/cmd/cudly-mcp/main.go @@ -0,0 +1,38 @@ +// Command cudly-mcp runs the CUDly MCP server on stdio, exposing CUDly's +// RI/SP/CUD search and purchase tools to any MCP client (e.g. Claude Code +// via ~/.claude/mcp.json). See mcp/README.md for setup and usage. +// +// This binary is intentionally separate from the ri-helper CLI (cmd/main.go): +// it is a local/desktop MCP server, never a Lambda handler, and is kept out +// of iac/ and terraform/ (see mcp/README.md "Deployment model"). +package main + +import ( + "context" + "log" + "os" + + gosdk "github.com/modelcontextprotocol/go-sdk/mcp" + + cudlymcp "github.com/LeanerCloud/CUDly/mcp" + _ "github.com/LeanerCloud/CUDly/providers/aws" + _ "github.com/LeanerCloud/CUDly/providers/azure" + _ "github.com/LeanerCloud/CUDly/providers/gcp" +) + +// version is overridable at build time via: +// +// go build -ldflags "-X main.version=1.2.3" ./cmd/cudly-mcp +var version = "dev" + +func main() { + server, err := cudlymcp.NewServer(version) + if err != nil { + log.Fatalf("cudly-mcp: failed to build server: %v", err) + } + + if err := server.Run(context.Background(), &gosdk.StdioTransport{}); err != nil { + log.Printf("cudly-mcp: server exited with error: %v", err) + os.Exit(1) + } +} diff --git a/cmd/cudly-mcp/main_test.go b/cmd/cudly-mcp/main_test.go new file mode 100644 index 000000000..1ec678940 --- /dev/null +++ b/cmd/cudly-mcp/main_test.go @@ -0,0 +1,103 @@ +package main + +import ( + "context" + "path/filepath" + "strings" + "testing" + "time" + + gosdk "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + cudlymcp "github.com/LeanerCloud/CUDly/mcp" +) + +// isolateFromAmbientAWS points the AWS SDK at deliberately nonexistent +// profile/config/credentials so config.LoadDefaultConfig cannot resolve any +// real credentials -- neither from a dev machine's ~/.aws files nor from the +// network (IMDS, ECS/EKS container credential endpoints, web identity). This +// is required for TestRealPurchasePastProviderRegistration below: that test +// drives a real (non-dry-run) purchase call, so it must be impossible for it +// to reach an actual AWS account or make an actual network call, in this or +// any other environment the test happens to run in. +func isolateFromAmbientAWS(t *testing.T) { + t.Helper() + t.Setenv("AWS_PROFILE", "cudly-mcp-regression-test-nonexistent-profile") + t.Setenv("AWS_SHARED_CREDENTIALS_FILE", filepath.Join(t.TempDir(), "no-credentials")) + t.Setenv("AWS_CONFIG_FILE", filepath.Join(t.TempDir(), "no-config")) + t.Setenv("AWS_ACCESS_KEY_ID", "") + t.Setenv("AWS_SECRET_ACCESS_KEY", "") + t.Setenv("AWS_SESSION_TOKEN", "") + t.Setenv("AWS_EC2_METADATA_DISABLED", "true") + t.Setenv("AWS_CONTAINER_CREDENTIALS_RELATIVE_URI", "") + t.Setenv("AWS_CONTAINER_CREDENTIALS_FULL_URI", "") + t.Setenv("AWS_ROLE_ARN", "") + t.Setenv("AWS_WEB_IDENTITY_TOKEN_FILE", "") +} + +// TestRealPurchasePastProviderRegistration is the regression guard for the +// bug this file's blank imports fix: cudly-mcp never imported +// providers/aws|azure|gcp, so their init()-registered factories were never +// added to provider.CreateProvider's registry, and every real (non-dry-run) +// purchase failed at ResolveClient with "provider aws is not registered" +// before ever reaching AWS. +// +// This test MUST live in package main under cmd/cudly-mcp/ -- go test +// ./mcp/... does not catch this bug even with the blank imports reverted, +// because a test binary for the mcp or mcp/tools package never pulls in +// cmd/cudly-mcp's imports. Only a test in this package has the blank +// imports in its own dependency graph, so only here does reverting them +// actually flip provider.CreateProvider("aws") back to unregistered. +// +// The test asserts the purchase attempt gets PAST registration and fails for +// a completely different, credentials-shaped reason ("AWS is not +// configured", from providers/aws/provider.go's GetServiceClient) instead of +// "not registered". It never reaches AWS: isolateFromAmbientAWS makes +// config.LoadDefaultConfig fail to resolve the (deliberately nonexistent) +// named profile before any credential lookup or network call happens. +func TestRealPurchasePastProviderRegistration(t *testing.T) { + isolateFromAmbientAWS(t) + + server, err := cudlymcp.NewServer("test-regression") + require.NoError(t, err) + + clientTransport, serverTransport := gosdk.NewInMemoryTransports() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + go func() { + _ = server.Run(ctx, serverTransport) + }() + + client := gosdk.NewClient(&gosdk.Implementation{Name: "test-client"}, nil) + session, err := client.Connect(ctx, clientTransport, nil) + require.NoError(t, err) + defer session.Close() + + result, err := session.CallTool(ctx, &gosdk.CallToolParams{ + Name: "cudly_aws_ec2_ri_purchase", + Arguments: map[string]any{ + "region": "us-east-1", + "instance_type": "m5.large", + "count": 1, + "term_years": 1, + "payment_option": "no-upfront", + "dry_run": false, + "confirm": true, + }, + }) + require.NoError(t, err, "CallTool itself must not return a transport-level error") + require.True(t, result.IsError, "a failed real purchase must surface as a tool error, not a transport error") + + text := result.Content[0].(*gosdk.TextContent).Text + assert.NotContains(t, strings.ToLower(text), "not registered", + "provider must be registered for the cudly-mcp binary: got %q", text) + // providers/aws/provider.go's GetServiceClient (via AWSProvider.IsConfigured) + // returns exactly this string when config.LoadDefaultConfig cannot resolve + // the requested profile -- observed and confirmed stable in this test run. + assert.Contains(t, text, "AWS is not configured", + "expected a credentials/config-shaped failure once past registration: got %q", text) +} diff --git a/go.mod b/go.mod index 4e151560d..5dbda56fd 100644 --- a/go.mod +++ b/go.mod @@ -104,9 +104,11 @@ require ( github.com/coreos/go-oidc/v3 v3.18.0 github.com/go-jose/go-jose/v4 v4.1.4 github.com/golang-migrate/migrate/v4 v4.19.1 + github.com/google/jsonschema-go v0.4.3 github.com/google/uuid v1.6.0 github.com/jackc/pgx/v5 v5.9.2 github.com/microsoftgraph/msgraph-sdk-go v1.99.0 + github.com/modelcontextprotocol/go-sdk v1.6.1 github.com/pashagolub/pgxmock/v4 v4.9.0 github.com/testcontainers/testcontainers-go v0.42.0 github.com/testcontainers/testcontainers-go/modules/postgres v0.42.0 @@ -175,12 +177,15 @@ require ( github.com/opencontainers/image-spec v1.1.1 // indirect github.com/planetscale/vtprotobuf v0.6.1-0.20240319094008-0393e58bdf10 // indirect github.com/power-devops/perfstat v0.0.0-20240221224432-82ca36839d55 // indirect + github.com/segmentio/asm v1.1.3 // indirect + github.com/segmentio/encoding v0.5.4 // indirect github.com/shirou/gopsutil/v4 v4.26.3 // indirect github.com/sirupsen/logrus v1.9.4 // indirect github.com/spiffe/go-spiffe/v2 v2.6.0 // indirect github.com/std-uritemplate/std-uritemplate/go/v2 v2.0.3 // indirect github.com/tklauser/go-sysconf v0.3.16 // indirect github.com/tklauser/numcpus v0.11.0 // indirect + github.com/yosida95/uritemplate/v3 v3.0.2 // indirect github.com/yusufpapurcu/wmi v1.2.4 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/contrib/detectors/gcp v1.39.0 // indirect diff --git a/go.sum b/go.sum index c412d33fa..32659d76d 100644 --- a/go.sum +++ b/go.sum @@ -222,6 +222,8 @@ github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6 github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0= +github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE= github.com/google/martian/v3 v3.3.3 h1:DIhPTQrbPkgs2yJYdXU/eNACCG5DVQjySNRNlflZ9Fc= github.com/google/martian/v3 v3.3.3/go.mod h1:iEPrYcgCF7jA9OtScMFQyAlZZ4YXTKEtJ1E6RWzmBA0= github.com/google/s2a-go v0.1.9 h1:LGD7gtMgezd8a/Xak7mEWL0PjoTQFvpRudN895yqKW0= @@ -296,6 +298,8 @@ github.com/moby/sys/userns v0.1.0 h1:tVLXkFOxVu9A64/yh59slHVv9ahO9UIev4JZusOLG/g github.com/moby/sys/userns v0.1.0/go.mod h1:IHUYgu/kao6N8YZlp9Cf444ySSvCmDlmzUcYfDHOl28= github.com/moby/term v0.5.2 h1:6qk3FJAFDs6i/q3W/pQ97SX192qKfZgGjCQqfCJkgzQ= github.com/moby/term v0.5.2/go.mod h1:d3djjFCrjnB+fl8NJux+EJzu0msscUP+f8it8hPkFLc= +github.com/modelcontextprotocol/go-sdk v1.6.1 h1:0zOSupjKUxPKSocPT1Wtago+mUHU2/uZ4xSOY0FGReU= +github.com/modelcontextprotocol/go-sdk v1.6.1/go.mod h1:kzm3kzFL1/+AziGOE0nUs3gvPoNxMCvkxokMkuFapXQ= github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A= github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc= github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= @@ -318,6 +322,10 @@ github.com/power-devops/perfstat v0.0.0-20240221224432-82ca36839d55/go.mod h1:Om github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= +github.com/segmentio/asm v1.1.3 h1:WM03sfUOENvvKexOLp+pCqgb/WDjsi7EK8gIsICtzhc= +github.com/segmentio/asm v1.1.3/go.mod h1:Ld3L4ZXGNcSLRg4JBsZ3//1+f/TjYl0Mzen/DQy1EJg= +github.com/segmentio/encoding v0.5.4 h1:OW1VRern8Nw6ITAtwSZ7Idrl3MXCFwXHPgqESYfvNt0= +github.com/segmentio/encoding v0.5.4/go.mod h1:HS1ZKa3kSN32ZHVZ7ZLPLXWvOVIiZtyJnO1gPH1sKt0= github.com/shirou/gopsutil/v4 v4.26.3 h1:2ESdQt90yU3oXF/CdOlRCJxrP+Am1aBYubTMTfxJ1qc= github.com/shirou/gopsutil/v4 v4.26.3/go.mod h1:LZ6ewCSkBqUpvSOf+LsTGnRinC6iaNUNMGBtDkJBaLQ= github.com/sirupsen/logrus v1.9.4 h1:TsZE7l11zFCLZnZ+teH4Umoq5BhEIfIzfRDZ1Uzql2w= @@ -345,6 +353,8 @@ github.com/tklauser/go-sysconf v0.3.16 h1:frioLaCQSsF5Cy1jgRBrzr6t502KIIwQ0MArYI github.com/tklauser/go-sysconf v0.3.16/go.mod h1:/qNL9xxDhc7tx3HSRsLWNnuzbVfh3e7gh/BmM179nYI= github.com/tklauser/numcpus v0.11.0 h1:nSTwhKH5e1dMNsCdVBukSZrURJRoHbSEQjdEbY+9RXw= github.com/tklauser/numcpus v0.11.0/go.mod h1:z+LwcLq54uWZTX0u/bGobaV34u6V7KNlTZejzM6/3MQ= +github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= +github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4= github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0= github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0= go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= @@ -387,6 +397,8 @@ golang.org/x/text v0.39.0 h1:UbZz4pLOvn600D6Oh6GGEI6VAmndrEBLv8/6BEXzyus= golang.org/x/text v0.39.0/go.mod h1:3UwRclnC2g0TU9x8PZiyfOajCd1zaUNHF9cvqcQZ+ZM= golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= +golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= +golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= diff --git a/go.work.sum b/go.work.sum index 9890d079d..77b0e5af9 100644 --- a/go.work.sum +++ b/go.work.sum @@ -489,6 +489,7 @@ golang.org/x/mod v0.4.1/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c= golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY= golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= +golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= @@ -697,6 +698,7 @@ golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxb golang.org/x/tools v0.26.0/go.mod h1:TPVVj70c7JJ3WCazhD8OdXcZg/og+b9+tH/KxylGwH0= golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI= golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= +golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= golang.org/x/tools/godoc v0.1.0-deprecated/go.mod h1:qM63CriJ961IHWmnWa9CjZnBndniPt4a3CK0PVB9bIg= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= diff --git a/mcp/README.md b/mcp/README.md new file mode 100644 index 000000000..78342949f --- /dev/null +++ b/mcp/README.md @@ -0,0 +1,121 @@ +# CUDly MCP Server + +`cudly-mcp` exposes CUDly's reserved-capacity search and purchase tools (AWS EC2/RDS/ElastiCache/OpenSearch/Redshift/MemoryDB/Savings Plans, Azure VM Reservations, GCP Compute Engine CUDs) to any MCP client -- Claude Code, Claude Desktop, or another MCP-speaking agent -- as a local process. It is a thin wrapper around the same in-process Go packages (`pkg/provider`, `pkg/common`) the `ri-helper` CLI uses; it never shells out to `ri-helper`, and it is not deployed anywhere (see [Deployment model](#deployment-model)). + +Every purchase tool is dry-run by default (`dry_run=true`) and requires an explicit `confirm=true` alongside `dry_run=false` before it spends money. See [Safety model](#safety-model). + +## Install + +From the repository root: + +```bash +go install ./cmd/cudly-mcp +``` + +This installs the `cudly-mcp` binary to `$(go env GOBIN)` (or `$(go env GOPATH)/bin` if `GOBIN` is unset); make sure that directory is on your `PATH` so `cudly-mcp` resolves without a full path. + +Or run directly without a separate install step: + +```bash +go run ./cmd/cudly-mcp +``` + +There is no separate module or release artifact for `cudly-mcp` yet -- install it from a checkout of this repository. + +## Configure credentials + +The server takes no CUDly-specific configuration of its own. Each provider tool authenticates the same way the corresponding CUDly CLI path does, and every tool call also accepts a per-call override (`aws_profile`, `azure_subscription_id`, `gcp_project_id`) so one running server instance can serve requests against different accounts/subscriptions/projects without a restart. + +### AWS + +One of, in the usual SDK precedence order: + +- `AWS_PROFILE` (matches a profile in `~/.aws/config` / `~/.aws/credentials`) +- `AWS_ACCESS_KEY_ID` + `AWS_SECRET_ACCESS_KEY` (+ optional `AWS_SESSION_TOKEN`) +- An IAM role (EC2 instance profile, ECS task role, etc.) + +Per-call override: pass `aws_profile` on any AWS tool call to use a specific named profile for that call only. + +### Azure + +One of: + +- A service principal via `AZURE_CLIENT_ID`, `AZURE_CLIENT_SECRET`, `AZURE_TENANT_ID` +- Local `az login` state (picked up by `azidentity.NewDefaultAzureCredential`) + +Set `AZURE_SUBSCRIPTION_ID` to select the subscription, or pass `azure_subscription_id` on a per-call basis to override it. + +### GCP + +One of: + +- `GOOGLE_APPLICATION_CREDENTIALS` pointing at a service-account JSON key file +- Application Default Credentials (`gcloud auth application-default login`) + +Pass `gcp_project_id` on a per-call basis to override the ambient project. + +## Launch + +```bash +cudly-mcp +``` + +The server speaks MCP over stdio and logs diagnostics to stderr; it does not print anything to stdout other than protocol traffic, so it is safe to launch directly from an MCP client's process-spawning config (below) rather than through a wrapper script. + +## Register with an MCP client + +Add an entry to your client's MCP server config. For Claude Code, this is `~/.claude/mcp.json`. If the client does not inherit your shell's `PATH`, use the absolute path `go install` reported: `$(go env GOBIN)/cudly-mcp` if `GOBIN` is set, otherwise `$(go env GOPATH)/bin/cudly-mcp` (`$(go env GOBIN)` expands to an empty string when `GOBIN` is unset, so that path alone is not a valid binary location): + +```json +{ + "mcpServers": { + "cudly": { + "command": "/absolute/path/to/cudly-mcp", + "env": { + "AWS_PROFILE": "my-aws-profile", + "AZURE_SUBSCRIPTION_ID": "00000000-0000-0000-0000-000000000000" + } + } + } +} +``` + +`env` is optional -- omit it entirely to rely on whatever ambient credentials are already active in the shell that launches the client, or set only the provider(s) you actually use. + +## Worked example + +A typical session searches for a recommendation, previews the purchase, then executes it: + +1. **Search**: call `cudly_search_recommendations` with `provider="aws"`, `service="ec2"`, `region="us-east-1"` to see what AWS Cost Explorer currently recommends reserving. +2. **Preview**: take a result's `region`/`resource_type`/`count` and call `cudly_aws_ec2_ri_purchase` with those values and `term_years`/`payment_option` of your choice. Leave `dry_run` at its default (`true`) -- the response validates your parameters without contacting AWS or spending anything, and reports `cost`/`on_demand_cost`/`estimated_savings`/`savings_percentage` only when a real figure is actually known (omitted otherwise, never a fabricated `0`). +3. **Execute**: once the preview looks right, call the same tool again with `dry_run=false, confirm=true`. This is the only combination that performs a real purchase; any other combination either previews or returns an explicit refusal error (see [Safety model](#safety-model)). + +Every other provider's purchase tool (`cudly_aws_savingsplans_purchase`, `cudly_aws_rds_ri_purchase`, `cudly_azure_compute_ri_purchase`, `cudly_gcp_computeengine_cud_purchase`, ...) follows the identical dry_run-then-confirm pattern. Call `cudly_list_commitment_actions` at any point for the full, always-current list of tools, which ones can spend real money today, and 2-3 example prompts per tool. + +## Safety model + +- `dry_run` defaults to `true` on every purchase tool. A dry-run call never contacts the cloud provider and never spends money -- it only validates your parameters. It reports pricing (`cost`/`on_demand_cost`/`estimated_savings`/`savings_percentage`) only when a real figure is genuinely known; those fields are omitted, not zeroed, when it isn't. +- A real purchase requires **both** `dry_run=false` **and** `confirm=true`. `dry_run=false` with `confirm=false` (or vice versa) is refused with a structured error, not silently downgraded to a preview or silently ignored. +- Every money-affecting parameter (region, resource type, count, term, payment option, and any provider-specific dimension such as RDS's `az_config`) is validated against an explicit enum or non-empty check before anything is built or sent. There is no silent default for a value that materially changes what gets purchased. +- Every real purchase is tagged with a source identifying it came from this MCP server (never a user-suppliable string) and a deterministic idempotency token derived from the request's own parameters. By default, retrying an identical tool call -- however long after the original, and regardless of any clock boundary -- always derives the same token, so the provider dedupes the retry instead of buying twice; this is a fail-safe default, since the worst case of a false dedupe is a skipped intentional repeat, never a double purchase. To deliberately make a second, otherwise-identical purchase (e.g. "buy 3 RIs now" and "buy 3 more next week"), pass a fresh `idempotency_nonce` value on the second call; passing the same nonce on a retry of that same call still dedupes correctly. +- Provider/SDK failures surface their full error text back to the caller; nothing is swallowed. + +## Caveats and known gaps + +These are pre-existing behaviours in the underlying purchase clients, not something introduced by or specific to the MCP server -- flagged here so you know what to expect: + +- **Azure VM Reservations have no partial-upfront billing plan.** Azure honors exactly two billing plans, all-upfront and no-upfront (billed monthly, same total price -- Azure charges no premium for spreading payments), and `cudly_azure_compute_ri_purchase` defaults `payment_option` to no-upfront when omitted. `payment_option=partial-upfront` has no Azure equivalent and is rejected with an explicit error rather than silently purchased under all-upfront or no-upfront instead. +- **GCP Compute Engine CUDs commit resources, not instances.** `cudly_gcp_computeengine_cud_purchase` takes `vcpu_count` and `memory_gb` directly (a CUD is a vCPU+memory commitment), not an instance count -- there is no implicit vCPU-per-instance conversion. +- **AWS Savings Plans searches require term/payment/lookback; `cudly_search_recommendations` defaults them.** Unlike EC2/RDS reservation searches (which AWS defaults server-side when `term_years`/`payment_option`/`lookback_period` are omitted), `GetSavingsPlansPurchaseRecommendation` rejects the call unless all three are set. When `service` targets a Savings Plans search on AWS, the tool fills in unset fields with `term_years=1` (1yr), `payment_option=no-upfront`, `lookback_period=30d` -- an explicitly supplied value is never overridden. Reservation searches (EC2/RDS/etc) are unaffected: omitted still means "search all". +- **`region` on `cudly_search_recommendations` is a client-side post-filter for AWS reservation/Savings Plans searches.** `GetReservationPurchaseRecommendation` and `GetSavingsPlansPurchaseRecommendation` are account-level Cost Explorer calls with no region parameter -- AWS returns recommendations from every region the account has usage in. The tool filters the returned recommendations down to `region`/`include_regions` (minus `exclude_regions`) after the call; when no region constraint is supplied, all regions are returned as before. + +## Deployment model + +`cudly-mcp` is a local/desktop process, not a deployed service: it is intentionally kept out of `iac/`, `terraform/`, and the `internal/api` Lambda packaging path. Run it on the same machine as your MCP client. + +## Troubleshooting + +- **"provider ... is not configured" / credential errors**: confirm the relevant environment variable(s) from [Configure credentials](#configure-credentials) are set in the shell (or the client's `env` block) that launches `cudly-mcp`, or pass the matching per-call override (`aws_profile` / `azure_subscription_id` / `gcp_project_id`). +- **Azure/GCP purchase calls appear to hang**: Azure Reservations and GCP Compute Commitments both provision asynchronously after the purchase call returns; a `success=true` response means the purchase request was accepted, not necessarily that the resource is already active in the portal/console. Re-run `cudly_search_recommendations` or check the provider console if you need to confirm activation state. +- **Rate limits / throttling from the cloud provider**: retry the same tool call with the same parameters -- the idempotency token guarantees a retry cannot double-purchase no matter how long you wait before retrying (see [Safety model](#safety-model)). If you genuinely want a second, separate purchase with the same parameters instead of a retry, pass a fresh `idempotency_nonce` value. +- **"invalid ... must be one of ..." errors**: every enum-typed parameter (term, payment option, engine, az_config, sp_type, scope, tenancy, platform) is validated against an explicit allow-list; call `cudly_list_commitment_actions` or re-check this README's per-tool schema for the exact accepted values. diff --git a/mcp/server.go b/mcp/server.go new file mode 100644 index 000000000..65142099c --- /dev/null +++ b/mcp/server.go @@ -0,0 +1,68 @@ +// Package mcp wires the CUDly MCP server: it registers every tool from +// mcp/tools onto a github.com/modelcontextprotocol/go-sdk/mcp.Server and +// builds the cudly_list_commitment_actions catalog from those same tools' +// descriptors, so the live tool set and the discoverability catalog can +// never drift apart. +package mcp + +import ( + "fmt" + + gosdk "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/mcp/tools" +) + +// ServerName is the MCP Implementation.Name this server identifies as. +const ServerName = "cudly-mcp" + +// registrations returns every purchase/search tool this server exposes, +// excluding cudly_list_commitment_actions itself (NewServer adds that one +// last, once it has every other tool's Descriptor to build the catalog +// from). +func registrations() []tools.Registration { + return []tools.Registration{ + tools.NewSearchRecommendationsTool(), + tools.NewAWSEC2RIPurchaseTool(), + tools.NewAWSOpenSearchRIPurchaseTool(), + tools.NewAWSRedshiftRIPurchaseTool(), + tools.NewAWSMemoryDBRIPurchaseTool(), + tools.NewAWSRDSRIPurchaseTool(), + tools.NewAWSElastiCacheRIPurchaseTool(), + tools.NewAWSSavingsPlansPurchaseTool(), + tools.NewAzureComputeRIPurchaseTool(), + tools.NewGCPComputeEngineCUDPurchaseTool(), + } +} + +// NewServer builds the CUDly MCP server with every tool registered. version +// is reported to clients as the server's Implementation.Version (pass the +// build-time version string, or "dev" for unreleased builds). Callers run +// the returned server on a transport, e.g.: +// +// server, err := mcp.NewServer("1.0.0") +// server.Run(ctx, &gosdk.StdioTransport{}) +func NewServer(version string) (*gosdk.Server, error) { + s := gosdk.NewServer(&gosdk.Implementation{Name: ServerName, Version: version}, nil) + + regs := registrations() + descriptors := make([]tools.Descriptor, 0, len(regs)+1) + for _, r := range regs { + descriptors = append(descriptors, r.Descriptor()) + } + // Include cudly_list_commitment_actions' own entry so the catalog it + // returns really does list "every tool available on this MCP server" as + // its description promises, itself included. + descriptors = append(descriptors, tools.ListCommitmentActionsDescriptor()) + + listTool := tools.NewListCommitmentActions(descriptors) + regs = append(regs, listTool) + + for _, r := range regs { + if err := r.Register(s); err != nil { + return nil, fmt.Errorf("register tool %q: %w", r.Descriptor().Name, err) + } + } + + return s, nil +} diff --git a/mcp/server_test.go b/mcp/server_test.go new file mode 100644 index 000000000..f2b44964c --- /dev/null +++ b/mcp/server_test.go @@ -0,0 +1,223 @@ +package mcp + +import ( + "context" + "encoding/json" + "strings" + "testing" + + gosdk "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/mcp/tools" +) + +func TestNewServerBuildsWithoutError(t *testing.T) { + t.Parallel() + s, err := NewServer("test") + require.NoError(t, err) + require.NotNil(t, s) +} + +// TestRegistryNonEmpty proves the tool registry is never accidentally empty: +// cudly_list_commitment_actions is always present, even before any +// purchase/search tool has been registered. +func TestRegistryNonEmpty(t *testing.T) { + t.Parallel() + regs := registrations() + descriptors := make([]tools.Descriptor, 0, len(regs)+1) + for _, r := range regs { + descriptors = append(descriptors, r.Descriptor()) + } + listTool := tools.NewListCommitmentActions(descriptors) + descriptors = append(descriptors, listTool.Descriptor()) + + require.NotEmpty(t, descriptors) + names := make(map[string]bool, len(descriptors)) + for _, d := range descriptors { + names[d.Name] = true + } + assert.True(t, names["cudly_list_commitment_actions"]) +} + +// TestRealPurchaseToolsDocumentMoneyImpactAndDryRun proves every tool that +// can execute a real purchase leads its description with a money-impact +// statement and a dry-run recommendation (design doc §3/§5) -- so any future +// purchase tool that forgets one fails this test rather than shipping a +// tool description that quietly omits the safety framing every other +// purchase tool carries. +// TestEndToEndSearchThenDryRunPurchase drives the real MCP protocol path +// end to end -- a real Client connected to the real NewServer over an +// in-memory transport, not a bare Go function call -- through the exact +// chain the README's worked example describes: list the catalog, then call +// cudly_aws_ec2_ri_purchase with dry_run=true. It proves the tool schema +// registered without error (AddTool's schema inference/validation runs at +// connect time) and that a dry-run purchase returns structured cost JSON +// without any AWS credentials configured in this test environment -- +// confirming dry_run=true never reaches the real purchase path even through +// the full protocol stack, not just the Go-level unit tests in +// mcp/tools/aws_ec2_ri_test.go. +func TestEndToEndSearchThenDryRunPurchase(t *testing.T) { + t.Parallel() + ctx := context.Background() + + server, err := NewServer("test") + require.NoError(t, err) + + clientTransport, serverTransport := gosdk.NewInMemoryTransports() + go func() { + _ = server.Run(ctx, serverTransport) + }() + + client := gosdk.NewClient(&gosdk.Implementation{Name: "test-client"}, nil) + session, err := client.Connect(ctx, clientTransport, nil) + require.NoError(t, err) + defer session.Close() + + toolsList, err := session.ListTools(ctx, nil) + require.NoError(t, err) + names := make(map[string]bool, len(toolsList.Tools)) + for _, tl := range toolsList.Tools { + names[tl.Name] = true + } + assert.True(t, names["cudly_list_commitment_actions"]) + assert.True(t, names["cudly_aws_ec2_ri_purchase"]) + + result, err := session.CallTool(ctx, &gosdk.CallToolParams{ + Name: "cudly_aws_ec2_ri_purchase", + Arguments: map[string]any{ + "region": "us-east-1", + "instance_type": "m5.large", + "count": 3, + "term_years": 3, + "payment_option": "no-upfront", + // dry_run/confirm omitted: must default to true/false. + }, + }) + require.NoError(t, err) + require.False(t, result.IsError, "dry_run purchase must not be a tool error") + + structured, err := json.Marshal(result.StructuredContent) + require.NoError(t, err) + var resp tools.PurchaseResponse + require.NoError(t, json.Unmarshal(structured, &resp)) + assert.True(t, resp.DryRun) + assert.True(t, resp.Success) + assert.Empty(t, resp.Error) +} + +// TestListCommitmentActionsIncludesItself is the regression guard for the +// CodeRabbit finding that NewServer built the descriptors slice passed to +// cudly_list_commitment_actions before appending the list tool itself to +// regs, so the catalog the tool actually returns at runtime never included +// its own entry despite its description claiming to list "every" tool. +// Unlike TestRegistryNonEmpty (which builds its own descriptors slice by +// hand) this drives the real NewServer wiring and calls the live tool over +// the protocol, so it fails if NewServer regresses even if the manual +// helper above stays correct. +func TestListCommitmentActionsIncludesItself(t *testing.T) { + t.Parallel() + ctx := context.Background() + + server, err := NewServer("test") + require.NoError(t, err) + + clientTransport, serverTransport := gosdk.NewInMemoryTransports() + go func() { + _ = server.Run(ctx, serverTransport) + }() + + client := gosdk.NewClient(&gosdk.Implementation{Name: "test-client"}, nil) + session, err := client.Connect(ctx, clientTransport, nil) + require.NoError(t, err) + defer session.Close() + + result, err := session.CallTool(ctx, &gosdk.CallToolParams{ + Name: "cudly_list_commitment_actions", + Arguments: map[string]any{}, + }) + require.NoError(t, err) + require.False(t, result.IsError, "cudly_list_commitment_actions must not itself error") + + structured, err := json.Marshal(result.StructuredContent) + require.NoError(t, err) + var catalog struct { + Actions []tools.ActionEntry `json:"actions"` + } + require.NoError(t, json.Unmarshal(structured, &catalog)) + + names := make([]string, 0, len(catalog.Actions)) + for _, a := range catalog.Actions { + names = append(names, a.Name) + } + assert.Contains(t, names, "cudly_list_commitment_actions", + "the catalog cudly_list_commitment_actions returns must include its own entry, got: %v", names) +} + +func TestRealPurchaseToolsDocumentMoneyImpactAndDryRun(t *testing.T) { + t.Parallel() + for _, r := range registrations() { + d := r.Descriptor() + if !d.RealPurchaseEnabled { + continue + } + lower := strings.ToLower(d.Description) + assert.Truef(t, strings.Contains(lower, "real money") || strings.Contains(lower, "spends money"), + "tool %q description must state its money impact: %q", d.Name, d.Description) + assert.Truef(t, strings.Contains(lower, "dry_run") || strings.Contains(lower, "dry run"), + "tool %q description must recommend dry_run first: %q", d.Name, d.Description) + } +} + +// TestAzureComputeRIPurchaseSchemaExcludesPartialUpfront proves the +// cudly_azure_compute_ri_purchase tool's advertised payment_option enum +// never offers partial-upfront: Azure has no partial-upfront billing plan +// (mcp/tools/azure_compute_ri.go rejects it at runtime as defense in depth), +// so the schema itself must not invite a caller to pick a value the tool +// can only ever refuse. Drives the real MCP protocol (ListTools), not a +// bare Go function call, so it catches a regression in what a real client +// actually sees, not just in the tool's internal validation. +func TestAzureComputeRIPurchaseSchemaExcludesPartialUpfront(t *testing.T) { + t.Parallel() + ctx := context.Background() + + server, err := NewServer("test") + require.NoError(t, err) + + clientTransport, serverTransport := gosdk.NewInMemoryTransports() + go func() { + _ = server.Run(ctx, serverTransport) + }() + + client := gosdk.NewClient(&gosdk.Implementation{Name: "test-client"}, nil) + session, err := client.Connect(ctx, clientTransport, nil) + require.NoError(t, err) + defer session.Close() + + toolsList, err := session.ListTools(ctx, nil) + require.NoError(t, err) + + var azureTool *gosdk.Tool + for _, tl := range toolsList.Tools { + if tl.Name == "cudly_azure_compute_ri_purchase" { + azureTool = tl + break + } + } + require.NotNil(t, azureTool, "cudly_azure_compute_ri_purchase must be registered") + + schema, ok := azureTool.InputSchema.(map[string]any) + require.True(t, ok, "client-side InputSchema must be the default JSON marshaling (map[string]any)") + properties, ok := schema["properties"].(map[string]any) + require.True(t, ok) + paymentOption, ok := properties["payment_option"].(map[string]any) + require.True(t, ok) + enum, ok := paymentOption["enum"].([]any) + require.True(t, ok) + + assert.NotContains(t, enum, "partial-upfront", + "Azure schema must not advertise partial-upfront: Azure has no such billing plan") + assert.Contains(t, enum, "all-upfront") + assert.Contains(t, enum, "no-upfront") +} diff --git a/mcp/tools/aws_ec2_ri.go b/mcp/tools/aws_ec2_ri.go new file mode 100644 index 000000000..dce1e6229 --- /dev/null +++ b/mcp/tools/aws_ec2_ri.go @@ -0,0 +1,224 @@ +package tools + +import ( + "context" + "fmt" + + ec2types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +const awsEC2RIPurchaseName = "cudly_aws_ec2_ri_purchase" + +const awsEC2RIPurchaseDescription = "Purchase AWS EC2 Reserved Instances. THIS SPENDS REAL MONEY when " + + "dry_run=false and confirm=true. Always call with dry_run=true first (the default) to validate your " + + "parameters before committing; a dry_run response never contacts AWS and never spends money. Search first " + + "with cudly_search_recommendations to find a region/instance_type/count worth reserving." + +// ec2RIPurchaseArgs is the input schema for cudly_aws_ec2_ri_purchase. term +// and payment_option map onto common.Recommendation.Term/PaymentOption; +// platform/tenancy/scope map onto common.ComputeDetails, which +// providers/aws/services/ec2/client.go:401-424 requires (Platform empty is a +// hard error there) -- they default to the overwhelmingly common case +// (on-demand Linux, shared tenancy, region-scoped) but are always visible in +// the schema and can be overridden per call. +type ec2RIPurchaseArgs struct { + Region string `json:"region" jsonschema:"AWS region, e.g. us-east-1"` + InstanceType string `json:"instance_type" jsonschema:"EC2 instance type, e.g. m5.large"` + Count int `json:"count" jsonschema:"number of instances to reserve, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + PaymentOption string `json:"payment_option" jsonschema:"payment schedule"` + Platform string `json:"platform,omitempty" jsonschema:"RI product description (operating system); defaults to Linux/UNIX"` + Tenancy string `json:"tenancy,omitempty" jsonschema:"instance tenancy; defaults to default (shared)"` + Scope string `json:"scope,omitempty" jsonschema:"region or availability-zone; defaults to region"` + AWSProfile string `json:"aws_profile,omitempty" jsonschema:"AWS named profile override (~/.aws/config); default uses ambient credentials"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type awsEC2RIPurchaseTool struct { + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewAWSEC2RIPurchaseTool builds the cudly_aws_ec2_ri_purchase tool. +func NewAWSEC2RIPurchaseTool() Registration { + return &awsEC2RIPurchaseTool{createProvider: provider.CreateProvider} +} + +func (t *awsEC2RIPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: awsEC2RIPurchaseName, + Provider: "aws", + Product: "ec2", + Action: "ri_purchase", + Description: awsEC2RIPurchaseDescription, + RealPurchaseEnabled: true, + ExamplePrompts: []string{ + "Preview buying 3 m5.large 3-year no-upfront RIs in us-east-1", + "Buy 2 r6g.large Reserved Instances in eu-west-1 for real, 1-year all-upfront", + }, + } +} + +func (t *awsEC2RIPurchaseTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[ec2RIPurchaseArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + "payment_option": {Enum: []any{string(PaymentOptionAllUpfront), string(PaymentOptionPartialUpfront), string(PaymentOptionNoUpfront)}}, + "scope": {Enum: []any{string(ScopeRegion), string(ScopeAvailabilityZone)}, Default: string(ScopeRegion)}, + "tenancy": {Enum: []any{string(TenancyDefault), string(TenancyDedicated)}, Default: string(TenancyDefault)}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: awsEC2RIPurchaseName, + Description: awsEC2RIPurchaseDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *awsEC2RIPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args ec2RIPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := ec2RecommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +// ec2ComputeDimensions holds the validated, possibly-defaulted EC2 RI +// dimensions that do not vary by resource/count/term -- split out of +// ec2RecommendationFromArgs to keep that function under the pre-commit +// gocyclo threshold. +type ec2ComputeDimensions struct { + platform ec2types.RIProductDescription + tenancy Tenancy + scope Scope +} + +// resolveEC2ComputeDimensions validates the optional platform/tenancy/scope +// fields, applying their documented defaults (Linux/UNIX, default tenancy, +// region scope) when the caller omits them. +func resolveEC2ComputeDimensions(args ec2RIPurchaseArgs) (ec2ComputeDimensions, error) { + dims := ec2ComputeDimensions{ + platform: ec2types.RIProductDescriptionLinuxUnix, + tenancy: TenancyDefault, + scope: ScopeRegion, + } + var err error + if args.Platform != "" { + if dims.platform, err = ValidatePlatform(args.Platform); err != nil { + return ec2ComputeDimensions{}, err + } + } + if args.Tenancy != "" { + if dims.tenancy, err = ValidateTenancy(args.Tenancy); err != nil { + return ec2ComputeDimensions{}, err + } + } + if args.Scope != "" { + if dims.scope, err = ValidateScope(args.Scope); err != nil { + return ec2ComputeDimensions{}, err + } + } + return dims, nil +} + +// effectiveDryRunConfirm applies the dry_run=true / confirm=false defaults: +// Go's zero value for bool cannot distinguish "caller omitted the field" +// from "caller explicitly set it false", so both flags are pointers and this +// is the single place that resolves them to concrete booleans. +func effectiveDryRunConfirm(args ec2RIPurchaseArgs) (dryRun, confirm bool) { + dryRun = true + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + return dryRun, confirm +} + +// ec2RecommendationFromArgs validates every field of args and builds the +// common.Recommendation to purchase, the effective region (trimmed of any +// surrounding whitespace), and the effective dry_run/confirm booleans. +func ec2RecommendationFromArgs(args ec2RIPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + region, err = requireNonBlank("region", args.Region) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + instanceType, err := requireNonBlank("instance_type", args.InstanceType) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + if args.Count <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("count must be > 0, got %d", args.Count) + } + term, err := ValidateTermYears(args.TermYears) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + paymentOption, err := ValidatePaymentOption(args.PaymentOption) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + dims, err := resolveEC2ComputeDimensions(args) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + rec = common.Recommendation{ + Provider: common.ProviderAWS, + Service: common.ServiceEC2, + Region: region, + ResourceType: instanceType, + Count: args.Count, + CommitmentType: common.CommitmentReservedInstance, + Term: term.RecommendationTerm(), + PaymentOption: string(paymentOption), + Details: &common.ComputeDetails{ + InstanceType: instanceType, + Platform: string(dims.platform), + Tenancy: string(dims.tenancy), + Scope: string(dims.scope), + }, + } + + dryRun, confirm = effectiveDryRunConfirm(args) + return rec, region, dryRun, confirm, nil +} + +// resolveClient returns the ResolveClientFunc that ExecutePurchase invokes +// only for a real purchase, so provider/credential resolution is deferred +// until after the dry_run/confirm gate has already decided to execute. +// region is the effective, already-validated-and-trimmed region returned by +// ec2RecommendationFromArgs -- not args.Region -- so a real purchase never +// resolves the provider/service client against a raw, un-trimmed value. +func (t *awsEC2RIPurchaseTool) resolveClient(args ec2RIPurchaseArgs, region string) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderAWS), AWSProfile: args.AWSProfile, Region: region} + prov, err := t.createProvider(string(common.ProviderAWS), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, common.ServiceEC2, region) + } +} diff --git a/mcp/tools/aws_ec2_ri_test.go b/mcp/tools/aws_ec2_ri_test.go new file mode 100644 index 000000000..2d9e9ff9f --- /dev/null +++ b/mcp/tools/aws_ec2_ri_test.go @@ -0,0 +1,225 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +func boolPtr(b bool) *bool { return &b } + +func validEC2Args() ec2RIPurchaseArgs { + return ec2RIPurchaseArgs{ + Region: "us-east-1", + InstanceType: "m5.large", + Count: 3, + TermYears: 3, + PaymentOption: "no-upfront", + } +} + +func TestEC2RecommendationFromArgsDefaults(t *testing.T) { + t.Parallel() + rec, region, dryRun, confirm, err := ec2RecommendationFromArgs(validEC2Args()) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region) + assert.True(t, dryRun, "dry_run must default to true") + assert.False(t, confirm, "confirm must default to false") + assert.Equal(t, common.ProviderAWS, rec.Provider) + assert.Equal(t, common.ServiceEC2, rec.Service) + assert.Equal(t, "m5.large", rec.ResourceType) + assert.Equal(t, 3, rec.Count) + assert.Equal(t, "3yr", rec.Term) + assert.Equal(t, "no-upfront", rec.PaymentOption) + details, ok := rec.Details.(*common.ComputeDetails) + require.True(t, ok, "Details must be *common.ComputeDetails") + assert.Equal(t, "Linux/UNIX", details.Platform) + assert.Equal(t, "default", details.Tenancy) + assert.Equal(t, "region", details.Scope) +} + +func TestEC2RecommendationFromArgsExplicitFlags(t *testing.T) { + t.Parallel() + args := validEC2Args() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + args.Platform = "Windows" + args.Tenancy = "dedicated" + args.Scope = "availability-zone" + + rec, _, dryRun, confirm, err := ec2RecommendationFromArgs(args) + require.NoError(t, err) + assert.False(t, dryRun) + assert.True(t, confirm) + details, ok := rec.Details.(*common.ComputeDetails) + require.True(t, ok) + assert.Equal(t, "Windows", details.Platform) + assert.Equal(t, "dedicated", details.Tenancy) + assert.Equal(t, "availability-zone", details.Scope) +} + +// TestEC2RecommendationFromArgsTrimsSurroundingWhitespace is the regression +// guard for the CodeRabbit finding: requireNonBlank rejected an all-whitespace +// value but let a value with surrounding whitespace (e.g. " us-east-1 ") +// through unchanged, which then flowed into rec.Region/rec.ResourceType and +// (via the returned region) into resolveClient's ProviderConfig.Region and +// GetServiceClient call. Both the region returned to the caller and every +// identifier field on rec must be the trimmed form. +func TestEC2RecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + args := validEC2Args() + args.Region = " us-east-1 " + args.InstanceType = " m5.large " + + rec, region, _, _, err := ec2RecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region, "returned region must be trimmed") + assert.Equal(t, "us-east-1", rec.Region, "rec.Region must be trimmed") + assert.Equal(t, "m5.large", rec.ResourceType, "rec.ResourceType must be trimmed") + details, ok := rec.Details.(*common.ComputeDetails) + require.True(t, ok) + assert.Equal(t, "m5.large", details.InstanceType, "Details.InstanceType must be trimmed") +} + +func TestEC2RecommendationFromArgsMissingRequiredFields(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*ec2RIPurchaseArgs) + errSub string + }{ + {"missing region", func(a *ec2RIPurchaseArgs) { a.Region = "" }, "region is required"}, + {"whitespace-only region", func(a *ec2RIPurchaseArgs) { a.Region = " " }, "region is required"}, + {"missing instance_type", func(a *ec2RIPurchaseArgs) { a.InstanceType = "" }, "instance_type is required"}, + {"whitespace-only instance_type", func(a *ec2RIPurchaseArgs) { a.InstanceType = "\t " }, "instance_type is required"}, + {"zero count", func(a *ec2RIPurchaseArgs) { a.Count = 0 }, "count must be"}, + {"negative count", func(a *ec2RIPurchaseArgs) { a.Count = -1 }, "count must be"}, + {"invalid term", func(a *ec2RIPurchaseArgs) { a.TermYears = 2 }, "invalid term_years"}, + {"invalid payment option", func(a *ec2RIPurchaseArgs) { a.PaymentOption = "bogus" }, "invalid payment_option"}, + {"invalid platform", func(a *ec2RIPurchaseArgs) { a.Platform = "MacOS" }, "invalid platform"}, + {"invalid tenancy", func(a *ec2RIPurchaseArgs) { a.Tenancy = "host" }, "invalid tenancy"}, + {"invalid scope", func(a *ec2RIPurchaseArgs) { a.Scope = "zonal" }, "invalid scope"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validEC2Args() + tc.mutate(&args) + _, _, _, _, err := ec2RecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestAWSEC2RIPurchaseHandleConfirmFalseRefuses proves the end-to-end tool +// handler refuses a dry_run=false, confirm=false call with a structured +// error, never touching the provider. +func TestAWSEC2RIPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsEC2RIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validEC2Args() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +// TestAWSEC2RIPurchaseHandleDryRunNeverCallsProvider proves the default +// dry_run=true path never resolves a provider, even with confirm=true. +func TestAWSEC2RIPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsEC2RIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validEC2Args() + args.Confirm = boolPtr(true) // dry_run stays at its true default + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) + assert.True(t, resp.Success) +} + +// TestAWSEC2RIPurchaseHandleInvalidArgsNeverCallsProvider proves a boundary +// validation failure (bad enum) short-circuits before any provider call. +func TestAWSEC2RIPurchaseHandleInvalidArgsNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsEC2RIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validEC2Args() + args.TermYears = 5 // invalid + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "invalid term_years") +} + +// TestAWSEC2RIPurchaseHandleRealPurchaseResolvesEC2Client proves a +// confirmed real purchase resolves the AWS EC2 service client for the +// requested region. +func TestAWSEC2RIPurchaseHandleRealPurchaseResolvesEC2Client(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "ri-abc"}} + var gotService common.ServiceType + var gotRegion string + fp := &fakeProvider{ + name: "aws", + } + tool := &awsEC2RIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{fakeProvider: fp, client: fake, gotService: &gotService, gotRegion: &gotRegion}, nil + }, + } + args := validEC2Args() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, "ri-abc", resp.CommitmentID) + assert.Equal(t, common.ServiceEC2, gotService) + assert.Equal(t, "us-east-1", gotRegion) + assert.Equal(t, 1, fake.purchaseCalls) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) +} + +// recordingProvider wraps fakeProvider to capture the service/region passed +// to GetServiceClient and always return a fixed ServiceClient. +type recordingProvider struct { + *fakeProvider + client provider.ServiceClient + gotService *common.ServiceType + gotRegion *string +} + +func (r *recordingProvider) GetServiceClient(_ context.Context, service common.ServiceType, region string) (provider.ServiceClient, error) { + *r.gotService = service + *r.gotRegion = region + return r.client, nil +} diff --git a/mcp/tools/aws_elasticache_ri.go b/mcp/tools/aws_elasticache_ri.go new file mode 100644 index 000000000..46023cfc7 --- /dev/null +++ b/mcp/tools/aws_elasticache_ri.go @@ -0,0 +1,166 @@ +package tools + +import ( + "context" + "fmt" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +const awsElastiCacheRIPurchaseName = "cudly_aws_elasticache_ri_purchase" + +const awsElastiCacheRIPurchaseDescription = "Purchase AWS ElastiCache Reserved Cache Nodes. THIS SPENDS REAL " + + "MONEY when dry_run=false and confirm=true. Always call with dry_run=true first (the default) to validate " + + "your parameters before committing; a dry_run response never contacts AWS and never spends money." + +// elasticacheRIPurchaseArgs is the input schema for +// cudly_aws_elasticache_ri_purchase. engine maps onto common.CacheDetails, +// which providers/aws/services/elasticache/client.go:273-283 requires for +// the offering lookup. +type elasticacheRIPurchaseArgs struct { + Region string `json:"region" jsonschema:"AWS region, e.g. us-east-1"` + NodeType string `json:"node_type" jsonschema:"ElastiCache cache node type, e.g. cache.r6g.large"` + Count int `json:"count" jsonschema:"number of cache nodes to reserve, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + PaymentOption string `json:"payment_option" jsonschema:"payment schedule"` + Engine string `json:"engine" jsonschema:"cache engine"` + AWSProfile string `json:"aws_profile,omitempty" jsonschema:"AWS named profile override (~/.aws/config); default uses ambient credentials"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type awsElastiCacheRIPurchaseTool struct { + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewAWSElastiCacheRIPurchaseTool builds the cudly_aws_elasticache_ri_purchase tool. +func NewAWSElastiCacheRIPurchaseTool() Registration { + return &awsElastiCacheRIPurchaseTool{createProvider: provider.CreateProvider} +} + +func (t *awsElastiCacheRIPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: awsElastiCacheRIPurchaseName, + Provider: "aws", + Product: "elasticache", + Action: "ri_purchase", + Description: awsElastiCacheRIPurchaseDescription, + RealPurchaseEnabled: true, + ExamplePrompts: []string{ + "Preview buying 3 cache.r6g.large redis ElastiCache RIs in us-east-1 for 1 year", + "Buy an ElastiCache Reserved Cache Node for memcached in eu-west-1 for real", + }, + } +} + +func (t *awsElastiCacheRIPurchaseTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[elasticacheRIPurchaseArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + "payment_option": {Enum: []any{string(PaymentOptionAllUpfront), string(PaymentOptionPartialUpfront), string(PaymentOptionNoUpfront)}}, + "engine": {Enum: []any{string(CacheEngineRedis), string(CacheEngineMemcached)}}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: awsElastiCacheRIPurchaseName, + Description: awsElastiCacheRIPurchaseDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *awsElastiCacheRIPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args elasticacheRIPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := elasticacheRecommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +// elasticacheRecommendationFromArgs validates args and builds the +// common.Recommendation to purchase, the effective region (trimmed of any +// surrounding whitespace), and the effective dry_run/confirm booleans. +func elasticacheRecommendationFromArgs(args elasticacheRIPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + region, err = requireNonBlank("region", args.Region) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + nodeType, err := requireNonBlank("node_type", args.NodeType) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + if args.Count <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("count must be > 0, got %d", args.Count) + } + term, err := ValidateTermYears(args.TermYears) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + paymentOption, err := ValidatePaymentOption(args.PaymentOption) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + engine, err := ValidateCacheEngine(args.Engine) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + rec = common.Recommendation{ + Provider: common.ProviderAWS, + Service: common.ServiceElastiCache, + Region: region, + ResourceType: nodeType, + Count: args.Count, + CommitmentType: common.CommitmentReservedInstance, + Term: term.RecommendationTerm(), + PaymentOption: string(paymentOption), + Details: &common.CacheDetails{ + Engine: string(engine), + NodeType: nodeType, + }, + } + + dryRun, confirm = true, false + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + return rec, region, dryRun, confirm, nil +} + +// resolveClient returns the ResolveClientFunc that ExecutePurchase invokes +// only for a real purchase. region is the effective, already-validated-and- +// trimmed region returned by elasticacheRecommendationFromArgs -- not +// args.Region -- so a real purchase never resolves the provider/service +// client against a raw, un-trimmed value. +func (t *awsElastiCacheRIPurchaseTool) resolveClient(args elasticacheRIPurchaseArgs, region string) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderAWS), AWSProfile: args.AWSProfile, Region: region} + prov, err := t.createProvider(string(common.ProviderAWS), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, common.ServiceElastiCache, region) + } +} diff --git a/mcp/tools/aws_elasticache_ri_test.go b/mcp/tools/aws_elasticache_ri_test.go new file mode 100644 index 000000000..acb435789 --- /dev/null +++ b/mcp/tools/aws_elasticache_ri_test.go @@ -0,0 +1,149 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +func validElastiCacheArgs() elasticacheRIPurchaseArgs { + return elasticacheRIPurchaseArgs{ + Region: "us-east-1", + NodeType: "cache.r6g.large", + Count: 3, + TermYears: 1, + PaymentOption: "all-upfront", + Engine: "redis", + } +} + +func TestElastiCacheRecommendationFromArgs(t *testing.T) { + t.Parallel() + rec, region, dryRun, confirm, err := elasticacheRecommendationFromArgs(validElastiCacheArgs()) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region) + assert.True(t, dryRun) + assert.False(t, confirm) + assert.Equal(t, common.ServiceElastiCache, rec.Service) + details, ok := rec.Details.(*common.CacheDetails) + require.True(t, ok) + assert.Equal(t, "redis", details.Engine) + assert.Equal(t, "cache.r6g.large", details.NodeType) +} + +func TestElastiCacheRecommendationFromArgsInvalid(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*elasticacheRIPurchaseArgs) + errSub string + }{ + {"missing region", func(a *elasticacheRIPurchaseArgs) { a.Region = "" }, "region is required"}, + {"whitespace-only region", func(a *elasticacheRIPurchaseArgs) { a.Region = " " }, "region is required"}, + {"missing node_type", func(a *elasticacheRIPurchaseArgs) { a.NodeType = "" }, "node_type is required"}, + {"whitespace-only node_type", func(a *elasticacheRIPurchaseArgs) { a.NodeType = "\t " }, "node_type is required"}, + {"zero count", func(a *elasticacheRIPurchaseArgs) { a.Count = 0 }, "count must be"}, + {"invalid term", func(a *elasticacheRIPurchaseArgs) { a.TermYears = 5 }, "invalid term_years"}, + {"invalid payment option", func(a *elasticacheRIPurchaseArgs) { a.PaymentOption = "bogus" }, "invalid payment_option"}, + {"missing engine", func(a *elasticacheRIPurchaseArgs) { a.Engine = "" }, "invalid engine"}, + {"invalid engine", func(a *elasticacheRIPurchaseArgs) { a.Engine = "postgres" }, "invalid engine"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validElastiCacheArgs() + tc.mutate(&args) + _, _, _, _, err := elasticacheRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestElastiCacheRecommendationFromArgsTrimsSurroundingWhitespace is the +// regression guard for the CodeRabbit finding: requireNonBlank rejected an +// all-whitespace value but let a value with surrounding whitespace (e.g. +// " us-east-1 ") pass through unchanged into rec.Region/rec.ResourceType/ +// Details and the returned region (which resolveClient uses for +// ProviderConfig.Region and GetServiceClient). +func TestElastiCacheRecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + args := validElastiCacheArgs() + args.Region = " us-east-1 " + args.NodeType = " cache.r6g.large " + + rec, region, _, _, err := elasticacheRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region, "returned region must be trimmed") + assert.Equal(t, "us-east-1", rec.Region, "rec.Region must be trimmed") + assert.Equal(t, "cache.r6g.large", rec.ResourceType, "rec.ResourceType must be trimmed") + details, ok := rec.Details.(*common.CacheDetails) + require.True(t, ok) + assert.Equal(t, "cache.r6g.large", details.NodeType, "Details.NodeType must be trimmed") +} + +func TestAWSElastiCacheRIPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsElastiCacheRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validElastiCacheArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +func TestAWSElastiCacheRIPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsElastiCacheRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validElastiCacheArgs() + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) +} + +func TestAWSElastiCacheRIPurchaseHandleRealPurchase(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "ec-ri-1"}} + var gotService common.ServiceType + tool := &awsElastiCacheRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "aws"}, + client: fake, + gotService: &gotService, + gotRegion: new(string), + }, nil + }, + } + args := validElastiCacheArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, common.ServiceElastiCache, gotService) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) +} diff --git a/mcp/tools/aws_rds_ri.go b/mcp/tools/aws_rds_ri.go new file mode 100644 index 000000000..645c9a1b8 --- /dev/null +++ b/mcp/tools/aws_rds_ri.go @@ -0,0 +1,174 @@ +package tools + +import ( + "context" + "fmt" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +const awsRDSRIPurchaseName = "cudly_aws_rds_ri_purchase" + +const awsRDSRIPurchaseDescription = "Purchase AWS RDS Reserved Instances. THIS SPENDS REAL MONEY when " + + "dry_run=false and confirm=true. Always call with dry_run=true first (the default) to validate your " + + "parameters before committing; a dry_run response never contacts AWS and never spends money." + +// rdsRIPurchaseArgs is the input schema for cudly_aws_rds_ri_purchase. +// engine and az_config map onto common.DatabaseDetails, which +// providers/aws/services/rds/client.go:301-322 requires -- az_config in +// particular has no safe default (single-AZ and multi-AZ RIs have different +// prices and do not cover each other's demand), so unlike EC2's +// platform/tenancy/scope it is a required field here, not a defaulted one. +type rdsRIPurchaseArgs struct { + Region string `json:"region" jsonschema:"AWS region, e.g. us-east-1"` + InstanceClass string `json:"instance_class" jsonschema:"RDS DB instance class, e.g. db.r6g.large"` + Count int `json:"count" jsonschema:"number of instances to reserve, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + PaymentOption string `json:"payment_option" jsonschema:"payment schedule"` + Engine string `json:"engine" jsonschema:"RDS database engine, e.g. mysql, postgres, mariadb, oracle-se2, sqlserver-ee"` + AZConfig string `json:"az_config" jsonschema:"single-az or multi-az; must match the recommendation exactly (different price, no cross-coverage)"` + AWSProfile string `json:"aws_profile,omitempty" jsonschema:"AWS named profile override (~/.aws/config); default uses ambient credentials"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type awsRDSRIPurchaseTool struct { + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewAWSRDSRIPurchaseTool builds the cudly_aws_rds_ri_purchase tool. +func NewAWSRDSRIPurchaseTool() Registration { + return &awsRDSRIPurchaseTool{createProvider: provider.CreateProvider} +} + +func (t *awsRDSRIPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: awsRDSRIPurchaseName, + Provider: "aws", + Product: "rds", + Action: "ri_purchase", + Description: awsRDSRIPurchaseDescription, + RealPurchaseEnabled: true, + ExamplePrompts: []string{ + "Preview buying 2 db.r6g.large multi-az postgres RDS RIs in us-east-1 for 3 years", + "Buy an RDS Reserved Instance for a single-az mysql db.t3.medium in eu-west-1 for real", + }, + } +} + +func (t *awsRDSRIPurchaseTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[rdsRIPurchaseArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + "payment_option": {Enum: []any{string(PaymentOptionAllUpfront), string(PaymentOptionPartialUpfront), string(PaymentOptionNoUpfront)}}, + "az_config": {Enum: []any{string(AZConfigSingleAZ), string(AZConfigMultiAZ)}}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: awsRDSRIPurchaseName, + Description: awsRDSRIPurchaseDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *awsRDSRIPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args rdsRIPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := rdsRecommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +// rdsRecommendationFromArgs validates args and builds the +// common.Recommendation to purchase, the effective region (trimmed of any +// surrounding whitespace), and the effective dry_run/confirm booleans. +func rdsRecommendationFromArgs(args rdsRIPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + region, err = requireNonBlank("region", args.Region) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + instanceClass, err := requireNonBlank("instance_class", args.InstanceClass) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + if args.Count <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("count must be > 0, got %d", args.Count) + } + engine, err := requireNonBlank("engine", args.Engine) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + term, err := ValidateTermYears(args.TermYears) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + paymentOption, err := ValidatePaymentOption(args.PaymentOption) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + azConfig, err := ValidateAZConfig(args.AZConfig) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + rec = common.Recommendation{ + Provider: common.ProviderAWS, + Service: common.ServiceRDS, + Region: region, + ResourceType: instanceClass, + Count: args.Count, + CommitmentType: common.CommitmentReservedInstance, + Term: term.RecommendationTerm(), + PaymentOption: string(paymentOption), + Details: &common.DatabaseDetails{ + Engine: engine, + AZConfig: string(azConfig), + InstanceClass: instanceClass, + }, + } + + dryRun, confirm = true, false + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + return rec, region, dryRun, confirm, nil +} + +// resolveClient returns the ResolveClientFunc that ExecutePurchase invokes +// only for a real purchase. region is the effective, already-validated-and- +// trimmed region returned by rdsRecommendationFromArgs -- not args.Region -- +// so a real purchase never resolves the provider/service client against a +// raw, un-trimmed value. +func (t *awsRDSRIPurchaseTool) resolveClient(args rdsRIPurchaseArgs, region string) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderAWS), AWSProfile: args.AWSProfile, Region: region} + prov, err := t.createProvider(string(common.ProviderAWS), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, common.ServiceRDS, region) + } +} diff --git a/mcp/tools/aws_rds_ri_test.go b/mcp/tools/aws_rds_ri_test.go new file mode 100644 index 000000000..664b69053 --- /dev/null +++ b/mcp/tools/aws_rds_ri_test.go @@ -0,0 +1,155 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +func validRDSArgs() rdsRIPurchaseArgs { + return rdsRIPurchaseArgs{ + Region: "us-east-1", + InstanceClass: "db.r6g.large", + Count: 2, + TermYears: 3, + PaymentOption: "no-upfront", + Engine: "postgres", + AZConfig: "multi-az", + } +} + +func TestRDSRecommendationFromArgs(t *testing.T) { + t.Parallel() + rec, region, dryRun, confirm, err := rdsRecommendationFromArgs(validRDSArgs()) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region) + assert.True(t, dryRun) + assert.False(t, confirm) + assert.Equal(t, common.ServiceRDS, rec.Service) + details, ok := rec.Details.(*common.DatabaseDetails) + require.True(t, ok) + assert.Equal(t, "postgres", details.Engine) + assert.Equal(t, "multi-az", details.AZConfig) + assert.Equal(t, "db.r6g.large", details.InstanceClass) +} + +func TestRDSRecommendationFromArgsInvalid(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*rdsRIPurchaseArgs) + errSub string + }{ + {"missing region", func(a *rdsRIPurchaseArgs) { a.Region = "" }, "region is required"}, + {"whitespace-only region", func(a *rdsRIPurchaseArgs) { a.Region = " " }, "region is required"}, + {"missing instance_class", func(a *rdsRIPurchaseArgs) { a.InstanceClass = "" }, "instance_class is required"}, + {"whitespace-only instance_class", func(a *rdsRIPurchaseArgs) { a.InstanceClass = "\t " }, "instance_class is required"}, + {"missing engine", func(a *rdsRIPurchaseArgs) { a.Engine = "" }, "engine is required"}, + {"whitespace-only engine", func(a *rdsRIPurchaseArgs) { a.Engine = "\t " }, "engine is required"}, + {"zero count", func(a *rdsRIPurchaseArgs) { a.Count = 0 }, "count must be"}, + {"invalid term", func(a *rdsRIPurchaseArgs) { a.TermYears = 2 }, "invalid term_years"}, + {"invalid payment option", func(a *rdsRIPurchaseArgs) { a.PaymentOption = "bogus" }, "invalid payment_option"}, + {"missing az_config refuses to guess", func(a *rdsRIPurchaseArgs) { a.AZConfig = "" }, "invalid az_config"}, + {"invalid az_config", func(a *rdsRIPurchaseArgs) { a.AZConfig = "triple-az" }, "invalid az_config"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validRDSArgs() + tc.mutate(&args) + _, _, _, _, err := rdsRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestRDSRecommendationFromArgsTrimsSurroundingWhitespace is the regression +// guard for the CodeRabbit finding: requireNonBlank rejected an +// all-whitespace value but let a value with surrounding whitespace (e.g. +// " us-east-1 ") pass through unchanged into rec.Region/rec.ResourceType/ +// Details and the returned region (which resolveClient uses for +// ProviderConfig.Region and GetServiceClient). +func TestRDSRecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + args := validRDSArgs() + args.Region = " us-east-1 " + args.InstanceClass = " db.r6g.large " + args.Engine = " postgres " + + rec, region, _, _, err := rdsRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region, "returned region must be trimmed") + assert.Equal(t, "us-east-1", rec.Region, "rec.Region must be trimmed") + assert.Equal(t, "db.r6g.large", rec.ResourceType, "rec.ResourceType must be trimmed") + details, ok := rec.Details.(*common.DatabaseDetails) + require.True(t, ok) + assert.Equal(t, "db.r6g.large", details.InstanceClass, "Details.InstanceClass must be trimmed") + assert.Equal(t, "postgres", details.Engine, "Details.Engine must be trimmed") +} + +func TestAWSRDSRIPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsRDSRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validRDSArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +func TestAWSRDSRIPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsRDSRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validRDSArgs() + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) +} + +func TestAWSRDSRIPurchaseHandleRealPurchase(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "rds-ri-1"}} + var gotService common.ServiceType + tool := &awsRDSRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "aws"}, + client: fake, + gotService: &gotService, + gotRegion: new(string), + }, nil + }, + } + args := validRDSArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, common.ServiceRDS, gotService) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) +} diff --git a/mcp/tools/aws_savingsplans.go b/mcp/tools/aws_savingsplans.go new file mode 100644 index 000000000..dbd09fd81 --- /dev/null +++ b/mcp/tools/aws_savingsplans.go @@ -0,0 +1,255 @@ +package tools + +import ( + "context" + "fmt" + "strings" + + spTypes "github.com/aws/aws-sdk-go-v2/service/savingsplans/types" + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" + "github.com/LeanerCloud/CUDly/providers/aws/services/savingsplans" +) + +const awsSavingsPlansPurchaseName = "cudly_aws_savingsplans_purchase" + +const awsSavingsPlansPurchaseDescription = "Purchase an AWS Savings Plan (Compute, EC2Instance, SageMaker, or " + + "Database). THIS SPENDS REAL MONEY when dry_run=false and confirm=true. Always call with dry_run=true " + + "first (the default) to validate your parameters before committing; a dry_run response never contacts AWS " + + "and never spends money. Unlike RI purchases this is dollar-denominated: you specify hourly_commitment " + + "(USD/hour), not an instance count. CAVEAT: sp_type=Database only supports term_years=1 and " + + "payment_option=no-upfront; AWS does not offer a 3-year Database Savings Plan or all-upfront/" + + "partial-upfront billing for it." + +// savingsPlansAccountLevelRegion is the region used to resolve the account- +// level Savings Plans service client when the caller omits region -- Compute, +// SageMaker, and Database plans are global, and cmd/multi_service_helpers.go +// already establishes this same "single query, us-east-1" convention for +// account-level Savings Plans recommendations. +const savingsPlansAccountLevelRegion = "us-east-1" + +// savingsPlansPurchaseArgs is the input schema for +// cudly_aws_savingsplans_purchase. instance_family and region are only +// meaningful for EC2Instance plans (common.SavingsPlanDetails); Compute, +// SageMaker, and Database plans are family-agnostic and account-level. +type savingsPlansPurchaseArgs struct { + SPType string `json:"sp_type" jsonschema:"AWS Savings Plans type"` + HourlyCommitment float64 `json:"hourly_commitment" jsonschema:"USD/hour commitment amount, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + PaymentOption string `json:"payment_option" jsonschema:"payment schedule"` + InstanceFamily string `json:"instance_family,omitempty" jsonschema:"EC2 instance family, e.g. m5; only meaningful for sp_type=EC2Instance"` + Region string `json:"region,omitempty" jsonschema:"AWS region; required for sp_type=EC2Instance, ignored for account-level plan types"` + AWSProfile string `json:"aws_profile,omitempty" jsonschema:"AWS named profile override (~/.aws/config); default uses ambient credentials"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type awsSavingsPlansPurchaseTool struct { + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewAWSSavingsPlansPurchaseTool builds the cudly_aws_savingsplans_purchase tool. +func NewAWSSavingsPlansPurchaseTool() Registration { + return &awsSavingsPlansPurchaseTool{createProvider: provider.CreateProvider} +} + +func (t *awsSavingsPlansPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: awsSavingsPlansPurchaseName, + Provider: "aws", + Product: "savingsplans", + Action: "purchase", + Description: awsSavingsPlansPurchaseDescription, + RealPurchaseEnabled: true, + ExamplePrompts: []string{ + "Preview a $10/hour Compute Savings Plan, 3-year no-upfront", + "Buy a $5/hour EC2Instance Savings Plan for the m5 family in us-east-1 for real", + }, + } +} + +func (t *awsSavingsPlansPurchaseTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[savingsPlansPurchaseArgs](map[string]FieldOverride{ + "sp_type": {Enum: []any{ + string(SPTypeCompute), string(SPTypeEC2Instance), string(SPTypeSageMaker), string(SPTypeDatabase), + }}, + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + "payment_option": {Enum: []any{string(PaymentOptionAllUpfront), string(PaymentOptionPartialUpfront), string(PaymentOptionNoUpfront)}}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: awsSavingsPlansPurchaseName, + Description: awsSavingsPlansPurchaseDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *awsSavingsPlansPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args savingsPlansPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := savingsPlanRecommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region, rec.Service), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +// validateSavingsPlanArgs validates every field of args that does not +// depend on the effective region, returning the typed sp_type, term, and +// payment_option. Split out of savingsPlanRecommendationFromArgs so that +// function's cyclomatic complexity stays under the repo's gocyclo gate as +// validation branches (e.g. validateDatabaseSPConstraints) are added. +func validateSavingsPlanArgs(args savingsPlansPurchaseArgs) (spType SPType, term TermYears, paymentOption PaymentOption, err error) { + if args.HourlyCommitment <= 0 { + return "", 0, "", fmt.Errorf("hourly_commitment must be > 0, got %v", args.HourlyCommitment) + } + spType, err = ValidateSPType(args.SPType) + if err != nil { + return "", 0, "", err + } + term, err = ValidateTermYears(args.TermYears) + if err != nil { + return "", 0, "", err + } + paymentOption, err = ValidatePaymentOption(args.PaymentOption) + if err != nil { + return "", 0, "", err + } + if spType == SPTypeEC2Instance && strings.TrimSpace(args.Region) == "" { + return "", 0, "", fmt.Errorf("region is required for sp_type=%s", SPTypeEC2Instance) + } + // instance_family is the filter that stops DescribeSavingsPlansOfferings + // from resolving to an arbitrary EC2Instance offering across every family + // in the region. providers/aws/services/savingsplans/client.go's + // lookupEC2OfferingIDStrict does fail loud when the resulting offerings + // span more than one family, but that is defense in depth at the API + // boundary; requiring the family here, at the tool boundary, catches the + // missing value before a real purchase attempt is even made. + if spType == SPTypeEC2Instance && strings.TrimSpace(args.InstanceFamily) == "" { + return "", 0, "", fmt.Errorf("instance_family is required for sp_type=%s", SPTypeEC2Instance) + } + if err := validateDatabaseSPConstraints(spType, term, paymentOption); err != nil { + return "", 0, "", err + } + return spType, term, paymentOption, nil +} + +// savingsPlanRecommendationFromArgs validates args and builds the +// common.Recommendation to purchase, the effective region to resolve the +// service client against, and the effective dry_run/confirm booleans. +func savingsPlanRecommendationFromArgs(args savingsPlansPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + spType, term, paymentOption, err := validateSavingsPlanArgs(args) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + // Trim once and use the trimmed value everywhere region/instance_family + // are stored or forwarded (resolved region, rec.Region, Details.Region, + // Details.InstanceFamily): validateSavingsPlanArgs above only trims for + // the blank-check, so a caller-supplied " us-east-1 " would otherwise + // flow raw into ProviderConfig/GetServiceClient and the real + // DescribeSavingsPlansOfferings lookup for an EC2Instance purchase. + trimmedRegion := strings.TrimSpace(args.Region) + trimmedInstanceFamily := strings.TrimSpace(args.InstanceFamily) + + region = trimmedRegion + if region == "" { + region = savingsPlansAccountLevelRegion + } + + // Resolve the precise per-plan-type ServiceType (e.g. + // ServiceSavingsPlansCompute) rather than the ServiceSavingsPlansAll + // umbrella sentinel, so GetServiceClient returns a client scoped to + // spType: providers/aws/services/savingsplans/client.go's + // resolveSPPlanType then rejects a mismatched Details.PlanType instead of + // silently buying whatever plan type happens to be in Details (defense in + // depth on top of the ValidateSPType check above). + service := savingsplans.ServiceTypeForPlanType(spTypes.SavingsPlanType(spType)) + + details := &common.SavingsPlanDetails{ + PlanType: string(spType), + HourlyCommitment: args.HourlyCommitment, + } + // InstanceFamily and Region are only meaningful for EC2Instance plans + // (common.SavingsPlanDetails documents both as "only populated for + // EC2Instance"); leaving them unset for Compute/SageMaker/Database keeps + // that contract instead of leaking a caller-supplied region/family into + // an account-level, family-agnostic plan's Details. + if spType == SPTypeEC2Instance { + details.InstanceFamily = trimmedInstanceFamily + details.Region = trimmedRegion + } + + rec = common.Recommendation{ + Provider: common.ProviderAWS, + Service: service, + Region: region, + CommitmentType: common.CommitmentSavingsPlan, + Term: term.RecommendationTerm(), + PaymentOption: string(paymentOption), + Details: details, + } + + dryRun, confirm = true, false + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + return rec, region, dryRun, confirm, nil +} + +// validateDatabaseSPConstraints rejects a Database Savings Plan request +// AWS's purchase API would itself reject: per AWS's Database Savings Plans +// announcement (aws.amazon.com/about-aws/whats-new/2025/12/database-savings-plans-savings), +// Database Savings Plans support only a one-year term billed no-upfront -- +// unlike Compute, EC2Instance, and SageMaker plans, there is no three-year +// term and no all-upfront/partial-upfront option. Failing loud here, before +// building the recommendation, surfaces AWS's real constraint instead of +// letting a real purchase reach AWS only to be rejected there. +func validateDatabaseSPConstraints(spType SPType, term TermYears, paymentOption PaymentOption) error { + if spType != SPTypeDatabase { + return nil + } + if term != TermOneYear { + return fmt.Errorf("sp_type=%s only supports a %d-year term (got term_years=%d): "+ + "AWS Database Savings Plans do not offer a %d-year term", + SPTypeDatabase, TermOneYear, term, TermThreeYear) + } + if paymentOption != PaymentOptionNoUpfront { + return fmt.Errorf("sp_type=%s only supports payment_option=%s (got %q): "+ + "AWS Database Savings Plans do not offer all-upfront or partial-upfront billing", + SPTypeDatabase, PaymentOptionNoUpfront, paymentOption) + } + return nil +} + +func (t *awsSavingsPlansPurchaseTool) resolveClient(args savingsPlansPurchaseArgs, region string, service common.ServiceType) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderAWS), AWSProfile: args.AWSProfile, Region: region} + prov, err := t.createProvider(string(common.ProviderAWS), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, service, region) + } +} diff --git a/mcp/tools/aws_savingsplans_test.go b/mcp/tools/aws_savingsplans_test.go new file mode 100644 index 000000000..ed1c05d75 --- /dev/null +++ b/mcp/tools/aws_savingsplans_test.go @@ -0,0 +1,428 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +func validSavingsPlansArgs() savingsPlansPurchaseArgs { + return savingsPlansPurchaseArgs{ + SPType: "Compute", + HourlyCommitment: 10.50, + TermYears: 3, + PaymentOption: "no-upfront", + } +} + +func TestSavingsPlanRecommendationFromArgsAccountLevel(t *testing.T) { + t.Parallel() + rec, region, dryRun, confirm, err := savingsPlanRecommendationFromArgs(validSavingsPlansArgs()) + require.NoError(t, err) + assert.True(t, dryRun) + assert.False(t, confirm) + assert.Equal(t, savingsPlansAccountLevelRegion, region, "account-level plan defaults to the shared query region") + assert.Equal(t, common.ServiceSavingsPlansCompute, rec.Service) + assert.Equal(t, common.CommitmentSavingsPlan, rec.CommitmentType) + assert.Equal(t, "3yr", rec.Term) + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Equal(t, "Compute", details.PlanType) + assert.InDelta(t, 10.50, details.HourlyCommitment, 0.001) +} + +// TestSavingsPlanRecommendationFromArgsAccountLevelWhitespaceRegion proves a +// whitespace-only region (e.g. " ") for an account-level sp_type (Compute, +// SageMaker, Database) still falls back to savingsPlansAccountLevelRegion, +// the same as an empty region does. Before the fix, the fallback only +// triggered on region == "", so a whitespace-only region threaded the raw +// " " value into resolveClient instead of the account-level default. +func TestSavingsPlanRecommendationFromArgsAccountLevelWhitespaceRegion(t *testing.T) { + t.Parallel() + args := validSavingsPlansArgs() + args.Region = " " + rec, region, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, savingsPlansAccountLevelRegion, region, + "whitespace-only region must resolve to the account-level default, not be threaded through as-is") + assert.Equal(t, common.ServiceSavingsPlansCompute, rec.Service) +} + +func TestSavingsPlanRecommendationFromArgsEC2InstanceRequiresRegion(t *testing.T) { + t.Parallel() + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.InstanceFamily = "m5" + + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), "region is required") + + args.Region = "us-east-1" + rec, region, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region) + assert.Equal(t, common.ServiceSavingsPlansEC2Instance, rec.Service) + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Equal(t, "m5", details.InstanceFamily) + assert.Equal(t, "us-east-1", details.Region) +} + +// TestSavingsPlanRecommendationFromArgsEC2InstanceRejectsWhitespaceOnly +// proves region and instance_family are rejected when they contain only +// whitespace, not just when they are the empty string: a bare `== ""` check +// would let " " through to a real EC2Instance Savings Plan purchase. +func TestSavingsPlanRecommendationFromArgsEC2InstanceRejectsWhitespaceOnly(t *testing.T) { + t.Parallel() + + t.Run("whitespace-only region", func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.InstanceFamily = "m5" + args.Region = " " + + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), "region is required") + }) + + t.Run("whitespace-only instance_family", func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.Region = "us-east-1" + args.InstanceFamily = "\t " + + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), "instance_family is required") + }) +} + +// TestSavingsPlanRecommendationFromArgsEC2InstanceRequiresInstanceFamily is +// the regression guard for the CodeRabbit money-path finding: omitting +// instance_family for sp_type=EC2Instance lets DescribeSavingsPlansOfferings +// resolve across every instance family in the region instead of the one +// Cost Explorer actually recommended, risking a real purchase for the wrong +// workload. instance_family must be required exactly like region already is. +func TestSavingsPlanRecommendationFromArgsEC2InstanceRequiresInstanceFamily(t *testing.T) { + t.Parallel() + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.Region = "us-east-1" + + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.Error(t, err, "EC2Instance sp_type without instance_family must be rejected") + assert.Contains(t, err.Error(), "instance_family is required") + + args.InstanceFamily = "m5" + rec, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err, "EC2Instance sp_type with instance_family set must succeed") + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Equal(t, "m5", details.InstanceFamily) +} + +// TestSavingsPlanRecommendationFromArgsInstanceFamilyOptionalForOtherTypes +// proves the new instance_family requirement is scoped to sp_type=EC2Instance +// only: Compute, SageMaker, and Database plans are family-agnostic and +// account-level, so instance_family stays optional (and ignored) for them. +func TestSavingsPlanRecommendationFromArgsInstanceFamilyOptionalForOtherTypes(t *testing.T) { + t.Parallel() + for _, spType := range []string{"Compute", "SageMaker"} { + t.Run(spType, func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = spType + args.InstanceFamily = "" + + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err, "instance_family must remain optional for sp_type=%s", spType) + }) + } + + t.Run("Database", func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = "Database" + args.TermYears = 1 + args.PaymentOption = "no-upfront" + args.InstanceFamily = "" + + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err, "instance_family must remain optional for sp_type=Database") + }) +} + +// TestSavingsPlanRecommendationFromArgsDetailsRegionFamilyOnlyForEC2Instance +// is the regression guard for the CodeRabbit money-path finding: Details.Region +// and Details.InstanceFamily are documented on common.SavingsPlanDetails as +// "only populated for EC2Instance"; providers/aws/services/savingsplans/client.go +// uses Details.Region as a DescribeSavingsPlansOfferings filter. A +// caller-supplied region/instance_family for an account-level sp_type +// (Compute, SageMaker, Database) must NOT leak into Details -- before the +// fix, both fields were populated unconditionally from args, regardless of +// sp_type. +func TestSavingsPlanRecommendationFromArgsDetailsRegionFamilyOnlyForEC2Instance(t *testing.T) { + t.Parallel() + + for _, spType := range []string{"Compute", "SageMaker"} { + t.Run(spType, func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = spType + args.Region = "eu-west-1" + args.InstanceFamily = "m5" + + rec, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Empty(t, details.Region, "account-level sp_type=%s must not carry a caller-supplied region into Details", spType) + assert.Empty(t, details.InstanceFamily, "account-level sp_type=%s must not carry a caller-supplied instance_family into Details", spType) + }) + } + + t.Run("Database", func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = "Database" + args.TermYears = 1 + args.PaymentOption = "no-upfront" + args.Region = "eu-west-1" + args.InstanceFamily = "m5" + + rec, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Empty(t, details.Region, "sp_type=Database must not carry a caller-supplied region into Details") + assert.Empty(t, details.InstanceFamily, "sp_type=Database must not carry a caller-supplied instance_family into Details") + }) + + t.Run("EC2Instance still carries both", func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.Region = "us-east-1" + args.InstanceFamily = "m5" + + rec, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Equal(t, "us-east-1", details.Region) + assert.Equal(t, "m5", details.InstanceFamily) + }) +} + +// TestSavingsPlanRecommendationFromArgsTrimsSurroundingWhitespace is the +// regression guard for the CodeRabbit money-path finding: +// validateSavingsPlanArgs only trims region/instance_family for the +// blank-check (its `strings.TrimSpace(args.Region) == ""` guards), but +// savingsPlanRecommendationFromArgs then stored the RAW args.Region and +// args.InstanceFamily into the resolved region, rec.Region, and +// Details.Region/Details.InstanceFamily -- which flow into +// ProviderConfig/GetServiceClient and the DescribeSavingsPlansOfferings +// lookup for a real EC2Instance Savings Plan purchase. Before the fix, +// " us-east-1 " passed validation but reached a real purchase with the +// surrounding whitespace intact. +func TestSavingsPlanRecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.Region = " us-east-1 " + args.InstanceFamily = " m5 " + + rec, region, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region, "returned region must be trimmed") + assert.Equal(t, "us-east-1", rec.Region, "rec.Region must be trimmed") + details, ok := rec.Details.(*common.SavingsPlanDetails) + require.True(t, ok) + assert.Equal(t, "us-east-1", details.Region, "Details.Region must be trimmed") + assert.Equal(t, "m5", details.InstanceFamily, "Details.InstanceFamily must be trimmed") +} + +// TestAWSSavingsPlansPurchaseHandleForwardsTrimmedRegionToServiceClient +// proves the trimmed region -- not the raw, whitespace-padded args.Region -- +// is what actually reaches GetServiceClient on a real EC2Instance Savings +// Plan purchase, i.e. the fix holds through the full handle() path, not just +// the recommendation-building helper. +func TestAWSSavingsPlansPurchaseHandleForwardsTrimmedRegionToServiceClient(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "sp-2"}} + var gotService common.ServiceType + var gotRegion string + tool := &awsSavingsPlansPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "aws"}, + client: fake, + gotService: &gotService, + gotRegion: &gotRegion, + }, nil + }, + } + args := validSavingsPlansArgs() + args.SPType = "EC2Instance" + args.Region = " us-east-1 " + args.InstanceFamily = " m5 " + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, "us-east-1", gotRegion, "GetServiceClient must receive the trimmed region, not raw whitespace") +} + +func TestSavingsPlanRecommendationFromArgsInvalid(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*savingsPlansPurchaseArgs) + errSub string + }{ + {"zero hourly commitment", func(a *savingsPlansPurchaseArgs) { a.HourlyCommitment = 0 }, "hourly_commitment must be"}, + {"negative hourly commitment", func(a *savingsPlansPurchaseArgs) { a.HourlyCommitment = -5 }, "hourly_commitment must be"}, + {"invalid sp_type", func(a *savingsPlansPurchaseArgs) { a.SPType = "Storage" }, "invalid sp_type"}, + {"invalid term", func(a *savingsPlansPurchaseArgs) { a.TermYears = 2 }, "invalid term_years"}, + {"invalid payment option", func(a *savingsPlansPurchaseArgs) { a.PaymentOption = "bogus" }, "invalid payment_option"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validSavingsPlansArgs() + tc.mutate(&args) + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestSavingsPlanRecommendationFromArgsDatabaseConstraints proves the +// CodeRabbit-requested up-front validation: per AWS's Database Savings +// Plans announcement, sp_type=Database only supports a one-year term +// billed no-upfront -- unlike Compute/EC2Instance/SageMaker, there is no +// 3-year term and no all-upfront/partial-upfront option. A mismatched +// term_years or payment_option must be rejected before building the +// recommendation, not left for AWS's purchase API to reject. +func TestSavingsPlanRecommendationFromArgsDatabaseConstraints(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*savingsPlansPurchaseArgs) + errSub string + }{ + {"3yr term rejected", func(a *savingsPlansPurchaseArgs) { a.TermYears = 3 }, "only supports a 1-year term"}, + {"all-upfront rejected", func(a *savingsPlansPurchaseArgs) { a.PaymentOption = "all-upfront" }, "only supports payment_option=no-upfront"}, + {"partial-upfront rejected", func(a *savingsPlansPurchaseArgs) { a.PaymentOption = "partial-upfront" }, "only supports payment_option=no-upfront"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validSavingsPlansArgs() + args.SPType = "Database" + args.TermYears = 1 + args.PaymentOption = "no-upfront" + tc.mutate(&args) + _, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestSavingsPlanRecommendationFromArgsDatabaseAllowedCombo proves the one +// term/payment_option combination Database Savings Plans actually support +// still succeeds. +func TestSavingsPlanRecommendationFromArgsDatabaseAllowedCombo(t *testing.T) { + t.Parallel() + args := validSavingsPlansArgs() + args.SPType = "Database" + args.TermYears = 1 + args.PaymentOption = "no-upfront" + + rec, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, common.ServiceSavingsPlansDatabase, rec.Service) + assert.Equal(t, "1yr", rec.Term) +} + +// TestSavingsPlanRecommendationFromArgsNonDatabaseUnaffected proves the +// Database-only constraint does not leak onto other sp_types: Compute keeps +// supporting 3-year all-upfront, the combo Database rejects. +func TestSavingsPlanRecommendationFromArgsNonDatabaseUnaffected(t *testing.T) { + t.Parallel() + args := validSavingsPlansArgs() + args.SPType = "Compute" + args.TermYears = 3 + args.PaymentOption = "all-upfront" + + rec, _, _, _, err := savingsPlanRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "3yr", rec.Term) + assert.Equal(t, "all-upfront", rec.PaymentOption) +} + +func TestAWSSavingsPlansPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsSavingsPlansPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validSavingsPlansArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +func TestAWSSavingsPlansPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &awsSavingsPlansPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validSavingsPlansArgs() + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) +} + +func TestAWSSavingsPlansPurchaseHandleRealPurchaseUsesScopedService(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "sp-1"}} + var gotService common.ServiceType + tool := &awsSavingsPlansPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "aws"}, + client: fake, + gotService: &gotService, + gotRegion: new(string), + }, nil + }, + } + args := validSavingsPlansArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, common.ServiceSavingsPlansCompute, gotService, "must resolve the plan-type-scoped client, not the umbrella sentinel") + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) +} diff --git a/mcp/tools/aws_simple_ri.go b/mcp/tools/aws_simple_ri.go new file mode 100644 index 000000000..241baa166 --- /dev/null +++ b/mcp/tools/aws_simple_ri.go @@ -0,0 +1,215 @@ +package tools + +import ( + "context" + "fmt" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +// simpleAWSRIPurchaseSpec configures a region+resource_type+count+term+ +// payment_option AWS RI purchase tool -- the shape shared by OpenSearch, +// Redshift, and MemoryDB. None of their PurchaseCommitment implementations +// read rec.Details (providers/aws/services/{opensearch,redshift,memorydb}/ +// client.go), unlike EC2 (ComputeDetails) or RDS/ElastiCache (Database/ +// CacheDetails), so one generic tool type serves all three rather than +// three near-identical copies. +type simpleAWSRIPurchaseSpec struct { + name string + product string + displayName string // human-readable product name for the tool description, e.g. "OpenSearch" + service common.ServiceType + resourceTypeDesc string // jsonschema description for the resource_type field + examplePrompts []string +} + +// simpleAWSRIPurchaseArgs is the input schema shared by every +// simpleAWSRIPurchaseTool instance. +type simpleAWSRIPurchaseArgs struct { + Region string `json:"region" jsonschema:"AWS region, e.g. us-east-1"` + ResourceType string `json:"resource_type" jsonschema:"resource/node type to reserve"` + Count int `json:"count" jsonschema:"number of nodes/instances to reserve, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + PaymentOption string `json:"payment_option" jsonschema:"payment schedule"` + AWSProfile string `json:"aws_profile,omitempty" jsonschema:"AWS named profile override (~/.aws/config); default uses ambient credentials"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type simpleAWSRIPurchaseTool struct { + spec simpleAWSRIPurchaseSpec + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// newSimpleAWSRIPurchaseTool builds a Registration for spec. +func newSimpleAWSRIPurchaseTool(spec simpleAWSRIPurchaseSpec) Registration { + return &simpleAWSRIPurchaseTool{spec: spec, createProvider: provider.CreateProvider} +} + +// NewAWSOpenSearchRIPurchaseTool builds cudly_aws_opensearch_ri_purchase. +func NewAWSOpenSearchRIPurchaseTool() Registration { + return newSimpleAWSRIPurchaseTool(simpleAWSRIPurchaseSpec{ + name: "cudly_aws_opensearch_ri_purchase", + product: "opensearch", + displayName: "OpenSearch", + service: common.ServiceOpenSearch, + resourceTypeDesc: "OpenSearch instance type, e.g. r6g.large.search", + examplePrompts: []string{ + "Preview buying 2 r6g.large.search OpenSearch RIs in us-east-1 for 1 year", + "Buy an OpenSearch Reserved Instance in eu-west-1 for real", + }, + }) +} + +// NewAWSRedshiftRIPurchaseTool builds cudly_aws_redshift_ri_purchase. +func NewAWSRedshiftRIPurchaseTool() Registration { + return newSimpleAWSRIPurchaseTool(simpleAWSRIPurchaseSpec{ + name: "cudly_aws_redshift_ri_purchase", + product: "redshift", + displayName: "Redshift", + service: common.ServiceRedshift, + resourceTypeDesc: "Redshift node type, e.g. dc2.large", + examplePrompts: []string{ + "Preview buying 4 dc2.large Redshift RIs in us-east-1 for 3 years, all-upfront", + }, + }) +} + +// NewAWSMemoryDBRIPurchaseTool builds cudly_aws_memorydb_ri_purchase. +func NewAWSMemoryDBRIPurchaseTool() Registration { + return newSimpleAWSRIPurchaseTool(simpleAWSRIPurchaseSpec{ + name: "cudly_aws_memorydb_ri_purchase", + product: "memorydb", + displayName: "MemoryDB", + service: common.ServiceMemoryDB, + resourceTypeDesc: "MemoryDB node type, e.g. db.r6g.large", + examplePrompts: []string{ + "Preview buying 2 db.r6g.large MemoryDB RIs in us-east-1", + }, + }) +} + +func (t *simpleAWSRIPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: t.spec.name, + Provider: "aws", + Product: t.spec.product, + Action: "ri_purchase", + Description: fmt.Sprintf( + "Purchase AWS %s Reserved Instances. THIS SPENDS REAL MONEY when dry_run=false and confirm=true. "+ + "Always call with dry_run=true first (the default) to validate your parameters before "+ + "committing; a dry_run response never contacts AWS and never spends money.", + t.spec.displayName), + RealPurchaseEnabled: true, + ExamplePrompts: t.spec.examplePrompts, + } +} + +func (t *simpleAWSRIPurchaseTool) Register(s *mcp.Server) error { + desc := t.Descriptor().Description + schema, err := BuildInputSchema[simpleAWSRIPurchaseArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + "payment_option": {Enum: []any{string(PaymentOptionAllUpfront), string(PaymentOptionPartialUpfront), string(PaymentOptionNoUpfront)}}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + // resource_type's description is spec-specific (differs per product), + // so it is set directly rather than through a generic FieldOverride. + if prop, ok := schema.Properties["resource_type"]; ok { + prop.Description = t.spec.resourceTypeDesc + } + mcp.AddTool(s, &mcp.Tool{ + Name: t.spec.name, + Description: desc, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *simpleAWSRIPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args simpleAWSRIPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := t.recommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +// recommendationFromArgs validates args and builds the common.Recommendation +// to purchase, the effective region (trimmed of any surrounding whitespace), +// and the effective dry_run/confirm booleans. +func (t *simpleAWSRIPurchaseTool) recommendationFromArgs(args simpleAWSRIPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + region, err = requireNonBlank("region", args.Region) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + resourceType, err := requireNonBlank("resource_type", args.ResourceType) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + if args.Count <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("count must be > 0, got %d", args.Count) + } + term, err := ValidateTermYears(args.TermYears) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + paymentOption, err := ValidatePaymentOption(args.PaymentOption) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + rec = common.Recommendation{ + Provider: common.ProviderAWS, + Service: t.spec.service, + Region: region, + ResourceType: resourceType, + Count: args.Count, + CommitmentType: common.CommitmentReservedInstance, + Term: term.RecommendationTerm(), + PaymentOption: string(paymentOption), + } + + dryRun, confirm = true, false + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + return rec, region, dryRun, confirm, nil +} + +// resolveClient returns the ResolveClientFunc that ExecutePurchase invokes +// only for a real purchase. region is the effective, already-validated-and- +// trimmed region returned by recommendationFromArgs -- not args.Region -- +// so a real purchase never resolves the provider/service client against a +// raw, un-trimmed value. +func (t *simpleAWSRIPurchaseTool) resolveClient(args simpleAWSRIPurchaseArgs, region string) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderAWS), AWSProfile: args.AWSProfile, Region: region} + prov, err := t.createProvider(string(common.ProviderAWS), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, t.spec.service, region) + } +} diff --git a/mcp/tools/aws_simple_ri_test.go b/mcp/tools/aws_simple_ri_test.go new file mode 100644 index 000000000..a48c8d9bc --- /dev/null +++ b/mcp/tools/aws_simple_ri_test.go @@ -0,0 +1,210 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +// simpleToolConstructors covers every simpleAWSRIPurchaseTool instance so +// the shared safety-rail behavior (confirm gate, dry_run gate, boundary +// validation, real-purchase wiring) is proven once per product rather than +// hand-copied three times. +func simpleToolConstructors() map[string]func() Registration { + return map[string]func() Registration{ + "opensearch": NewAWSOpenSearchRIPurchaseTool, + "redshift": NewAWSRedshiftRIPurchaseTool, + "memorydb": NewAWSMemoryDBRIPurchaseTool, + } +} + +func validSimpleArgs() simpleAWSRIPurchaseArgs { + return simpleAWSRIPurchaseArgs{ + Region: "us-east-1", + ResourceType: "r6g.large", + Count: 2, + TermYears: 1, + PaymentOption: "all-upfront", + } +} + +func TestSimpleAWSRIPurchaseDescriptorsAreDistinctAndRealPurchaseEnabled(t *testing.T) { + t.Parallel() + names := map[string]bool{} + for product, ctor := range simpleToolConstructors() { + d := ctor().Descriptor() + assert.True(t, d.RealPurchaseEnabled, "%s must be real-purchase enabled", product) + assert.False(t, names[d.Name], "duplicate tool name %q", d.Name) + names[d.Name] = true + assert.NotEmpty(t, d.ExamplePrompts, "%s must document example prompts", product) + } +} + +// TestSimpleAWSRIPurchaseDescriptorUsesProperlyCasedDisplayName proves the +// CodeRabbit finding: t.spec.product ("opensearch", "redshift", "memorydb") +// used to be interpolated raw into the human-readable description, +// producing "AWS opensearch Reserved Instances" instead of the properly +// cased "AWS OpenSearch Reserved Instances". The identifier used in API +// calls (spec.product) must stay lowercase; only the description text uses +// the display name. +func TestSimpleAWSRIPurchaseDescriptorUsesProperlyCasedDisplayName(t *testing.T) { + t.Parallel() + wantDisplayName := map[string]string{ + "opensearch": "OpenSearch", + "redshift": "Redshift", + "memorydb": "MemoryDB", + } + for product, ctor := range simpleToolConstructors() { + t.Run(product, func(t *testing.T) { + d := ctor().Descriptor() + want := wantDisplayName[product] + require.NotEmpty(t, want, "test table missing a display name for %s", product) + assert.Contains(t, d.Description, want) + assert.NotContains(t, d.Description, "AWS "+product+" ", "description must not use the raw lowercase identifier") + }) + } +} + +func TestSimpleAWSRIPurchaseRecommendationFromArgs(t *testing.T) { + t.Parallel() + for product, ctor := range simpleToolConstructors() { + t.Run(product, func(t *testing.T) { + tool := ctor().(*simpleAWSRIPurchaseTool) + rec, region, dryRun, confirm, err := tool.recommendationFromArgs(validSimpleArgs()) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region) + assert.True(t, dryRun) + assert.False(t, confirm) + assert.Equal(t, common.ProviderAWS, rec.Provider) + assert.Equal(t, tool.spec.service, rec.Service) + assert.Equal(t, "r6g.large", rec.ResourceType) + assert.Equal(t, 2, rec.Count) + assert.Equal(t, "1yr", rec.Term) + assert.Equal(t, "all-upfront", rec.PaymentOption) + assert.Nil(t, rec.Details, "%s must not require service Details", product) + }) + } +} + +func TestSimpleAWSRIPurchaseInvalidArgs(t *testing.T) { + t.Parallel() + tool := NewAWSOpenSearchRIPurchaseTool().(*simpleAWSRIPurchaseTool) + cases := []struct { + name string + mutate func(*simpleAWSRIPurchaseArgs) + errSub string + }{ + {"missing region", func(a *simpleAWSRIPurchaseArgs) { a.Region = "" }, "region is required"}, + {"whitespace-only region", func(a *simpleAWSRIPurchaseArgs) { a.Region = " " }, "region is required"}, + {"missing resource_type", func(a *simpleAWSRIPurchaseArgs) { a.ResourceType = "" }, "resource_type is required"}, + {"whitespace-only resource_type", func(a *simpleAWSRIPurchaseArgs) { a.ResourceType = "\t " }, "resource_type is required"}, + {"zero count", func(a *simpleAWSRIPurchaseArgs) { a.Count = 0 }, "count must be"}, + {"invalid term", func(a *simpleAWSRIPurchaseArgs) { a.TermYears = 4 }, "invalid term_years"}, + {"invalid payment option", func(a *simpleAWSRIPurchaseArgs) { a.PaymentOption = "bogus" }, "invalid payment_option"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validSimpleArgs() + tc.mutate(&args) + _, _, _, _, err := tool.recommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestSimpleAWSRIPurchaseRecommendationFromArgsTrimsSurroundingWhitespace is +// the regression guard for the CodeRabbit finding: requireNonBlank rejected +// an all-whitespace value but let a value with surrounding whitespace (e.g. +// " us-east-1 ") pass through unchanged into rec.Region/rec.ResourceType and +// the returned region (which resolveClient uses for ProviderConfig.Region +// and GetServiceClient). +func TestSimpleAWSRIPurchaseRecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + tool := NewAWSOpenSearchRIPurchaseTool().(*simpleAWSRIPurchaseTool) + args := validSimpleArgs() + args.Region = " us-east-1 " + args.ResourceType = " r6g.large " + + rec, region, _, _, err := tool.recommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-east-1", region, "returned region must be trimmed") + assert.Equal(t, "us-east-1", rec.Region, "rec.Region must be trimmed") + assert.Equal(t, "r6g.large", rec.ResourceType, "rec.ResourceType must be trimmed") +} + +func TestSimpleAWSRIPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + for product, ctor := range simpleToolConstructors() { + t.Run(product, func(t *testing.T) { + tool := ctor().(*simpleAWSRIPurchaseTool) + resolveCalled := false + tool.createProvider = func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + } + args := validSimpleArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") + }) + } +} + +func TestSimpleAWSRIPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + for product, ctor := range simpleToolConstructors() { + t.Run(product, func(t *testing.T) { + tool := ctor().(*simpleAWSRIPurchaseTool) + resolveCalled := false + tool.createProvider = func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + } + args := validSimpleArgs() + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) + }) + } +} + +func TestSimpleAWSRIPurchaseHandleRealPurchaseCallsCorrectService(t *testing.T) { + t.Parallel() + for product, ctor := range simpleToolConstructors() { + t.Run(product, func(t *testing.T) { + tool := ctor().(*simpleAWSRIPurchaseTool) + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "res-1"}} + var gotService common.ServiceType + tool.createProvider = func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "aws"}, + client: fake, + gotService: &gotService, + gotRegion: new(string), + }, nil + } + args := validSimpleArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, tool.spec.service, gotService) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) + }) + } +} diff --git a/mcp/tools/azure_compute_ri.go b/mcp/tools/azure_compute_ri.go new file mode 100644 index 000000000..6f31dfe5c --- /dev/null +++ b/mcp/tools/azure_compute_ri.go @@ -0,0 +1,195 @@ +package tools + +import ( + "context" + "fmt" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +const azureComputeRIPurchaseName = "cudly_azure_compute_ri_purchase" + +// azureComputeRIPurchaseDescription documents Azure's actual billing-plan +// contract: providers/azure/services/compute/client.go's buildReservationBody +// sends properties.billingPlan (armreservations.ReservationBillingPlan -- +// Upfront or Monthly), so both all-upfront and no-upfront purchases are +// honored for real. Monthly costs the same total as Upfront -- Azure has no +// premium for spreading payments -- but there is no partial-upfront billing +// plan at all, so that value is rejected with an explicit error rather than +// silently purchased under a different schedule (see +// azureComputeRecommendationFromArgs). +const azureComputeRIPurchaseDescription = "Purchase an Azure VM Reserved Instance. THIS SPENDS REAL MONEY when " + + "dry_run=false and confirm=true. Always call with dry_run=true first (the default) to validate your " + + "parameters before committing; a dry_run response never contacts Azure and never spends money. Azure " + + "Reserved Instances support two billing plans: all-upfront and no-upfront (billed monthly, same total " + + "price as all-upfront -- Azure charges no premium for spreading payments). payment_option defaults to " + + "no-upfront when omitted. Azure has no partial-upfront billing plan, so that value is rejected with an " + + "explicit error rather than silently purchased under all-upfront or no-upfront instead." + +// azureComputeRIPurchaseArgs is the input schema for +// cudly_azure_compute_ri_purchase. Unlike EC2, Azure's purchase body needs no +// Recommendation.Details -- providers/azure/services/compute/client.go's +// buildReservationBody only reads Region/ResourceType/Count/Term/PaymentOption. +type azureComputeRIPurchaseArgs struct { + Region string `json:"region" jsonschema:"Azure region, e.g. eastus"` + VMSize string `json:"vm_size" jsonschema:"Azure VM size (SKU), e.g. Standard_D2s_v3"` + Count int `json:"count" jsonschema:"number of VM instances to reserve, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + PaymentOption string `json:"payment_option,omitempty" jsonschema:"payment schedule; Azure honors all-upfront and no-upfront (monthly, same total price); no partial-upfront; defaults to no-upfront"` + AzureSubscriptionID string `json:"azure_subscription_id,omitempty" jsonschema:"Azure subscription ID override; default uses AZURE_SUBSCRIPTION_ID"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type azureComputeRIPurchaseTool struct { + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewAzureComputeRIPurchaseTool builds the cudly_azure_compute_ri_purchase tool. +func NewAzureComputeRIPurchaseTool() Registration { + return &azureComputeRIPurchaseTool{createProvider: provider.CreateProvider} +} + +func (t *azureComputeRIPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: azureComputeRIPurchaseName, + Provider: "azure", + Product: "compute", + Action: "ri_purchase", + Description: azureComputeRIPurchaseDescription, + RealPurchaseEnabled: true, + ExamplePrompts: []string{ + "Preview buying 2 Standard_D2s_v3 Azure VM RIs in eastus for 3 years", + "Buy an Azure VM Reserved Instance for real in westeurope", + }, + } +} + +func (t *azureComputeRIPurchaseTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[azureComputeRIPurchaseArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + // Azure has no partial-upfront billing plan (see + // azureComputeRIPurchaseDescription and azureComputeRecommendationFromArgs + // below), so this tool's schema advertises only the two values Azure + // actually honors. The runtime check in azureComputeRecommendationFromArgs + // still rejects partial-upfront explicitly, as defense in depth for a + // caller that bypasses the schema. + "payment_option": {Enum: []any{string(PaymentOptionAllUpfront), string(PaymentOptionNoUpfront)}, Default: string(PaymentOptionNoUpfront)}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: azureComputeRIPurchaseName, + Description: azureComputeRIPurchaseDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *azureComputeRIPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args azureComputeRIPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := azureComputeRecommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +func azureComputeRecommendationFromArgs(args azureComputeRIPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + region, err = requireNonBlank("region", args.Region) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + vmSize, err := requireNonBlank("vm_size", args.VMSize) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + if args.Count <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("count must be > 0, got %d", args.Count) + } + term, err := ValidateTermYears(args.TermYears) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + // payment_option defaults to no-upfront (matching the CLI's --payment + // default, cmd/main.go) when the caller omits it -- an omitted string + // field arrives as "" and is never confused with an explicit, + // unrecognized value (feedback_no_silent_fallbacks: the default is + // applied here, explicitly, not fabricated deeper in the stack). + paymentOptionStr := args.PaymentOption + if paymentOptionStr == "" { + paymentOptionStr = string(PaymentOptionNoUpfront) + } + paymentOption, err := ValidatePaymentOption(paymentOptionStr) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + // Azure reservations support exactly two billing plans (Upfront, + // Monthly -- see providers/azure/services/internal/reservations. + // BillingPlanForPaymentOption); there is no partial-upfront at any + // layer of Azure's API. Rejecting it here, unconditionally (not just + // for a real purchase), means a dry_run preview never reports success + // for a request that could never be honored for real. + if paymentOption == PaymentOptionPartialUpfront { + return common.Recommendation{}, "", false, false, fmt.Errorf( + "azure reservations do not support payment_option=%q: azure billing plans are all-upfront or "+ + "no-upfront (monthly, same total price) only, with no partial-upfront option", + paymentOption) + } + + dryRun, confirm = true, false + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + + rec = common.Recommendation{ + Provider: common.ProviderAzure, + Service: common.ServiceCompute, + Region: region, + ResourceType: vmSize, + Count: args.Count, + CommitmentType: common.CommitmentReservedInstance, + Term: term.RecommendationTerm(), + PaymentOption: string(paymentOption), + } + + return rec, region, dryRun, confirm, nil +} + +// resolveClient returns the ResolveClientFunc that ExecutePurchase invokes +// only for a real purchase. region is the effective, already-validated-and- +// trimmed region returned by azureComputeRecommendationFromArgs -- not +// args.Region -- so a real purchase never resolves the provider/service +// client against a raw, un-trimmed value. +func (t *azureComputeRIPurchaseTool) resolveClient(args azureComputeRIPurchaseArgs, region string) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderAzure), AzureSubscriptionID: args.AzureSubscriptionID, Region: region} + prov, err := t.createProvider(string(common.ProviderAzure), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, common.ServiceCompute, region) + } +} diff --git a/mcp/tools/azure_compute_ri_test.go b/mcp/tools/azure_compute_ri_test.go new file mode 100644 index 000000000..c5c6719fc --- /dev/null +++ b/mcp/tools/azure_compute_ri_test.go @@ -0,0 +1,319 @@ +package tools + +import ( + "context" + "fmt" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +// validAzureComputeArgs uses payment_option=all-upfront. Azure Reserved +// Instances honor both all-upfront and no-upfront (see +// TestAzureComputeRecommendationFromArgsAcceptsNoUpfront); all-upfront is +// used here just to keep the baseline args deterministic across tests that +// don't care which honored schedule they exercise. +func validAzureComputeArgs() azureComputeRIPurchaseArgs { + return azureComputeRIPurchaseArgs{ + Region: "eastus", + VMSize: "Standard_D2s_v3", + Count: 2, + TermYears: 3, + PaymentOption: "all-upfront", + } +} + +func TestAzureComputeRecommendationFromArgs(t *testing.T) { + t.Parallel() + rec, region, dryRun, confirm, err := azureComputeRecommendationFromArgs(validAzureComputeArgs()) + require.NoError(t, err) + assert.Equal(t, "eastus", region) + assert.True(t, dryRun) + assert.False(t, confirm) + assert.Equal(t, common.ProviderAzure, rec.Provider) + assert.Equal(t, common.ServiceCompute, rec.Service) + assert.Equal(t, "Standard_D2s_v3", rec.ResourceType) + assert.Equal(t, 2, rec.Count) + assert.Equal(t, "3yr", rec.Term) + assert.Nil(t, rec.Details, "Azure VM purchase reads no Recommendation.Details") +} + +// TestAzureComputeRecommendationFromArgsTrimsSurroundingWhitespace is the +// regression guard for the CodeRabbit finding: requireNonBlank rejected an +// all-whitespace value but let a value with surrounding whitespace (e.g. +// " eastus ") pass through unchanged into rec.Region/rec.ResourceType and the +// returned region (which resolveClient uses for ProviderConfig.Region and +// GetServiceClient). +func TestAzureComputeRecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + args := validAzureComputeArgs() + args.Region = " eastus " + args.VMSize = " Standard_D2s_v3 " + + rec, region, _, _, err := azureComputeRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "eastus", region, "returned region must be trimmed") + assert.Equal(t, "eastus", rec.Region, "rec.Region must be trimmed") + assert.Equal(t, "Standard_D2s_v3", rec.ResourceType, "rec.ResourceType must be trimmed") +} + +func TestAzureComputeRecommendationFromArgsInvalid(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*azureComputeRIPurchaseArgs) + errSub string + }{ + {"missing region", func(a *azureComputeRIPurchaseArgs) { a.Region = "" }, "region is required"}, + {"whitespace-only region", func(a *azureComputeRIPurchaseArgs) { a.Region = " " }, "region is required"}, + {"missing vm_size", func(a *azureComputeRIPurchaseArgs) { a.VMSize = "" }, "vm_size is required"}, + {"whitespace-only vm_size", func(a *azureComputeRIPurchaseArgs) { a.VMSize = "\t\n " }, "vm_size is required"}, + {"zero count", func(a *azureComputeRIPurchaseArgs) { a.Count = 0 }, "count must be"}, + {"invalid term", func(a *azureComputeRIPurchaseArgs) { a.TermYears = 2 }, "invalid term_years"}, + {"invalid payment option", func(a *azureComputeRIPurchaseArgs) { a.PaymentOption = "bogus" }, "invalid payment_option"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validAzureComputeArgs() + tc.mutate(&args) + _, _, _, _, err := azureComputeRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestAzureComputeRecommendationFromArgsRejectsPartialUpfront proves Azure's +// billing-plan contract has exactly two members (Upfront, Monthly -- see +// providers/azure/services/internal/reservations.BillingPlanForPaymentOption): +// partial-upfront has no Azure equivalent at any layer, so it must be +// rejected with an explicit error rather than silently purchased under +// all-upfront or no-upfront instead. Unlike the former all-upfront-only gate +// (removed once billingPlan wiring landed), this rejection is unconditional: +// it fires for a dry_run preview too, because Azure can never honor +// partial-upfront for real, not just a gap in this tool's own behavior. +func TestAzureComputeRecommendationFromArgsRejectsPartialUpfront(t *testing.T) { + t.Parallel() + for _, dryRun := range []bool{true, false} { + t.Run(fmt.Sprintf("dry_run=%v", dryRun), func(t *testing.T) { + args := validAzureComputeArgs() + args.PaymentOption = "partial-upfront" + args.DryRun = boolPtr(dryRun) + args.Confirm = boolPtr(true) + _, _, _, _, err := azureComputeRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), "partial-upfront") + assert.Contains(t, err.Error(), "no-upfront") + }) + } +} + +// TestAzureComputeRecommendationFromArgsAcceptsAllUpfront proves the +// all-upfront billing plan is honored for real (dry_run=false, confirm=true). +func TestAzureComputeRecommendationFromArgsAcceptsAllUpfront(t *testing.T) { + t.Parallel() + args := validAzureComputeArgs() + args.PaymentOption = "all-upfront" + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + rec, _, dryRun, confirm, err := azureComputeRecommendationFromArgs(args) + require.NoError(t, err) + assert.False(t, dryRun) + assert.True(t, confirm) + assert.Equal(t, "all-upfront", rec.PaymentOption) +} + +// TestAzureComputeRecommendationFromArgsAcceptsNoUpfront proves the +// no-upfront billing plan (armreservations.ReservationBillingPlanMonthly) is +// honored for real, not just at preview time -- the gap this PR closes. +func TestAzureComputeRecommendationFromArgsAcceptsNoUpfront(t *testing.T) { + t.Parallel() + args := validAzureComputeArgs() + args.PaymentOption = "no-upfront" + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + rec, _, dryRun, confirm, err := azureComputeRecommendationFromArgs(args) + require.NoError(t, err) + assert.False(t, dryRun) + assert.True(t, confirm) + assert.Equal(t, "no-upfront", rec.PaymentOption) +} + +// TestAzureComputeRecommendationFromArgsDefaultsToNoUpfront proves omitting +// payment_option defaults to no-upfront (matching the CLI's --payment +// default, cmd/main.go), not to Azure's raw API default (all-upfront) and +// not to an error. +func TestAzureComputeRecommendationFromArgsDefaultsToNoUpfront(t *testing.T) { + t.Parallel() + args := validAzureComputeArgs() + args.PaymentOption = "" + rec, _, _, _, err := azureComputeRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "no-upfront", rec.PaymentOption) +} + +func TestAzureComputeRIPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validAzureComputeArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +// TestAzureComputeRIPurchaseHandleDryRunAcceptsNoUpfront proves a preview +// validates and accepts no-upfront (it is now an honored billing plan, not +// merely tolerated at preview time). +func TestAzureComputeRIPurchaseHandleDryRunAcceptsNoUpfront(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validAzureComputeArgs() + args.PaymentOption = "no-upfront" + args.DryRun = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) +} + +// TestAzureComputeRIPurchaseHandleRealPurchaseRejectsPartialUpfront proves +// the one payment_option Azure cannot express (partial-upfront) is still +// rejected for a real purchase after the billingPlan wiring landed. +func TestAzureComputeRIPurchaseHandleRealPurchaseRejectsPartialUpfront(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validAzureComputeArgs() + args.PaymentOption = "partial-upfront" + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "partial-upfront") +} + +// TestAzureComputeRIPurchaseHandleRealPurchaseNoUpfront proves a real +// purchase (dry_run=false, confirm=true) with payment_option=no-upfront +// reaches the provider -- the core gap this PR closes: before billingPlan +// wiring, only all-upfront could execute for real. +func TestAzureComputeRIPurchaseHandleRealPurchaseNoUpfront(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "azure-res-monthly"}} + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "azure"}, + client: fake, + gotService: new(common.ServiceType), + gotRegion: new(string), + }, nil + }, + } + args := validAzureComputeArgs() + args.PaymentOption = "no-upfront" + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, "no-upfront", fake.purchaseResult.Recommendation.PaymentOption) +} + +// TestAzureComputeRIPurchaseHandleOmittedPaymentOptionDefaultsToNoUpfront +// proves the tool-level default (payment_option omitted from the request) +// flows through the handler the same way an explicit no-upfront does. +func TestAzureComputeRIPurchaseHandleOmittedPaymentOptionDefaultsToNoUpfront(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "azure-res-default"}} + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "azure"}, + client: fake, + gotService: new(common.ServiceType), + gotRegion: new(string), + }, nil + }, + } + args := validAzureComputeArgs() + args.PaymentOption = "" + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, "no-upfront", fake.purchaseResult.Recommendation.PaymentOption) +} + +func TestAzureComputeRIPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validAzureComputeArgs() + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) +} + +func TestAzureComputeRIPurchaseHandleRealPurchase(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "azure-res-1"}} + var gotService common.ServiceType + tool := &azureComputeRIPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "azure"}, + client: fake, + gotService: &gotService, + gotRegion: new(string), + }, nil + }, + } + args := validAzureComputeArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, common.ServiceCompute, gotService) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) +} diff --git a/mcp/tools/enums.go b/mcp/tools/enums.go new file mode 100644 index 000000000..7b8157e2f --- /dev/null +++ b/mcp/tools/enums.go @@ -0,0 +1,241 @@ +// Package tools implements the individual MCP tool handlers exposed by the +// CUDly MCP server (mcp/server.go), plus the shared validation and purchase +// harness they all build on. +package tools + +import ( + "fmt" + "strings" + + ec2types "github.com/aws/aws-sdk-go-v2/service/ec2/types" +) + +// requireNonBlank returns the trimmed form of val, or an explicit +// " is required" error when val is empty or contains only whitespace. +// A whitespace-only value (e.g. " ") passes a bare `== ""` check but +// carries no real region/instance-type/etc information, and on a +// real-purchase path could still reach provider resolution instead of being +// rejected at the MCP tool boundary like an actually-empty value already +// is. Callers must store the returned trimmed value (not the raw val) back +// onto the field so surrounding whitespace (e.g. " us-east-1 ") is +// normalized before it reaches a provider call rather than passed through +// raw. +func requireNonBlank(field, val string) (string, error) { + trimmed := strings.TrimSpace(val) + if trimmed == "" { + return "", fmt.Errorf("%s is required", field) + } + return trimmed, nil +} + +// PaymentOption is the AWS/Azure/GCP-agnostic reserved-capacity payment +// schedule. It is validated at the MCP tool boundary before being copied +// onto common.Recommendation.PaymentOption (which stays a bare string there +// for backward compatibility with existing CSV/DB rows -- see +// pkg/common/types.go) so a caller can never smuggle an unrecognized payment +// term into a purchase. +type PaymentOption string + +const ( + PaymentOptionAllUpfront PaymentOption = "all-upfront" + PaymentOptionPartialUpfront PaymentOption = "partial-upfront" + PaymentOptionNoUpfront PaymentOption = "no-upfront" +) + +// ValidatePaymentOption returns the typed PaymentOption for s, or an explicit +// error when s is not one of the three allowed values. There is no default: +// an empty or unknown payment option is always an error, never silently +// coerced to a fallback term (feedback_no_silent_fallbacks). +func ValidatePaymentOption(s string) (PaymentOption, error) { + switch PaymentOption(s) { + case PaymentOptionAllUpfront, PaymentOptionPartialUpfront, PaymentOptionNoUpfront: + return PaymentOption(s), nil + default: + return "", fmt.Errorf("invalid payment_option %q: must be one of %s, %s, %s", + s, PaymentOptionAllUpfront, PaymentOptionPartialUpfront, PaymentOptionNoUpfront) + } +} + +// TermYears is the reserved-capacity commitment length, in years. +type TermYears int + +const ( + TermOneYear TermYears = 1 + TermThreeYear TermYears = 3 +) + +// ValidateTermYears returns the typed TermYears for n, or an explicit error +// when n is not 1 or 3 (the only terms AWS/Azure/GCP reserved-capacity +// products offer). +func ValidateTermYears(n int) (TermYears, error) { + switch TermYears(n) { + case TermOneYear, TermThreeYear: + return TermYears(n), nil + default: + return 0, fmt.Errorf("invalid term_years %d: must be %d or %d", n, TermOneYear, TermThreeYear) + } +} + +// RecommendationTerm renders t in the "1yr"/"3yr" vocabulary that +// common.Recommendation.Term and the provider clients expect. +func (t TermYears) RecommendationTerm() string { + return fmt.Sprintf("%dyr", int(t)) +} + +// SPType is the AWS Savings Plans product family (--include-sp-types in the +// CLI, cmd/main.go:112). +type SPType string + +const ( + SPTypeCompute SPType = "Compute" + SPTypeEC2Instance SPType = "EC2Instance" + SPTypeSageMaker SPType = "SageMaker" + SPTypeDatabase SPType = "Database" +) + +// ValidateSPType returns the typed SPType for s, or an explicit error when s +// is not one of the four AWS Savings Plans product families. +func ValidateSPType(s string) (SPType, error) { + switch SPType(s) { + case SPTypeCompute, SPTypeEC2Instance, SPTypeSageMaker, SPTypeDatabase: + return SPType(s), nil + default: + return "", fmt.Errorf("invalid sp_type %q: must be one of %s, %s, %s, %s", + s, SPTypeCompute, SPTypeEC2Instance, SPTypeSageMaker, SPTypeDatabase) + } +} + +// AZConfig is the RDS deployment topology (single-AZ vs multi-AZ), which +// carries a different price and offering catalog per +// providers/aws/services/rds/client.go:314-322. +type AZConfig string + +const ( + AZConfigSingleAZ AZConfig = "single-az" + AZConfigMultiAZ AZConfig = "multi-az" +) + +// ValidateAZConfig returns the typed AZConfig for s, or an explicit error +// when s is not single-az or multi-az. RDS's own client refuses to guess this +// value (see the comment at providers/aws/services/rds/client.go:306-322), so +// the MCP boundary must not default it either. +func ValidateAZConfig(s string) (AZConfig, error) { + switch AZConfig(s) { + case AZConfigSingleAZ, AZConfigMultiAZ: + return AZConfig(s), nil + default: + return "", fmt.Errorf("invalid az_config %q: must be %s or %s", s, AZConfigSingleAZ, AZConfigMultiAZ) + } +} + +// ValidatePlatform returns the AWS SDK's own ec2types.RIProductDescription +// enum member for s, or an explicit error when s does not match one of the +// four values that DescribeReservedInstancesOfferings accepts as +// ProductDescription (providers/aws/services/ec2/client.go:419). Reusing the +// SDK's own enum constants -- rather than inventing a "linux"/"windows" +// vocabulary -- means an outbound offering lookup can never carry a bare +// string literal that drifts from what the SDK actually recognizes +// (feedback_sdk_enum_string_literals). +func ValidatePlatform(s string) (ec2types.RIProductDescription, error) { + switch ec2types.RIProductDescription(s) { + case ec2types.RIProductDescriptionLinuxUnix, + ec2types.RIProductDescriptionLinuxUnixAmazonVpc, + ec2types.RIProductDescriptionWindows, + ec2types.RIProductDescriptionWindowsAmazonVpc: + return ec2types.RIProductDescription(s), nil + default: + return "", fmt.Errorf("invalid platform %q: must be one of %s, %s, %s, %s", s, + ec2types.RIProductDescriptionLinuxUnix, ec2types.RIProductDescriptionLinuxUnixAmazonVpc, + ec2types.RIProductDescriptionWindows, ec2types.RIProductDescriptionWindowsAmazonVpc) + } +} + +// Tenancy is the EC2 RI tenancy dimension. Values match ec2types.Tenancy +// (providers/aws/services/ec2/client.go:309-318 canonicalizes them further, +// but "default"/"dedicated" already pass through unchanged). +type Tenancy string + +const ( + TenancyDefault Tenancy = Tenancy(ec2types.TenancyDefault) + TenancyDedicated Tenancy = Tenancy(ec2types.TenancyDedicated) +) + +// ValidateTenancy returns the typed Tenancy for s, or an explicit error when +// s is neither default nor dedicated. +func ValidateTenancy(s string) (Tenancy, error) { + switch Tenancy(s) { + case TenancyDefault, TenancyDedicated: + return Tenancy(s), nil + default: + return "", fmt.Errorf("invalid tenancy %q: must be %s or %s", s, TenancyDefault, TenancyDedicated) + } +} + +// Scope is the EC2 RI applicability dimension. Values are the lowercase, +// hyphenated form that providers/aws/services/ec2/client.go:330-339 +// (canonicalizeEC2Scope) recognizes and normalizes to the SDK's +// ec2types.Scope casing ("Region" / "Availability Zone"). +type Scope string + +const ( + ScopeRegion Scope = "region" + ScopeAvailabilityZone Scope = "availability-zone" +) + +// ValidateScope returns the typed Scope for s, or an explicit error when s is +// neither region nor availability-zone. +func ValidateScope(s string) (Scope, error) { + switch Scope(s) { + case ScopeRegion, ScopeAvailabilityZone: + return Scope(s), nil + default: + return "", fmt.Errorf("invalid scope %q: must be %s or %s", s, ScopeRegion, ScopeAvailabilityZone) + } +} + +// LookbackPeriod is the cost/usage lookback window backing a +// cudly_search_recommendations call. Values match the enum +// search_recommendations.go's Register advertises in the tool's JSON schema +// (BuildInputSchema's "lookback_period" FieldOverride); re-validated here so +// a caller invoking the tool directly -- bypassing MCP schema enforcement -- +// cannot pass an unsupported window through to the provider. +type LookbackPeriod string + +const ( + LookbackPeriod7Days LookbackPeriod = "7d" + LookbackPeriod30Days LookbackPeriod = "30d" + LookbackPeriod60Days LookbackPeriod = "60d" +) + +// ValidateLookbackPeriod returns the typed LookbackPeriod for s, or an +// explicit error when s is non-empty and not one of the three supported +// windows. An empty s is valid: lookback_period is optional +// (json:"...,omitempty"), meaning "let the provider apply its own default". +func ValidateLookbackPeriod(s string) (LookbackPeriod, error) { + switch LookbackPeriod(s) { + case "", LookbackPeriod7Days, LookbackPeriod30Days, LookbackPeriod60Days: + return LookbackPeriod(s), nil + default: + return "", fmt.Errorf("invalid lookback_period %q: must be one of %s, %s, %s", + s, LookbackPeriod7Days, LookbackPeriod30Days, LookbackPeriod60Days) + } +} + +// CacheEngine is the ElastiCache engine dimension (common.CacheDetails.Engine). +type CacheEngine string + +const ( + CacheEngineRedis CacheEngine = "redis" + CacheEngineMemcached CacheEngine = "memcached" +) + +// ValidateCacheEngine returns the typed CacheEngine for s, or an explicit +// error when s is neither redis nor memcached. +func ValidateCacheEngine(s string) (CacheEngine, error) { + switch CacheEngine(s) { + case CacheEngineRedis, CacheEngineMemcached: + return CacheEngine(s), nil + default: + return "", fmt.Errorf("invalid engine %q: must be %s or %s", s, CacheEngineRedis, CacheEngineMemcached) + } +} diff --git a/mcp/tools/enums_test.go b/mcp/tools/enums_test.go new file mode 100644 index 000000000..234d485e3 --- /dev/null +++ b/mcp/tools/enums_test.go @@ -0,0 +1,231 @@ +package tools + +import ( + "testing" + + ec2types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/stretchr/testify/assert" +) + +func TestValidatePaymentOption(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want PaymentOption + wantErr bool + }{ + {"all-upfront", "all-upfront", PaymentOptionAllUpfront, false}, + {"partial-upfront", "partial-upfront", PaymentOptionPartialUpfront, false}, + {"no-upfront", "no-upfront", PaymentOptionNoUpfront, false}, + {"empty", "", "", true}, + {"unknown", "some-upfront", "", true}, + {"case sensitive", "All-Upfront", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidatePaymentOption(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestValidateTermYears(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in int + want TermYears + wantErr bool + }{ + {"one year", 1, TermOneYear, false}, + {"three year", 3, TermThreeYear, false}, + {"zero", 0, 0, true}, + {"two years unsupported", 2, 0, true}, + {"negative", -1, 0, true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidateTermYears(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestTermYearsRecommendationTerm(t *testing.T) { + t.Parallel() + assert.Equal(t, "1yr", TermOneYear.RecommendationTerm()) + assert.Equal(t, "3yr", TermThreeYear.RecommendationTerm()) +} + +func TestValidateSPType(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want SPType + wantErr bool + }{ + {"compute", "Compute", SPTypeCompute, false}, + {"ec2instance", "EC2Instance", SPTypeEC2Instance, false}, + {"sagemaker", "SageMaker", SPTypeSageMaker, false}, + {"database", "Database", SPTypeDatabase, false}, + {"lowercase rejected", "compute", "", true}, + {"empty", "", "", true}, + {"unknown", "Storage", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidateSPType(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestValidateAZConfig(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want AZConfig + wantErr bool + }{ + {"single-az", "single-az", AZConfigSingleAZ, false}, + {"multi-az", "multi-az", AZConfigMultiAZ, false}, + {"empty refuses to guess", "", "", true}, + {"unknown", "triple-az", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidateAZConfig(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestValidatePlatform(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want ec2types.RIProductDescription + wantErr bool + }{ + {"linux", "Linux/UNIX", ec2types.RIProductDescriptionLinuxUnix, false}, + {"linux vpc", "Linux/UNIX (Amazon VPC)", ec2types.RIProductDescriptionLinuxUnixAmazonVpc, false}, + {"windows", "Windows", ec2types.RIProductDescriptionWindows, false}, + {"windows vpc", "Windows (Amazon VPC)", ec2types.RIProductDescriptionWindowsAmazonVpc, false}, + {"lowercase rejected", "linux", "", true}, + {"empty", "", "", true}, + {"unknown os", "MacOS", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidatePlatform(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestValidateTenancy(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want Tenancy + wantErr bool + }{ + {"default", "default", TenancyDefault, false}, + {"dedicated", "dedicated", TenancyDedicated, false}, + {"empty", "", "", true}, + {"host unsupported", "host", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidateTenancy(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestValidateCacheEngine(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want CacheEngine + wantErr bool + }{ + {"redis", "redis", CacheEngineRedis, false}, + {"memcached", "memcached", CacheEngineMemcached, false}, + {"empty", "", "", true}, + {"unknown", "postgres", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidateCacheEngine(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +func TestValidateScope(t *testing.T) { + t.Parallel() + cases := []struct { + name string + in string + want Scope + wantErr bool + }{ + {"region", "region", ScopeRegion, false}, + {"availability-zone", "availability-zone", ScopeAvailabilityZone, false}, + {"empty", "", "", true}, + {"sdk casing rejected", "Region", "", true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ValidateScope(tc.in) + if tc.wantErr { + assert.Error(t, err) + return + } + assert.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} diff --git a/mcp/tools/gcp_computeengine_cud.go b/mcp/tools/gcp_computeengine_cud.go new file mode 100644 index 000000000..abf78826c --- /dev/null +++ b/mcp/tools/gcp_computeengine_cud.go @@ -0,0 +1,164 @@ +package tools + +import ( + "context" + "fmt" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +const gcpComputeEngineCUDPurchaseName = "cudly_gcp_computeengine_cud_purchase" + +const gcpComputeEngineCUDPurchaseDescription = "Purchase a GCP Compute Engine Committed Use Discount (CUD). THIS " + + "SPENDS REAL MONEY when dry_run=false and confirm=true. Always call with dry_run=true first (the default) " + + "to validate your parameters before committing; a dry_run response never contacts GCP and never spends " + + "money. A CUD commits vCPUs and memory directly (not an instance count): vcpu_count is the number of vCPUs " + + "and memory_gb is the amount of memory to commit." + +// gcpComputeEngineCUDPurchaseArgs is the input schema for +// cudly_gcp_computeengine_cud_purchase. memory_gb is required: unlike AWS/ +// Azure, providers/gcp/services/computeengine/client.go's buildInsertRequest +// reads Recommendation.Details as a *value* common.ComputeDetails (not a +// pointer, unlike every AWS Details assertion) and hard-errors when +// MemoryGB is absent or <= 0 rather than guessing a vCPU:memory ratio. +type gcpComputeEngineCUDPurchaseArgs struct { + Region string `json:"region" jsonschema:"GCP region, e.g. us-central1"` + MachineType string `json:"machine_type" jsonschema:"GCP machine type family for the commitment, e.g. n2-standard-4"` + VCPUCount int `json:"vcpu_count" jsonschema:"number of vCPUs to commit, must be > 0"` + MemoryGB float64 `json:"memory_gb" jsonschema:"amount of memory (GB) to commit, must be > 0"` + TermYears int `json:"term_years" jsonschema:"commitment length in years"` + GCPProjectID string `json:"gcp_project_id,omitempty" jsonschema:"GCP project ID override; default uses ambient project"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase; defaults to true"` + Confirm *bool `json:"confirm,omitempty" jsonschema:"required (with dry_run=false) to execute a real purchase; defaults to false"` + IdempotencyNonce string `json:"idempotency_nonce,omitempty" jsonschema:"optional; set to a fresh value to authorize a purchase that is otherwise identical to a previous one (e.g. buy 3 more RIs with the same parameters); leave empty (the default) so retries with identical parameters dedupe and never double-buy"` +} + +type gcpComputeEngineCUDPurchaseTool struct { + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewGCPComputeEngineCUDPurchaseTool builds the cudly_gcp_computeengine_cud_purchase tool. +func NewGCPComputeEngineCUDPurchaseTool() Registration { + return &gcpComputeEngineCUDPurchaseTool{createProvider: provider.CreateProvider} +} + +func (t *gcpComputeEngineCUDPurchaseTool) Descriptor() Descriptor { + return Descriptor{ + Name: gcpComputeEngineCUDPurchaseName, + Provider: "gcp", + Product: "computeengine", + Action: "cud_purchase", + Description: gcpComputeEngineCUDPurchaseDescription, + RealPurchaseEnabled: true, + ExamplePrompts: []string{ + "Preview a 3-year CUD for 8 vCPUs and 32 GB memory in us-central1", + "Buy a 1-year Compute Engine CUD for real: 4 vCPUs, 16 GB memory", + }, + } +} + +func (t *gcpComputeEngineCUDPurchaseTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[gcpComputeEngineCUDPurchaseArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{int(TermOneYear), int(TermThreeYear)}}, + "dry_run": {Default: true}, + "confirm": {Default: false}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: gcpComputeEngineCUDPurchaseName, + Description: gcpComputeEngineCUDPurchaseDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *gcpComputeEngineCUDPurchaseTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args gcpComputeEngineCUDPurchaseArgs) (*mcp.CallToolResult, PurchaseResponse, error) { + rec, region, dryRun, confirm, err := gcpComputeEngineRecommendationFromArgs(args) + if err != nil { + return nil, PurchaseResponse{}, err + } + + resp, err := ExecutePurchase(ctx, PurchaseRequest{ + Region: region, + Recommendation: rec, + DryRun: dryRun, + Confirm: confirm, + ResolveClient: t.resolveClient(args, region), + Nonce: args.IdempotencyNonce, + }) + if err != nil { + return nil, PurchaseResponse{}, err + } + return nil, *resp, nil +} + +// gcpComputeEngineRecommendationFromArgs validates args and builds the +// common.Recommendation to purchase, the effective region (trimmed of any +// surrounding whitespace), and the effective dry_run/confirm booleans. +// Details is set as a value (common.ComputeDetails{}), not a pointer, to +// match the value type assertion in +// providers/gcp/services/computeengine/client.go's memoryMBFromDetails. +func gcpComputeEngineRecommendationFromArgs(args gcpComputeEngineCUDPurchaseArgs) (rec common.Recommendation, region string, dryRun, confirm bool, err error) { + region, err = requireNonBlank("region", args.Region) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + machineType, err := requireNonBlank("machine_type", args.MachineType) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + if args.VCPUCount <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("vcpu_count must be > 0, got %d", args.VCPUCount) + } + if args.MemoryGB <= 0 { + return common.Recommendation{}, "", false, false, fmt.Errorf("memory_gb must be > 0, got %v", args.MemoryGB) + } + term, err := ValidateTermYears(args.TermYears) + if err != nil { + return common.Recommendation{}, "", false, false, err + } + + rec = common.Recommendation{ + Provider: common.ProviderGCP, + Service: common.ServiceCompute, + Region: region, + ResourceType: machineType, + Count: args.VCPUCount, + CommitmentType: common.CommitmentCUD, + Term: term.RecommendationTerm(), + Details: common.ComputeDetails{ + InstanceType: machineType, + MemoryGB: args.MemoryGB, + }, + } + + dryRun, confirm = true, false + if args.DryRun != nil { + dryRun = *args.DryRun + } + if args.Confirm != nil { + confirm = *args.Confirm + } + return rec, region, dryRun, confirm, nil +} + +// resolveClient returns the ResolveClientFunc that ExecutePurchase invokes +// only for a real purchase. region is the effective, already-validated-and- +// trimmed region returned by gcpComputeEngineRecommendationFromArgs -- not +// args.Region -- so a real purchase never resolves the provider/service +// client against a raw, un-trimmed value. +func (t *gcpComputeEngineCUDPurchaseTool) resolveClient(args gcpComputeEngineCUDPurchaseArgs, region string) ResolveClientFunc { + return func(ctx context.Context) (provider.ServiceClient, error) { + cfg := &provider.ProviderConfig{Name: string(common.ProviderGCP), GCPProjectID: args.GCPProjectID, Region: region} + prov, err := t.createProvider(string(common.ProviderGCP), cfg) + if err != nil { + return nil, err + } + return prov.GetServiceClient(ctx, common.ServiceCompute, region) + } +} diff --git a/mcp/tools/gcp_computeengine_cud_test.go b/mcp/tools/gcp_computeengine_cud_test.go new file mode 100644 index 000000000..9234ce0e5 --- /dev/null +++ b/mcp/tools/gcp_computeengine_cud_test.go @@ -0,0 +1,156 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +func validGCPCUDArgs() gcpComputeEngineCUDPurchaseArgs { + return gcpComputeEngineCUDPurchaseArgs{ + Region: "us-central1", + MachineType: "n2-standard-4", + VCPUCount: 4, + MemoryGB: 16, + TermYears: 1, + } +} + +func TestGCPComputeEngineRecommendationFromArgs(t *testing.T) { + t.Parallel() + rec, region, dryRun, confirm, err := gcpComputeEngineRecommendationFromArgs(validGCPCUDArgs()) + require.NoError(t, err) + assert.Equal(t, "us-central1", region) + assert.True(t, dryRun) + assert.False(t, confirm) + assert.Equal(t, common.ProviderGCP, rec.Provider) + assert.Equal(t, common.ServiceCompute, rec.Service) + assert.Equal(t, common.CommitmentCUD, rec.CommitmentType) + assert.Equal(t, "n2-standard-4", rec.ResourceType) + assert.Equal(t, 4, rec.Count) + assert.Equal(t, "1yr", rec.Term) + + // Details MUST be a value common.ComputeDetails, not a pointer: + // providers/gcp/services/computeengine/client.go's memoryMBFromDetails + // type-asserts rec.Details.(common.ComputeDetails), unlike every AWS + // Details assertion which expects a pointer. + details, ok := rec.Details.(common.ComputeDetails) + require.True(t, ok, "Details must be a value common.ComputeDetails, not *common.ComputeDetails") + assert.InDelta(t, 16.0, details.MemoryGB, 0.001) +} + +func TestGCPComputeEngineRecommendationFromArgsInvalid(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(*gcpComputeEngineCUDPurchaseArgs) + errSub string + }{ + {"missing region", func(a *gcpComputeEngineCUDPurchaseArgs) { a.Region = "" }, "region is required"}, + {"whitespace-only region", func(a *gcpComputeEngineCUDPurchaseArgs) { a.Region = " " }, "region is required"}, + {"missing machine_type", func(a *gcpComputeEngineCUDPurchaseArgs) { a.MachineType = "" }, "machine_type is required"}, + {"whitespace-only machine_type", func(a *gcpComputeEngineCUDPurchaseArgs) { a.MachineType = "\t " }, "machine_type is required"}, + {"zero vcpu_count", func(a *gcpComputeEngineCUDPurchaseArgs) { a.VCPUCount = 0 }, "vcpu_count must be"}, + {"zero memory_gb", func(a *gcpComputeEngineCUDPurchaseArgs) { a.MemoryGB = 0 }, "memory_gb must be"}, + {"negative memory_gb", func(a *gcpComputeEngineCUDPurchaseArgs) { a.MemoryGB = -1 }, "memory_gb must be"}, + {"invalid term", func(a *gcpComputeEngineCUDPurchaseArgs) { a.TermYears = 2 }, "invalid term_years"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + args := validGCPCUDArgs() + tc.mutate(&args) + _, _, _, _, err := gcpComputeEngineRecommendationFromArgs(args) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.errSub) + }) + } +} + +// TestGCPComputeEngineRecommendationFromArgsTrimsSurroundingWhitespace is the +// regression guard for the CodeRabbit finding: requireNonBlank rejected an +// all-whitespace value but let a value with surrounding whitespace (e.g. +// " us-central1 ") pass through unchanged into rec.Region/rec.ResourceType and +// the returned region (which resolveClient uses for ProviderConfig.Region and +// GetServiceClient). +func TestGCPComputeEngineRecommendationFromArgsTrimsSurroundingWhitespace(t *testing.T) { + t.Parallel() + args := validGCPCUDArgs() + args.Region = " us-central1 " + args.MachineType = " n2-standard-4 " + + rec, region, _, _, err := gcpComputeEngineRecommendationFromArgs(args) + require.NoError(t, err) + assert.Equal(t, "us-central1", region, "returned region must be trimmed") + assert.Equal(t, "us-central1", rec.Region, "rec.Region must be trimmed") + assert.Equal(t, "n2-standard-4", rec.ResourceType, "rec.ResourceType must be trimmed") + details, ok := rec.Details.(common.ComputeDetails) + require.True(t, ok) + assert.Equal(t, "n2-standard-4", details.InstanceType, "Details.InstanceType must be trimmed") +} + +func TestGCPComputeEngineCUDPurchaseHandleConfirmFalseRefuses(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &gcpComputeEngineCUDPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validGCPCUDArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(false) + + _, _, err := tool.handle(context.Background(), nil, args) + require.Error(t, err) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +func TestGCPComputeEngineCUDPurchaseHandleDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + tool := &gcpComputeEngineCUDPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + resolveCalled = true + return nil, nil + }, + } + args := validGCPCUDArgs() + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.False(t, resolveCalled) + assert.True(t, resp.DryRun) +} + +func TestGCPComputeEngineCUDPurchaseHandleRealPurchase(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "cud-1"}} + var gotService common.ServiceType + tool := &gcpComputeEngineCUDPurchaseTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return &recordingProvider{ + fakeProvider: &fakeProvider{name: "gcp"}, + client: fake, + gotService: &gotService, + gotRegion: new(string), + }, nil + }, + } + args := validGCPCUDArgs() + args.DryRun = boolPtr(false) + args.Confirm = boolPtr(true) + + _, resp, err := tool.handle(context.Background(), nil, args) + require.NoError(t, err) + assert.True(t, resp.Success) + assert.Equal(t, common.ServiceCompute, gotService) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) +} diff --git a/mcp/tools/list_commitment_actions.go b/mcp/tools/list_commitment_actions.go new file mode 100644 index 000000000..a9743124e --- /dev/null +++ b/mcp/tools/list_commitment_actions.go @@ -0,0 +1,94 @@ +package tools + +import ( + "context" + + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +const listCommitmentActionsName = "cudly_list_commitment_actions" + +const listCommitmentActionsDescription = "List every CUDly commitment-purchase and search tool available on this " + + "MCP server, including which ones can execute a REAL purchase (money-affecting) versus which are " + + "search/preview-only, plus example prompts for each. This tool never spends money and takes no parameters -- " + + "start here if you don't already know which cudly_* tool you need." + +// listCommitmentActionsArgs is empty: the tool takes no parameters. +type listCommitmentActionsArgs struct{} + +// ActionEntry is one catalog entry returned by cudly_list_commitment_actions, +// reshaping a Descriptor for JSON output. +type ActionEntry struct { + Name string `json:"name"` + Provider string `json:"provider,omitempty"` + Product string `json:"product,omitempty"` + Action string `json:"action,omitempty"` + Description string `json:"description"` + RealPurchaseEnabled bool `json:"real_purchase_enabled"` + ExamplePrompts []string `json:"example_prompts,omitempty"` +} + +// listCommitmentActionsResult is the tool's structured output. +type listCommitmentActionsResult struct { + Actions []ActionEntry `json:"actions"` +} + +type listCommitmentActionsTool struct { + descriptors []Descriptor +} + +// NewListCommitmentActions builds the cudly_list_commitment_actions tool +// from descriptors -- the same slice of Descriptor values mcp/server.go +// collects from every other tool's Descriptor() method, so this catalog is +// generated from the live registry rather than hand-duplicated in code or +// docs. +func NewListCommitmentActions(descriptors []Descriptor) Registration { + return &listCommitmentActionsTool{descriptors: descriptors} +} + +// ListCommitmentActionsDescriptor returns the static Descriptor for +// cudly_list_commitment_actions itself. It is exported so mcp/server.go can +// include this tool in its own catalog: the descriptors slice passed to +// NewListCommitmentActions must be assembled (and thus known) before the +// tool exists, so its own entry can't come from calling Descriptor() on an +// already-constructed instance the way every other tool's entry does. +func ListCommitmentActionsDescriptor() Descriptor { + return Descriptor{ + Name: listCommitmentActionsName, + Description: listCommitmentActionsDescription, + ExamplePrompts: []string{ + "What CUDly tools are available?", + "Which purchase tools can spend real money right now?", + "How do I buy AWS EC2 Reserved Instances through CUDly?", + }, + } +} + +func (t *listCommitmentActionsTool) Descriptor() Descriptor { + return ListCommitmentActionsDescriptor() +} + +func (t *listCommitmentActionsTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[listCommitmentActionsArgs](nil) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: listCommitmentActionsName, + Description: listCommitmentActionsDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *listCommitmentActionsTool) handle(_ context.Context, _ *mcp.CallToolRequest, _ listCommitmentActionsArgs) (*mcp.CallToolResult, listCommitmentActionsResult, error) { + actions := make([]ActionEntry, 0, len(t.descriptors)) + for _, d := range t.descriptors { + // ActionEntry's fields are identical in name, type, and order to + // Descriptor's -- only the json tags differ -- so a direct + // conversion is equivalent to (and clearer than) a field-by-field + // struct literal. + actions = append(actions, ActionEntry(d)) + } + return nil, listCommitmentActionsResult{Actions: actions}, nil +} diff --git a/mcp/tools/list_commitment_actions_test.go b/mcp/tools/list_commitment_actions_test.go new file mode 100644 index 000000000..e5b1cf891 --- /dev/null +++ b/mcp/tools/list_commitment_actions_test.go @@ -0,0 +1,48 @@ +package tools + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestListCommitmentActionsReturnsCatalog(t *testing.T) { + t.Parallel() + descriptors := []Descriptor{ + { + Name: "cudly_aws_ec2_ri_purchase", + Provider: "aws", + Product: "ec2", + Action: "ri_purchase", + Description: "spends real money. dry_run recommended first.", + RealPurchaseEnabled: true, + ExamplePrompts: []string{"buy 3 m5.large RIs in us-east-1"}, + }, + { + Name: "cudly_search_recommendations", + Description: "read-only search, never spends money.", + }, + } + + tool := NewListCommitmentActions(descriptors) + impl, ok := tool.(*listCommitmentActionsTool) + require.True(t, ok) + + _, result, err := impl.handle(context.Background(), nil, listCommitmentActionsArgs{}) + require.NoError(t, err) + require.Len(t, result.Actions, 2) + assert.Equal(t, "cudly_aws_ec2_ri_purchase", result.Actions[0].Name) + assert.True(t, result.Actions[0].RealPurchaseEnabled) + assert.False(t, result.Actions[1].RealPurchaseEnabled) +} + +func TestListCommitmentActionsDescriptorItself(t *testing.T) { + t.Parallel() + tool := NewListCommitmentActions(nil) + d := tool.Descriptor() + assert.Equal(t, "cudly_list_commitment_actions", d.Name) + assert.NotEmpty(t, d.Description) + assert.NotEmpty(t, d.ExamplePrompts) +} diff --git a/mcp/tools/purchase.go b/mcp/tools/purchase.go new file mode 100644 index 000000000..2e6b8ccdb --- /dev/null +++ b/mcp/tools/purchase.go @@ -0,0 +1,288 @@ +package tools + +import ( + "context" + "fmt" + "strconv" + "strings" + "time" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +// purchaseMode is the outcome of the dry_run/confirm safety gate: either a +// local preview (no provider call) or a real execution (provider call with +// PurchaseSourceMCP + an idempotency token). There is deliberately no third +// "no-op" outcome -- an ambiguous combination of flags is always an error, +// never a silent do-nothing (feedback_no_silent_fallbacks). +type purchaseMode int + +const ( + modePreview purchaseMode = iota + modeExecute +) + +// decidePurchaseMode applies the safety rail from the design doc (§7): a +// real purchase requires confirm=true AND dry_run=false. dry_run=true always +// wins and returns a preview, regardless of confirm, so a caller previewing +// a purchase can leave confirm at its default. The only refusal case is +// dry_run=false with confirm=false: the caller asked for a real purchase +// but did not confirm it, which must surface as an explicit error rather +// than silently downgrading to a preview or silently doing nothing. +func decidePurchaseMode(dryRun, confirm bool) (purchaseMode, error) { + if dryRun { + return modePreview, nil + } + if confirm { + return modeExecute, nil + } + return 0, fmt.Errorf("refusing real purchase: dry_run=false requires confirm=true (got confirm=false); " + + "set dry_run=true to preview this purchase instead, or confirm=true to execute it") +} + +// ResolveClientFunc lazily resolves the provider.ServiceClient that will +// receive the real PurchaseCommitment call. It is a func, not an +// already-resolved client, so ExecutePurchase can prove (and tests can +// assert) that a preview never triggers provider/credential resolution -- +// only modeExecute invokes it. +type ResolveClientFunc func(ctx context.Context) (provider.ServiceClient, error) + +// PurchaseRequest is the provider-agnostic input to ExecutePurchase. Each +// per-service tool handler builds one after validating its own typed +// parameters and constructing the common.Recommendation. +type PurchaseRequest struct { + Region string + Recommendation common.Recommendation + DryRun bool + Confirm bool + ResolveClient ResolveClientFunc + + // Nonce is optional. When non-empty, this call is treated as a + // DISTINCT purchase from an otherwise-identical one (authorizes a + // deliberate repeat, e.g. "buy 3 more RIs" on top of an earlier "buy 3 + // RIs" with the same parameters). When empty (the default), identical + // purchases dedupe so retries never double-buy. See idempotencyKeyFor. + Nonce string +} + +// PurchaseResponse is the structured result returned to the MCP caller for +// both preview and real-purchase outcomes. Error is a string (not the Go +// error) because it crosses the MCP JSON-RPC boundary as tool output, not a +// protocol-level error -- ExecutePurchase itself still returns a Go error +// for gate refusals and provider-call failures. +// +// Cost/OnDemandCost/EstimatedSavings/SavingsPercentage are pointers with +// omitempty: none of the *FromArgs constructors in this package populate +// Recommendation's cost fields (they build a fresh Recommendation from the +// caller's typed args, not from a priced search result), and some provider +// clients (e.g. AWS EC2 RIs, Savings Plans) never populate +// PurchaseResult.Cost either. A plain float64 could not distinguish "not +// known" from "genuinely $0", so every response reported 0 for money fields +// it never actually priced. A pointer that's nil (and omitted from the JSON +// payload entirely) when no real value exists lets a caller tell "unknown" +// apart from "confirmed zero" (feedback_nullable_not_zero). +type PurchaseResponse struct { + Success bool `json:"success"` + DryRun bool `json:"dry_run"` + CommitmentID string `json:"commitment_id,omitempty"` + Cost *float64 `json:"cost,omitempty"` + OnDemandCost *float64 `json:"on_demand_cost,omitempty"` + EstimatedSavings *float64 `json:"estimated_savings,omitempty"` + SavingsPercentage *float64 `json:"savings_percentage,omitempty"` + EffectiveDate string `json:"effective_date,omitempty"` + TermYears int `json:"term_years,omitempty"` + Error string `json:"error,omitempty"` +} + +// termYearsFromRecommendationTerm extracts the integer commitment length in +// years from a Recommendation.Term string in the "yr" format every +// *FromArgs constructor in this package writes via +// TermYears.RecommendationTerm() (enums.go). Returns 0 when term does not +// match that format (e.g. an empty Term), so PurchaseResponse.TermYears is +// simply omitted (it has `omitempty`) rather than reporting a fabricated +// value. +func termYearsFromRecommendationTerm(term string) int { + years, err := strconv.Atoi(strings.TrimSuffix(term, "yr")) + if err != nil { + return 0 + } + return years +} + +// nonZeroCostPtr returns a pointer to v, or nil when v is exactly zero. Cost +// and savings fields on common.Recommendation and common.PurchaseResult are +// plain (unpointered) float64s that upstream code sometimes never populates +// (see the PurchaseResponse doc comment above); this treats an unpopulated +// zero as "unknown" rather than fabricating a real $0 figure the caller +// never priced. +func nonZeroCostPtr(v float64) *float64 { + if v == 0 { + return nil + } + return &v +} + +// idempotencyKeyFor derives a stable per-request key from every field that +// identifies what is being bought: provider, region, service, resource type, +// count, term, payment option, plus every service-specific dimension held in +// rec.Details (see detailsKeyComponent) -- platform/tenancy/scope for EC2 +// RIs, engine/az_config for RDS RIs, engine for ElastiCache RIs, hourly +// commitment/instance family for Savings Plans, memory for GCP CUDs. A +// caller re-driving the exact same tool call (e.g. after a network timeout) +// reuses the same common.DeriveIdempotencyToken output and the provider +// dedupes the retry instead of double-purchasing; a request that differs in +// ANY price- or identity-affecting dimension derives a different key instead +// of silently colliding with an unrelated purchase (issue found in review: +// a $5/hr and $50/hr Compute Savings Plan previously shared a token because +// only HourlyCommitment differed and Details was never consulted). This is a +// request-scoped substitute for the purchase_executions row that the CLI/web +// paths use as their idempotency anchor (pkg/common/tokens.go) -- the MCP +// server has no such row, so the request's own identifying fields play that +// role. +// +// rec.Account is deliberately excluded: no *FromArgs constructor in this +// package populates it today, so folding it in would add an always-empty, +// misleading key component rather than real discrimination. +// +// This function is deliberately fail-safe with respect to time: it folds in +// no clock reading of any kind. When nonce is empty (the default), two calls +// with identical dimensions ALWAYS derive the same key, no matter how far +// apart in time they happen -- so a retry that straddles any time boundary +// still dedupes at the provider instead of risking a double purchase. An +// earlier version of this function instead folded in an automatic hourly +// time bucket to distinguish "buy 3 RIs now" from a genuinely separate "buy +// 3 more next week" with identical parameters; that inverted the safety +// direction of this money path, because a retry that happened to straddle +// an hour boundary (e.g. a slow request issued at 12:59:58 retried at +// 13:00:02) derived a different key and could double-buy. The worst case of +// the current, fail-safe default is a skipped intentional repeat -- a +// caller who genuinely wants a second, identical purchase must say so +// explicitly. nonce is that explicit opt-in: when the caller supplies a +// non-empty nonce, it is folded into the key so an otherwise-identical +// purchase becomes a distinct one (e.g. "buy 3 now" then "buy 3 more next +// week" by passing a fresh nonce on the second call); the same nonce plus +// the same dimensions still dedupes a nonce'd retry. +func idempotencyKeyFor(region string, rec common.Recommendation, nonce string) string { + return fmt.Sprintf("mcp:%s:%s:%s:%s:%d:%s:%s:%s:%s", + rec.Provider, region, rec.Service, rec.ResourceType, rec.Count, rec.Term, rec.PaymentOption, + detailsKeyComponent(rec.Details), nonce) +} + +// detailsKeyComponent returns a canonical, deterministic encoding of every +// field in rec.Details that the purchase tools in this package populate, so +// idempotencyKeyFor can fold service-specific price-affecting dimensions +// into the token. Each case lists every field of its concrete Details type +// explicitly (not a hand-picked subset) so a field added to one of these +// types later shows up here as a visible diff rather than a silent key gap. +// Returns "" for nil or an unrecognized Details (e.g. +// azure_compute_ri.go's tool, whose recommendation carries no Details at +// all). +func detailsKeyComponent(details common.ServiceDetails) string { + switch d := details.(type) { + case *common.ComputeDetails: + return computeDetailsKey(d) + case common.ComputeDetails: + return computeDetailsKey(&d) + case *common.DatabaseDetails: + return databaseDetailsKey(d) + case *common.CacheDetails: + return cacheDetailsKey(d) + case *common.SavingsPlanDetails: + return savingsPlanDetailsKey(d) + default: + return "" + } +} + +func computeDetailsKey(d *common.ComputeDetails) string { + if d == nil { + return "" + } + return fmt.Sprintf("instance_type=%s;platform=%s;tenancy=%s;scope=%s;vcpu=%d;memory_gb=%g", + d.InstanceType, d.Platform, d.Tenancy, d.Scope, d.VCPU, d.MemoryGB) +} + +func databaseDetailsKey(d *common.DatabaseDetails) string { + if d == nil { + return "" + } + return fmt.Sprintf("engine=%s;engine_version=%s;az_config=%s;instance_class=%s;deployment=%s", + d.Engine, d.EngineVersion, d.AZConfig, d.InstanceClass, d.Deployment) +} + +func cacheDetailsKey(d *common.CacheDetails) string { + if d == nil { + return "" + } + return fmt.Sprintf("engine=%s;node_type=%s;shards=%d", d.Engine, d.NodeType, d.Shards) +} + +func savingsPlanDetailsKey(d *common.SavingsPlanDetails) string { + if d == nil { + return "" + } + return fmt.Sprintf("plan_type=%s;hourly_commitment=%g;coverage=%s;instance_family=%s;region=%s;offering_id=%s", + d.PlanType, d.HourlyCommitment, d.Coverage, d.InstanceFamily, d.Region, d.OfferingID) +} + +// ExecutePurchase runs the shared dry_run/confirm safety gate and, for a +// real purchase, resolves the service client and calls PurchaseCommitment +// with PurchaseSourceMCP and a derived idempotency token. It never calls +// ResolveClient in preview mode, so a preview makes zero provider/SDK calls. +func ExecutePurchase(ctx context.Context, req PurchaseRequest) (*PurchaseResponse, error) { + mode, err := decidePurchaseMode(req.DryRun, req.Confirm) + if err != nil { + return nil, err + } + + rec := req.Recommendation + if mode == modePreview { + return &PurchaseResponse{ + Success: true, + DryRun: true, + Cost: nonZeroCostPtr(rec.CommitmentCost), + OnDemandCost: nonZeroCostPtr(rec.OnDemandCost), + EstimatedSavings: nonZeroCostPtr(rec.EstimatedSavings), + SavingsPercentage: nonZeroCostPtr(rec.SavingsPercentage), + TermYears: termYearsFromRecommendationTerm(rec.Term), + }, nil + } + + if req.ResolveClient == nil { + return nil, fmt.Errorf("internal error: no ResolveClient configured for real purchase") + } + client, err := req.ResolveClient(ctx) + if err != nil { + return nil, fmt.Errorf("resolve %s service client: %w", rec.Provider, err) + } + + token := common.DeriveIdempotencyToken(idempotencyKeyFor(req.Region, rec, req.Nonce), 0) + opts := common.PurchaseOptions{ + Source: common.PurchaseSourceMCP, + IdempotencyToken: token, + } + + result, err := client.PurchaseCommitment(ctx, rec, opts) + if err != nil { + // Full provider error text surfaces to the caller (feedback: + // providers must never swallow the underlying SDK/HTTP error). + return nil, fmt.Errorf("purchase commitment failed: %w", err) + } + + resp := &PurchaseResponse{ + Success: result.Success, + DryRun: result.DryRun, + CommitmentID: result.CommitmentID, + Cost: nonZeroCostPtr(result.Cost), + OnDemandCost: nonZeroCostPtr(rec.OnDemandCost), + EstimatedSavings: nonZeroCostPtr(rec.EstimatedSavings), + SavingsPercentage: nonZeroCostPtr(rec.SavingsPercentage), + EffectiveDate: result.Timestamp.Format(time.RFC3339), + TermYears: termYearsFromRecommendationTerm(rec.Term), + } + if result.Error != nil { + resp.Error = result.Error.Error() + } + return resp, nil +} diff --git a/mcp/tools/purchase_test.go b/mcp/tools/purchase_test.go new file mode 100644 index 000000000..282fd5b48 --- /dev/null +++ b/mcp/tools/purchase_test.go @@ -0,0 +1,489 @@ +package tools + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +// fakeServiceClient is a minimal provider.ServiceClient test double. Only +// PurchaseCommitment is exercised by these tests; the rest of the interface +// is implemented trivially to satisfy the type. +type fakeServiceClient struct { + purchaseCalls int + purchaseResult common.PurchaseResult + purchaseErr error + lastOpts common.PurchaseOptions +} + +func (f *fakeServiceClient) GetServiceType() common.ServiceType { return common.ServiceEC2 } +func (f *fakeServiceClient) GetRegion() string { return "us-east-1" } +func (f *fakeServiceClient) GetRecommendations(_ context.Context, _ *common.RecommendationParams) ([]common.Recommendation, error) { + return nil, nil +} +func (f *fakeServiceClient) GetExistingCommitments(_ context.Context) ([]common.Commitment, error) { + return nil, nil +} +func (f *fakeServiceClient) PurchaseCommitment(_ context.Context, rec common.Recommendation, opts common.PurchaseOptions) (common.PurchaseResult, error) { + f.purchaseCalls++ + f.lastOpts = opts + f.purchaseResult.Recommendation = rec + return f.purchaseResult, f.purchaseErr +} +func (f *fakeServiceClient) ValidateOffering(_ context.Context, _ common.Recommendation) error { + return nil +} +func (f *fakeServiceClient) GetOfferingDetails(_ context.Context, _ common.Recommendation) (*common.OfferingDetails, error) { + return nil, nil +} +func (f *fakeServiceClient) GetValidResourceTypes(_ context.Context) ([]string, error) { + return nil, nil +} + +var _ provider.ServiceClient = (*fakeServiceClient)(nil) + +// testRecommendation mirrors what a real purchase tool's *FromArgs +// constructor actually builds: none of them populate +// OnDemandCost/CommitmentCost/EstimatedSavings/SavingsPercentage (they build +// a fresh Recommendation from the caller's typed args, not from a priced +// search result), so this fixture leaves those fields at their zero value +// too. An earlier version of this fixture hand-set those fields, which +// masked the all-responses-report-0 finding from review -- see +// TestExecutePurchasePreviewOmitsUnknownCostFields. +func testRecommendation() common.Recommendation { + return common.Recommendation{ + Provider: common.ProviderAWS, + Account: "123456789012", + Service: common.ServiceEC2, + Region: "us-east-1", + ResourceType: "m5.large", + Count: 3, + Term: "3yr", + PaymentOption: "no-upfront", + } +} + +// testRecommendationWithCost extends testRecommendation with real cost +// figures, used only to prove ExecutePurchase passes a genuinely-known cost +// through to the response when one is present. +func testRecommendationWithCost() common.Recommendation { + rec := testRecommendation() + rec.OnDemandCost = 1000 + rec.CommitmentCost = 600 + rec.EstimatedSavings = 400 + rec.SavingsPercentage = 40 + return rec +} + +func TestDecidePurchaseMode(t *testing.T) { + t.Parallel() + cases := []struct { + name string + dryRun bool + confirm bool + want purchaseMode + wantErr bool + }{ + {"dry run wins regardless of confirm", true, false, modePreview, false}, + {"dry run with confirm still previews", true, true, modePreview, false}, + {"confirmed real purchase executes", false, true, modeExecute, false}, + {"unconfirmed real purchase refused", false, false, 0, true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := decidePurchaseMode(tc.dryRun, tc.confirm) + if tc.wantErr { + require.Error(t, err) + return + } + require.NoError(t, err) + assert.Equal(t, tc.want, got) + }) + } +} + +// TestExecutePurchaseDryRunNeverCallsProvider proves the safety rail from +// the design doc: dry_run=true must never invoke ResolveClient (and +// therefore never PurchaseCommitment), even when confirm=true. ResolveClient +// here returns an error if called at all, so any invocation fails the test. +func TestExecutePurchaseDryRunNeverCallsProvider(t *testing.T) { + t.Parallel() + resolveCalled := false + resolve := func(_ context.Context) (provider.ServiceClient, error) { + resolveCalled = true + return nil, errors.New("ResolveClient must not be called in dry_run mode") + } + + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendationWithCost(), + DryRun: true, + Confirm: true, + ResolveClient: resolve, + }) + + require.NoError(t, err) + require.NotNil(t, resp) + assert.False(t, resolveCalled, "dry_run=true must never resolve a service client") + assert.True(t, resp.DryRun) + assert.True(t, resp.Success) + require.NotNil(t, resp.Cost, "a genuinely-known cost must be passed through, not dropped") + assert.Equal(t, 600.0, *resp.Cost) + require.NotNil(t, resp.OnDemandCost) + assert.Equal(t, 1000.0, *resp.OnDemandCost) + require.NotNil(t, resp.EstimatedSavings) + assert.Equal(t, 400.0, *resp.EstimatedSavings) + require.NotNil(t, resp.SavingsPercentage) + assert.Equal(t, 40.0, *resp.SavingsPercentage) +} + +// TestExecutePurchasePreviewOmitsUnknownCostFields proves finding 2 of the +// adversarial review: a dry-run preview built from a Recommendation that +// mirrors what real purchase tools actually construct (no cost fields set, +// since no *FromArgs constructor in this package populates them) must not +// report cost/on_demand_cost/estimated_savings/savings_percentage as a real +// 0 -- that would be indistinguishable from a confirmed $0 purchase. The +// pointer fields must be nil, and therefore omitted from the JSON payload +// entirely rather than serialized as 0. +func TestExecutePurchasePreviewOmitsUnknownCostFields(t *testing.T) { + t.Parallel() + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: true, + Confirm: false, + }) + require.NoError(t, err) + require.NotNil(t, resp) + + assert.Nil(t, resp.Cost) + assert.Nil(t, resp.OnDemandCost) + assert.Nil(t, resp.EstimatedSavings) + assert.Nil(t, resp.SavingsPercentage) + + raw, err := json.Marshal(resp) + require.NoError(t, err) + body := string(raw) + assert.NotContains(t, body, `"cost"`, "unknown cost must be omitted from the JSON payload, not reported as 0") + assert.NotContains(t, body, `"on_demand_cost"`) + assert.NotContains(t, body, `"estimated_savings"`) + assert.NotContains(t, body, `"savings_percentage"`) +} + +// TestExecutePurchasePreviewPopulatesTermYears is the regression guard for +// the CodeRabbit finding that PurchaseResponse.TermYears was declared in the +// JSON contract but never set in either ExecutePurchase branch, so it was +// always zero/omitted even though the term is known from the recommendation. +// testRecommendation() carries Term: "3yr", the same "yr" format every +// *FromArgs constructor in this package writes. +func TestExecutePurchasePreviewPopulatesTermYears(t *testing.T) { + t.Parallel() + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: true, + Confirm: false, + }) + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, 3, resp.TermYears, "a preview response must carry the term the caller specified") +} + +// TestExecutePurchaseRealPurchasePopulatesTermYears is the real-purchase +// counterpart of TestExecutePurchasePreviewPopulatesTermYears: the term must +// be populated on the modeExecute branch too, not only the preview branch. +func TestExecutePurchaseRealPurchasePopulatesTermYears(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{ + purchaseResult: common.PurchaseResult{Success: true, CommitmentID: "ri-term-test"}, + } + + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: false, + Confirm: true, + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake, nil }, + }) + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, 3, resp.TermYears, "a real-purchase response must carry the term the caller specified") +} + +// TestExecutePurchaseUnconfirmedRealPurchaseRefused proves confirm=false +// refuses a real purchase (dry_run=false) with a structured error rather +// than a silent no-op, and that ResolveClient is never invoked either. +func TestExecutePurchaseUnconfirmedRealPurchaseRefused(t *testing.T) { + t.Parallel() + resolveCalled := false + resolve := func(_ context.Context) (provider.ServiceClient, error) { + resolveCalled = true + return nil, nil + } + + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: false, + Confirm: false, + ResolveClient: resolve, + }) + + require.Error(t, err) + assert.Nil(t, resp) + assert.False(t, resolveCalled) + assert.Contains(t, err.Error(), "confirm=true") +} + +// TestExecutePurchaseRealPurchaseCallsProviderWithMCPSource proves a +// confirmed real purchase resolves the client, calls PurchaseCommitment +// exactly once, and stamps PurchaseSourceMCP + a non-empty idempotency +// token -- never a caller-suppliable source string. +func TestExecutePurchaseRealPurchaseCallsProviderWithMCPSource(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{ + purchaseResult: common.PurchaseResult{ + Success: true, + CommitmentID: "ri-12345", + Cost: 600, + }, + } + resolve := func(_ context.Context) (provider.ServiceClient, error) { + return fake, nil + } + + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: false, + Confirm: true, + ResolveClient: resolve, + }) + + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, 1, fake.purchaseCalls) + assert.Equal(t, common.PurchaseSourceMCP, fake.lastOpts.Source) + assert.NotEmpty(t, fake.lastOpts.IdempotencyToken) + assert.True(t, resp.Success) + assert.Equal(t, "ri-12345", resp.CommitmentID) + assert.False(t, resp.DryRun) +} + +// TestExecutePurchaseSameRequestDerivesSameToken proves idempotencyKeyFor +// (and therefore the derived token) is deterministic for the same +// identifying fields, so a retried call with identical arguments dedupes at +// the provider rather than double-purchasing. +func TestExecutePurchaseSameRequestDerivesSameToken(t *testing.T) { + t.Parallel() + rec := testRecommendation() + fake1 := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true}} + fake2 := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true}} + + _, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", Recommendation: rec, DryRun: false, Confirm: true, + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake1, nil }, + }) + require.NoError(t, err) + + _, err = ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", Recommendation: rec, DryRun: false, Confirm: true, + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake2, nil }, + }) + require.NoError(t, err) + + assert.Equal(t, fake1.lastOpts.IdempotencyToken, fake2.lastOpts.IdempotencyToken) + + // A materially different request (different count) must derive a + // different token. + rec2 := rec + rec2.Count = 4 + fake3 := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true}} + _, err = ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", Recommendation: rec2, DryRun: false, Confirm: true, + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake3, nil }, + }) + require.NoError(t, err) + assert.NotEqual(t, fake1.lastOpts.IdempotencyToken, fake3.lastOpts.IdempotencyToken) +} + +// TestExecutePurchaseProviderErrorSurfaced proves a provider-side purchase +// failure surfaces the full underlying error text rather than being +// swallowed. +func TestExecutePurchaseProviderErrorSurfaced(t *testing.T) { + t.Parallel() + fake := &fakeServiceClient{purchaseErr: errors.New("AWS API: InsufficientInstanceCapacity")} + resolve := func(_ context.Context) (provider.ServiceClient, error) { return fake, nil } + + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: false, + Confirm: true, + ResolveClient: resolve, + }) + + require.Error(t, err) + assert.Nil(t, resp) + assert.Contains(t, err.Error(), "InsufficientInstanceCapacity") +} + +// TestExecutePurchaseResolveClientErrorSurfaced proves a client-resolution +// failure (e.g. bad credentials) surfaces its error text too. +func TestExecutePurchaseResolveClientErrorSurfaced(t *testing.T) { + t.Parallel() + resolve := func(_ context.Context) (provider.ServiceClient, error) { + return nil, errors.New("no AWS credentials found") + } + + resp, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", + Recommendation: testRecommendation(), + DryRun: false, + Confirm: true, + ResolveClient: resolve, + }) + + require.Error(t, err) + assert.Nil(t, resp) + assert.Contains(t, err.Error(), "no AWS credentials found") +} + +// TestIdempotencyKeyDistinguishesSavingsPlanHourlyCommitment proves finding +// 1 of the adversarial review of the purchase feature: two Savings Plans +// requests that differ only in hourly_commitment (a $5/hr vs a $50/hr +// Compute Savings Plan) must derive different idempotency tokens. Before the +// fix, idempotencyKeyFor never consulted rec.Details at all, so these two +// materially different purchases collided on the same token and AWS would +// have silently deduped the second call as a "retry" of the first instead +// of buying a second, larger plan. +func TestIdempotencyKeyDistinguishesSavingsPlanHourlyCommitment(t *testing.T) { + t.Parallel() + + cheapArgs := validSavingsPlansArgs() + cheapArgs.HourlyCommitment = 5 + expensiveArgs := validSavingsPlansArgs() + expensiveArgs.HourlyCommitment = 50 + + cheapRec, region, _, _, err := savingsPlanRecommendationFromArgs(cheapArgs) + require.NoError(t, err) + expensiveRec, _, _, _, err := savingsPlanRecommendationFromArgs(expensiveArgs) + require.NoError(t, err) + + cheapKey := idempotencyKeyFor(region, cheapRec, "") + expensiveKey := idempotencyKeyFor(region, expensiveRec, "") + assert.NotEqual(t, cheapKey, expensiveKey, + "a $5/hr and a $50/hr Compute Savings Plan must not derive the same idempotency key") +} + +// TestIdempotencyKeyDistinguishesEC2Platform proves the second half of +// finding 1: an EC2 RI purchase for Linux vs Windows, with every other field +// (region/instance_type/count/term/payment_option) identical, must not +// collide on the same idempotency key -- Platform is a price- and +// product-affecting dimension carried in rec.Details, and the pre-fix key +// derivation ignored Details entirely. +func TestIdempotencyKeyDistinguishesEC2Platform(t *testing.T) { + t.Parallel() + + linuxArgs := validEC2Args() + linuxArgs.Platform = "Linux/UNIX" + windowsArgs := validEC2Args() + windowsArgs.Platform = "Windows" + + linuxRec, linuxRegion, _, _, err := ec2RecommendationFromArgs(linuxArgs) + require.NoError(t, err) + windowsRec, windowsRegion, _, _, err := ec2RecommendationFromArgs(windowsArgs) + require.NoError(t, err) + + linuxKey := idempotencyKeyFor(linuxRegion, linuxRec, "") + windowsKey := idempotencyKeyFor(windowsRegion, windowsRec, "") + assert.NotEqual(t, linuxKey, windowsKey, + "a Linux and a Windows EC2 RI purchase must not derive the same idempotency key") +} + +// TestIdempotencyKeySameDimensionsNoNonceAlwaysMatch is the regression guard +// for the fail-safe design: identical purchase dimensions with no nonce must +// ALWAYS derive the same key, with no dependence on time at all. This is the +// inverse of, and replaces, a prior design that folded an automatic hourly +// time bucket into the key when no nonce was supplied -- under that design a +// retry that happened to straddle an hour boundary (e.g. issued at +// 12:59:58, retried four seconds later at 13:00:02) derived a DIFFERENT key, +// so the provider could treat the retry as a brand new purchase instead of +// deduping it, resulting in a double purchase. idempotencyKeyFor no longer +// reads a clock at all when nonce is empty, so this is not merely "same +// bucket" but unconditionally the same key for the life of the process. +func TestIdempotencyKeySameDimensionsNoNonceAlwaysMatch(t *testing.T) { + t.Parallel() + rec := testRecommendation() + region := "us-east-1" + + key1 := idempotencyKeyFor(region, rec, "") + key2 := idempotencyKeyFor(region, rec, "") + assert.Equal(t, key1, key2, + "identical dimensions with no nonce must always derive the same key, so a retry never double-buys") +} + +// TestIdempotencyKeyNonceAuthorizesDistinctRepeat proves the nonce is the +// caller's explicit opt-in to a deliberate repeat purchase: a non-empty +// nonce derives a key different from the no-nonce key and from a different +// nonce, but the SAME nonce with the SAME dimensions still dedupes (a +// nonce'd retry is still safe against double-buying). +func TestIdempotencyKeyNonceAuthorizesDistinctRepeat(t *testing.T) { + t.Parallel() + rec := testRecommendation() + region := "us-east-1" + + noNonceKey := idempotencyKeyFor(region, rec, "") + nonceAKey1 := idempotencyKeyFor(region, rec, "nonce-a") + nonceAKey2 := idempotencyKeyFor(region, rec, "nonce-a") + nonceBKey := idempotencyKeyFor(region, rec, "nonce-b") + + assert.NotEqual(t, noNonceKey, nonceAKey1, + "supplying a nonce must authorize a purchase distinct from the no-nonce default") + assert.NotEqual(t, nonceAKey1, nonceBKey, + "two different nonces must derive two different keys") + assert.Equal(t, nonceAKey1, nonceAKey2, + "the same nonce with the same dimensions must still dedupe a nonce'd retry") +} + +// TestExecutePurchaseNonceThreadedThroughToToken proves PurchaseRequest.Nonce +// is actually wired end to end into ExecutePurchase's derived token, not +// just exercised at the idempotencyKeyFor level in isolation. +func TestExecutePurchaseNonceThreadedThroughToToken(t *testing.T) { + t.Parallel() + rec := testRecommendation() + + fake1 := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true}} + _, err := ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", Recommendation: rec, DryRun: false, Confirm: true, Nonce: "call-1", + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake1, nil }, + }) + require.NoError(t, err) + + fake2 := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true}} + _, err = ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", Recommendation: rec, DryRun: false, Confirm: true, Nonce: "call-2", + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake2, nil }, + }) + require.NoError(t, err) + + assert.NotEqual(t, fake1.lastOpts.IdempotencyToken, fake2.lastOpts.IdempotencyToken, + "different nonces must derive different idempotency tokens") + + fake3 := &fakeServiceClient{purchaseResult: common.PurchaseResult{Success: true}} + _, err = ExecutePurchase(context.Background(), PurchaseRequest{ + Region: "us-east-1", Recommendation: rec, DryRun: false, Confirm: true, Nonce: "call-1", + ResolveClient: func(_ context.Context) (provider.ServiceClient, error) { return fake3, nil }, + }) + require.NoError(t, err) + + assert.Equal(t, fake1.lastOpts.IdempotencyToken, fake3.lastOpts.IdempotencyToken, + "the same nonce must derive the same idempotency token") +} diff --git a/mcp/tools/registry.go b/mcp/tools/registry.go new file mode 100644 index 000000000..d661a1c0a --- /dev/null +++ b/mcp/tools/registry.go @@ -0,0 +1,41 @@ +package tools + +import "github.com/modelcontextprotocol/go-sdk/mcp" + +// Descriptor is the source-of-truth metadata for one MCP tool. Every tool +// file builds one and both mcp/server.go (to know what to register) and +// cudly_list_commitment_actions (to know what to advertise) read it, so the +// live tool set and the discoverability catalog can never drift apart -- +// there is exactly one place each tool's name/description/example prompts +// are written. +type Descriptor struct { + // Name is the MCP tool name, e.g. "cudly_aws_ec2_ri_purchase". + Name string + // Provider is "aws", "azure", "gcp", or "" for provider-agnostic + // meta-tools (cudly_list_commitment_actions, cudly_search_recommendations). + Provider string + // Product is the service the tool acts on, e.g. "ec2", "rds", "compute". + Product string + // Action is what the tool does, e.g. "ri_purchase", "cud_purchase", "search". + Action string + // Description is the tool's full MCP description, shared verbatim with + // the live mcp.Tool registration so the two can never disagree. + Description string + // RealPurchaseEnabled reports whether this tool can execute a real, + // money-spending purchase today (dry_run=false, confirm=true). false for + // read-only tools and for tools shipped dry-run-only pending a + // prerequisite fix (see the Azure/GCP tool comments). + RealPurchaseEnabled bool + // ExamplePrompts are 2-3 natural-language prompts that would plausibly + // invoke this tool, surfaced by cudly_list_commitment_actions so a + // session that doesn't know the tool name yet can find it. + ExamplePrompts []string +} + +// Registration is implemented by every tool file. Descriptor feeds the +// catalog; Register performs the live mcp.AddTool (or mcp.Server.AddTool) +// call that wires the tool's schema and handler onto the server. +type Registration interface { + Descriptor() Descriptor + Register(s *mcp.Server) error +} diff --git a/mcp/tools/schema.go b/mcp/tools/schema.go new file mode 100644 index 000000000..b6227b661 --- /dev/null +++ b/mcp/tools/schema.go @@ -0,0 +1,57 @@ +package tools + +import ( + "encoding/json" + "fmt" + + "github.com/google/jsonschema-go/jsonschema" +) + +// FieldOverride declares JSON Schema refinements -- an explicit enum +// membership and/or a documented default -- for one property of an +// otherwise auto-inferred schema. Centralizing this in one helper +// (BuildInputSchema) means every tool declares its enum/default once, next +// to its Go struct, instead of re-implementing schema post-processing per +// tool. +type FieldOverride struct { + // Enum, when non-empty, restricts the property to these exact values. + Enum []any + // Default, when non-nil, is recorded on the schema as the property's + // documented default so a caller inspecting the tool (or an MCP client + // that surfaces schema defaults in its UI) can see it without reading + // the tool description prose. It does NOT, by itself, cause the value to + // be applied when the caller omits the field -- each tool's handler + // applies its own default explicitly (see the dry_run/confirm pattern in + // purchase.go) so "omitted" is never silently confused with "false". + Default any +} + +// BuildInputSchema infers the JSON Schema for T via jsonschema.For, then +// applies the given per-field overrides by JSON field name. It returns an +// error -- rather than silently skipping -- when an override names a field +// that does not exist on T, so a typo in the override map is caught at +// server-startup / test time instead of quietly shipping an unconstrained +// schema for a money-affecting field. +func BuildInputSchema[T any](overrides map[string]FieldOverride) (*jsonschema.Schema, error) { + schema, err := jsonschema.For[T](nil) + if err != nil { + return nil, fmt.Errorf("infer schema for %T: %w", *new(T), err) + } + for field, ov := range overrides { + prop, ok := schema.Properties[field] + if !ok { + return nil, fmt.Errorf("schema override for unknown field %q (does the json tag match?)", field) + } + if len(ov.Enum) > 0 { + prop.Enum = ov.Enum + } + if ov.Default != nil { + b, err := json.Marshal(ov.Default) + if err != nil { + return nil, fmt.Errorf("marshal default for field %q: %w", field, err) + } + prop.Default = b + } + } + return schema, nil +} diff --git a/mcp/tools/schema_test.go b/mcp/tools/schema_test.go new file mode 100644 index 000000000..153acad68 --- /dev/null +++ b/mcp/tools/schema_test.go @@ -0,0 +1,58 @@ +package tools + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type schemaTestArgs struct { + Region string `json:"region" jsonschema:"AWS region"` + TermYears int `json:"term_years" jsonschema:"commitment term in years"` + PaymentOption string `json:"payment_option" jsonschema:"payment schedule"` + DryRun *bool `json:"dry_run,omitempty" jsonschema:"preview only, no purchase"` +} + +func TestBuildInputSchemaAppliesEnumAndDefault(t *testing.T) { + t.Parallel() + trueDefault := true + schema, err := BuildInputSchema[schemaTestArgs](map[string]FieldOverride{ + "term_years": {Enum: []any{1, 3}}, + "payment_option": {Enum: []any{"all-upfront", "partial-upfront", "no-upfront"}}, + "dry_run": {Default: trueDefault}, + }) + require.NoError(t, err) + + require.Contains(t, schema.Properties, "term_years") + assert.Equal(t, []any{1, 3}, schema.Properties["term_years"].Enum) + + require.Contains(t, schema.Properties, "payment_option") + assert.Equal(t, []any{"all-upfront", "partial-upfront", "no-upfront"}, schema.Properties["payment_option"].Enum) + + require.Contains(t, schema.Properties, "dry_run") + var gotDefault bool + require.NoError(t, json.Unmarshal(schema.Properties["dry_run"].Default, &gotDefault)) + assert.True(t, gotDefault) + + // region carries no override and must stay unconstrained. + require.Contains(t, schema.Properties, "region") + assert.Empty(t, schema.Properties["region"].Enum) +} + +func TestBuildInputSchemaUnknownFieldErrors(t *testing.T) { + t.Parallel() + _, err := BuildInputSchema[schemaTestArgs](map[string]FieldOverride{ + "instance_type": {Enum: []any{"m5.large"}}, // not a field on schemaTestArgs + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "instance_type") +} + +func TestBuildInputSchemaNilOverridesIsNoOp(t *testing.T) { + t.Parallel() + schema, err := BuildInputSchema[schemaTestArgs](nil) + require.NoError(t, err) + require.Contains(t, schema.Properties, "region") +} diff --git a/mcp/tools/search_recommendations.go b/mcp/tools/search_recommendations.go new file mode 100644 index 000000000..77e6ee47f --- /dev/null +++ b/mcp/tools/search_recommendations.go @@ -0,0 +1,345 @@ +package tools + +import ( + "context" + "errors" + "fmt" + "strings" + + "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +const searchRecommendationsName = "cudly_search_recommendations" + +const searchRecommendationsDescription = "Search for reserved-capacity purchase recommendations (RI/SP/CUD) " + + "across AWS, Azure, or GCP. Read-only: makes no purchase and spends no money -- there is no dry_run or " + + "confirm parameter because nothing is ever bought. Use this first to find what to buy, then feed a result's " + + "region/resource_type/count into the matching cudly____purchase tool." + +// searchRecommendationsArgs mirrors common.RecommendationParams, adding the +// provider selector and the optional per-call credential overrides from the +// design doc's §4 config-exposure model (aws_profile / +// azure_subscription_id / gcp_project_id). +type searchRecommendationsArgs struct { + Provider string `json:"provider" jsonschema:"cloud provider to search"` + Service string `json:"service" jsonschema:"service to search, e.g. ec2, rds, elasticache, compute, computeengine"` + Region string `json:"region,omitempty" jsonschema:"region to search; omit for account/global-level services such as Savings Plans"` + IncludeRegions []string `json:"include_regions,omitempty" jsonschema:"restrict the search to these regions, in addition to (or instead of) region"` + ExcludeRegions []string `json:"exclude_regions,omitempty" jsonschema:"exclude these regions from the search"` + LookbackPeriod string `json:"lookback_period,omitempty" jsonschema:"cost/usage lookback window backing the recommendation; AWS Savings Plans searches default to 30d when omitted (AWS requires a value); reservation searches (EC2/RDS/etc) leave it omitted to search all lookback windows"` + TermYears int `json:"term_years,omitempty" jsonschema:"commitment term filter; AWS Savings Plans searches default to 1 (1yr, no-upfront) when omitted (AWS requires a value); reservation searches (EC2/RDS/etc) leave it omitted to search all terms"` + PaymentOption string `json:"payment_option,omitempty" jsonschema:"payment schedule filter; AWS Savings Plans searches default to no-upfront when omitted (AWS requires a value); reservation searches (EC2/RDS/etc) leave it omitted to search all payment options"` + AccountFilter []string `json:"account_filter,omitempty" jsonschema:"restrict the search to these account/subscription/project IDs"` + IncludeSPTypes []string `json:"include_sp_types,omitempty" jsonschema:"AWS Savings Plans types to include; omit for all"` + ExcludeSPTypes []string `json:"exclude_sp_types,omitempty" jsonschema:"AWS Savings Plans types to exclude"` + AWSProfile string `json:"aws_profile,omitempty" jsonschema:"AWS named profile override (~/.aws/config); default uses ambient credentials"` + AzureSubscriptionID string `json:"azure_subscription_id,omitempty" jsonschema:"Azure subscription ID override; default uses AZURE_SUBSCRIPTION_ID"` + GCPProjectID string `json:"gcp_project_id,omitempty" jsonschema:"GCP project ID override; default uses ambient project"` +} + +// searchRecommendationsResult is the tool's structured output. +type searchRecommendationsResult struct { + Count int `json:"count"` + Recommendations []common.Recommendation `json:"recommendations"` +} + +type searchRecommendationsTool struct { + // createProvider is a seam over provider.CreateProvider so tests can + // inject a fake Provider without resolving real cloud credentials. + createProvider func(name string, cfg *provider.ProviderConfig) (provider.Provider, error) +} + +// NewSearchRecommendationsTool builds the cudly_search_recommendations tool. +func NewSearchRecommendationsTool() Registration { + return &searchRecommendationsTool{createProvider: provider.CreateProvider} +} + +func (t *searchRecommendationsTool) Descriptor() Descriptor { + return Descriptor{ + Name: searchRecommendationsName, + Description: searchRecommendationsDescription, + Action: "search", + ExamplePrompts: []string{ + "Search for AWS EC2 RI recommendations in us-east-1", + "What RDS Reserved Instance recommendations exist for account 123456789012?", + "Find GCP Compute Engine committed-use discount recommendations", + }, + } +} + +func (t *searchRecommendationsTool) Register(s *mcp.Server) error { + schema, err := BuildInputSchema[searchRecommendationsArgs](map[string]FieldOverride{ + "provider": {Enum: []any{string(common.ProviderAWS), string(common.ProviderAzure), string(common.ProviderGCP)}}, + "lookback_period": {Enum: []any{"7d", "30d", "60d"}}, + }) + if err != nil { + return err + } + mcp.AddTool(s, &mcp.Tool{ + Name: searchRecommendationsName, + Description: searchRecommendationsDescription, + InputSchema: schema, + }, t.handle) + return nil +} + +func (t *searchRecommendationsTool) handle(ctx context.Context, _ *mcp.CallToolRequest, args searchRecommendationsArgs) (*mcp.CallToolResult, searchRecommendationsResult, error) { + providerType, term, args, err := validateSearchArgs(args) + if err != nil { + return nil, searchRecommendationsResult{}, err + } + args = trimSearchArgsIdentifiers(args) + + prov, err := t.createProvider(string(providerType), providerConfigFromArgs(providerType, args)) + if err != nil { + return nil, searchRecommendationsResult{}, fmt.Errorf("create %s provider: %w", providerType, err) + } + + service, err := validateSupportedService(prov, args.Service) + if err != nil { + return nil, searchRecommendationsResult{}, err + } + + recClient, err := prov.GetRecommendationsClient(ctx) + if err != nil { + return nil, searchRecommendationsResult{}, fmt.Errorf("get %s recommendations client: %w", providerType, err) + } + + recs, err := recClient.GetRecommendations(ctx, recommendationParamsFromArgs(service, term, args)) + if err != nil { + return nil, searchRecommendationsResult{}, fmt.Errorf("get recommendations: %w", err) + } + + return nil, searchRecommendationsResult{Count: len(recs), Recommendations: recs}, nil +} + +// validateSearchArgs validates every money-neutral-but-still-typed field on +// args that does not require a live provider (provider name, payment +// option, lookback_period, term, Savings Plans type filters), returning the +// typed provider name, the normalised Recommendation term string +// ("1yr"/"3yr", or "" when args.TermYears was omitted), and args with +// applySavingsPlansSearchDefaults applied. +// +// Every invalid or missing field is collected and joined into a single +// error rather than returning on the first one found, so a caller missing +// several required fields (see requireSavingsPlansSearchFields) sees all of +// them at once instead of fixing them one call at a time. +func validateSearchArgs(args searchRecommendationsArgs) (common.ProviderType, string, searchRecommendationsArgs, error) { + providerType, err := validateProviderName(args.Provider) + if err != nil { + return "", "", args, err + } + args = applySavingsPlansSearchDefaults(providerType, args) + + var errs []error + + if args.PaymentOption != "" { + if _, err := ValidatePaymentOption(args.PaymentOption); err != nil { + errs = append(errs, err) + } + } + + if _, err := ValidateLookbackPeriod(args.LookbackPeriod); err != nil { + errs = append(errs, err) + } + + term := "" + if args.TermYears != 0 { + if ty, err := ValidateTermYears(args.TermYears); err != nil { + errs = append(errs, err) + } else { + term = ty.RecommendationTerm() + } + } + + if err := validateSPTypeFilters(args.IncludeSPTypes, args.ExcludeSPTypes); err != nil { + errs = append(errs, err) + } + + if err := requireSavingsPlansSearchFields(providerType, args); err != nil { + errs = append(errs, err) + } + + if len(errs) > 0 { + return "", "", args, errors.Join(errs...) + } + + return providerType, term, args, nil +} + +// isAWSSavingsPlansSearch reports whether args targets an AWS Savings Plans +// search -- the only combination where AWS's +// GetSavingsPlansPurchaseRecommendation requires term_years, payment_option, +// and lookback_period on every call. GetReservationPurchaseRecommendation +// (EC2/RDS/etc) defaults these server-side when omitted, so detection must +// never broaden past the Savings Plans family. Checked straight off the raw +// args -- no live provider needed -- via the same common.IsSavingsPlan +// family predicate the AWS recommendations client itself dispatches on +// (providers/aws/recommendations/client.go). +func isAWSSavingsPlansSearch(providerType common.ProviderType, service string) bool { + return providerType == common.ProviderAWS && common.IsSavingsPlan(common.ServiceType(service)) +} + +// applySavingsPlansSearchDefaults defaults payment_option, term_years, and +// lookback_period to no-upfront/1yr/30d when args targets an AWS Savings +// Plans search and the caller omitted them. Without this, an AWS SP search +// with these fields blank fails against Cost Explorer one field at a time +// (issue #1506): GetSavingsPlansPurchaseRecommendation requires all three, +// unlike the EC2/RDS reservation path this tool otherwise shares, which +// defaults them server-side. A caller-supplied value always wins -- this +// only fills in what was left blank. No-op for every other provider/service +// combination. +func applySavingsPlansSearchDefaults(providerType common.ProviderType, args searchRecommendationsArgs) searchRecommendationsArgs { + if !isAWSSavingsPlansSearch(providerType, args.Service) { + return args + } + if args.PaymentOption == "" { + args.PaymentOption = string(PaymentOptionNoUpfront) + } + if args.TermYears == 0 { + args.TermYears = int(TermOneYear) + } + if args.LookbackPeriod == "" { + args.LookbackPeriod = string(LookbackPeriod30Days) + } + return args +} + +// requireSavingsPlansSearchFields is a validate-all-at-once safety net: after +// applySavingsPlansSearchDefaults runs, an AWS Savings Plans search should +// never still be missing payment_option/term_years/lookback_period. If a +// future change to the defaulting logic leaves one blank anyway, this fails +// loud -- naming every still-missing field in one error -- rather than +// letting the request reach Cost Explorer and fail one field at a time (the +// cascade issue #1506 fixes). +func requireSavingsPlansSearchFields(providerType common.ProviderType, args searchRecommendationsArgs) error { + if !isAWSSavingsPlansSearch(providerType, args.Service) { + return nil + } + var missing []string + if args.PaymentOption == "" { + missing = append(missing, "payment_option") + } + if args.TermYears == 0 { + missing = append(missing, "term_years") + } + if args.LookbackPeriod == "" { + missing = append(missing, "lookback_period") + } + if len(missing) == 0 { + return nil + } + return fmt.Errorf("savings plans search missing required field(s): %s", strings.Join(missing, ", ")) +} + +// validateSPTypeFilters validates every entry of include/exclude against +// the AWS Savings Plans type enum, naming which filter a bad entry came +// from. +func validateSPTypeFilters(include, exclude []string) error { + for _, sp := range include { + if _, err := ValidateSPType(sp); err != nil { + return fmt.Errorf("include_sp_types: %w", err) + } + } + for _, sp := range exclude { + if _, err := ValidateSPType(sp); err != nil { + return fmt.Errorf("exclude_sp_types: %w", err) + } + } + return nil +} + +// trimSearchArgsIdentifiers returns args with surrounding whitespace +// stripped from every free-text identifier field that flows into +// providerConfigFromArgs/recommendationParamsFromArgs (region, +// include_regions, exclude_regions, account_filter). Unlike the purchase +// tools' requireNonBlank, region is optional here, so trimming rather than +// rejecting a blank/whitespace value is correct: " us-east-1 " must resolve +// and search the same region as "us-east-1" instead of silently searching +// the wrong (or account-default) region. +func trimSearchArgsIdentifiers(args searchRecommendationsArgs) searchRecommendationsArgs { + args.Region = strings.TrimSpace(args.Region) + args.IncludeRegions = trimAll(args.IncludeRegions) + args.ExcludeRegions = trimAll(args.ExcludeRegions) + args.AccountFilter = trimAll(args.AccountFilter) + return args +} + +// trimAll returns a copy of ss with strings.TrimSpace applied to every +// entry, preserving a nil input as nil. +func trimAll(ss []string) []string { + if ss == nil { + return nil + } + out := make([]string, len(ss)) + for i, s := range ss { + out[i] = strings.TrimSpace(s) + } + return out +} + +// providerConfigFromArgs builds the provider.ProviderConfig for the given +// provider from the tool's per-call credential override fields (design §4). +func providerConfigFromArgs(providerType common.ProviderType, args searchRecommendationsArgs) *provider.ProviderConfig { + return &provider.ProviderConfig{ + Name: string(providerType), + AWSProfile: args.AWSProfile, + AzureSubscriptionID: args.AzureSubscriptionID, + GCPProjectID: args.GCPProjectID, + Region: args.Region, + } +} + +// recommendationParamsFromArgs builds the common.RecommendationParams for +// the already-validated service and term. +func recommendationParamsFromArgs(service common.ServiceType, term string, args searchRecommendationsArgs) *common.RecommendationParams { + return &common.RecommendationParams{ + Service: service, + Region: args.Region, + LookbackPeriod: args.LookbackPeriod, + Term: term, + PaymentOption: args.PaymentOption, + AccountFilter: args.AccountFilter, + IncludeRegions: args.IncludeRegions, + ExcludeRegions: args.ExcludeRegions, + IncludeSPTypes: args.IncludeSPTypes, + ExcludeSPTypes: args.ExcludeSPTypes, + } +} + +// validateProviderName returns the typed common.ProviderType for s, or an +// explicit error when s is not aws, azure, or gcp. +func validateProviderName(s string) (common.ProviderType, error) { + switch common.ProviderType(s) { + case common.ProviderAWS, common.ProviderAzure, common.ProviderGCP: + return common.ProviderType(s), nil + default: + return "", fmt.Errorf("invalid provider %q: must be one of %s, %s, %s", + s, common.ProviderAWS, common.ProviderAzure, common.ProviderGCP) + } +} + +// validateSupportedService checks service against prov's own +// GetSupportedServices() -- the provider's live list, not a hardcoded +// mirror of it -- so this tool never drifts from what each provider +// actually supports. +func validateSupportedService(prov provider.Provider, service string) (common.ServiceType, error) { + if service == "" { + return "", fmt.Errorf("service is required") + } + want := common.ServiceType(service) + supported := prov.GetSupportedServices() + for _, s := range supported { + if s == want { + return want, nil + } + } + names := make([]string, len(supported)) + for i, s := range supported { + names[i] = s.String() + } + return "", fmt.Errorf("invalid service %q for provider %s: must be one of %s", service, prov.Name(), strings.Join(names, ", ")) +} diff --git a/mcp/tools/search_recommendations_test.go b/mcp/tools/search_recommendations_test.go new file mode 100644 index 000000000..ffbae03bb --- /dev/null +++ b/mcp/tools/search_recommendations_test.go @@ -0,0 +1,360 @@ +package tools + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/LeanerCloud/CUDly/pkg/common" + "github.com/LeanerCloud/CUDly/pkg/provider" +) + +// fakeRecommendationsClient is a minimal provider.RecommendationsClient test +// double; only GetRecommendations is exercised by search_recommendations. +type fakeRecommendationsClient struct { + lastParams *common.RecommendationParams + recs []common.Recommendation + err error + calls int +} + +func (f *fakeRecommendationsClient) GetRecommendations(_ context.Context, params *common.RecommendationParams) ([]common.Recommendation, error) { + f.calls++ + f.lastParams = params + return f.recs, f.err +} +func (f *fakeRecommendationsClient) GetRecommendationsForService(_ context.Context, _ common.ServiceType) ([]common.Recommendation, error) { + return f.recs, f.err +} +func (f *fakeRecommendationsClient) GetAllRecommendations(_ context.Context) ([]common.Recommendation, error) { + return f.recs, f.err +} + +var _ provider.RecommendationsClient = (*fakeRecommendationsClient)(nil) + +// fakeProvider is a minimal provider.Provider test double. +type fakeProvider struct { + name string + services []common.ServiceType + recClient provider.RecommendationsClient + recErr error +} + +func (f *fakeProvider) Name() string { return f.name } +func (f *fakeProvider) DisplayName() string { return f.name } +func (f *fakeProvider) IsConfigured() bool { return true } +func (f *fakeProvider) GetCredentials() (provider.Credentials, error) { + return nil, nil +} +func (f *fakeProvider) ValidateCredentials(_ context.Context) error { return nil } +func (f *fakeProvider) GetAccounts(_ context.Context) ([]common.Account, error) { + return nil, nil +} +func (f *fakeProvider) GetRegions(_ context.Context) ([]common.Region, error) { + return nil, nil +} +func (f *fakeProvider) GetDefaultRegion() string { return "us-east-1" } +func (f *fakeProvider) GetSupportedServices() []common.ServiceType { + return f.services +} +func (f *fakeProvider) GetServiceClient(_ context.Context, _ common.ServiceType, _ string) (provider.ServiceClient, error) { + return nil, nil +} +func (f *fakeProvider) GetRecommendationsClient(_ context.Context) (provider.RecommendationsClient, error) { + return f.recClient, f.recErr +} + +var _ provider.Provider = (*fakeProvider)(nil) + +func newTestSearchTool(fp *fakeProvider) *searchRecommendationsTool { + return &searchRecommendationsTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return fp, nil + }, + } +} + +func TestSearchRecommendationsHappyPath(t *testing.T) { + t.Parallel() + recs := []common.Recommendation{{Provider: common.ProviderAWS, ResourceType: "m5.large", Count: 2}} + client := &fakeRecommendationsClient{recs: recs} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}, recClient: client} + tool := newTestSearchTool(fp) + + _, result, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + Region: "us-east-1", + }) + + require.NoError(t, err) + assert.Equal(t, 1, result.Count) + assert.Equal(t, recs, result.Recommendations) + require.NotNil(t, client.lastParams) + assert.Equal(t, common.ServiceEC2, client.lastParams.Service) + assert.Equal(t, "us-east-1", client.lastParams.Region) +} + +// TestSearchRecommendationsForwardsRegionFilters proves finding B of the +// CodeRabbit review: common.RecommendationParams has IncludeRegions and +// ExcludeRegions, but the tool neither accepted nor forwarded them, so a +// caller could not restrict a search to (or exclude) specific regions the +// way the CLI's config supports. include_regions/exclude_regions must reach +// the underlying RecommendationsClient call unchanged. +func TestSearchRecommendationsForwardsRegionFilters(t *testing.T) { + t.Parallel() + client := &fakeRecommendationsClient{} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}, recClient: client} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + IncludeRegions: []string{"us-east-1", "us-west-2"}, + ExcludeRegions: []string{"eu-west-1"}, + }) + + require.NoError(t, err) + require.NotNil(t, client.lastParams) + assert.Equal(t, []string{"us-east-1", "us-west-2"}, client.lastParams.IncludeRegions) + assert.Equal(t, []string{"eu-west-1"}, client.lastParams.ExcludeRegions) +} + +// TestSearchRecommendationsTrimsRegionFilters is the regression guard for +// the CodeRabbit finding: region (and include_regions/exclude_regions/ +// account_filter) were passed straight through to ProviderConfig.Region and +// RecommendationParams with no trim, unlike the purchase tools' +// requireNonBlank. " us-east-1 " must forward as "us-east-1" to both the +// provider config used to resolve the client and the params sent to the +// recommendations client -- region is optional here, so trimming (not +// rejecting) surrounding whitespace is the correct behavior. +func TestSearchRecommendationsTrimsRegionFilters(t *testing.T) { + t.Parallel() + client := &fakeRecommendationsClient{} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}, recClient: client} + var gotCfg *provider.ProviderConfig + tool := &searchRecommendationsTool{ + createProvider: func(_ string, cfg *provider.ProviderConfig) (provider.Provider, error) { + gotCfg = cfg + return fp, nil + }, + } + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + Region: " us-east-1 ", + IncludeRegions: []string{" us-east-1 ", " us-west-2 "}, + ExcludeRegions: []string{" eu-west-1 "}, + AccountFilter: []string{" 123456789012 "}, + }) + + require.NoError(t, err) + require.NotNil(t, gotCfg) + assert.Equal(t, "us-east-1", gotCfg.Region, "ProviderConfig.Region must be trimmed") + require.NotNil(t, client.lastParams) + assert.Equal(t, "us-east-1", client.lastParams.Region, "RecommendationParams.Region must be trimmed") + assert.Equal(t, []string{"us-east-1", "us-west-2"}, client.lastParams.IncludeRegions, "IncludeRegions entries must be trimmed") + assert.Equal(t, []string{"eu-west-1"}, client.lastParams.ExcludeRegions, "ExcludeRegions entries must be trimmed") + assert.Equal(t, []string{"123456789012"}, client.lastParams.AccountFilter, "AccountFilter entries must be trimmed") +} + +func TestSearchRecommendationsInvalidProvider(t *testing.T) { + t.Parallel() + tool := newTestSearchTool(&fakeProvider{}) + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "openstack", + Service: "ec2", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid provider") +} + +func TestSearchRecommendationsUnsupportedService(t *testing.T) { + t.Parallel() + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2, common.ServiceRDS}} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "cosmosdb", // not an AWS service + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid service") +} + +func TestSearchRecommendationsInvalidTermYears(t *testing.T) { + t.Parallel() + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + TermYears: 2, + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid term_years") +} + +// TestSearchRecommendationsInvalidLookbackPeriod is the regression guard for +// the CodeRabbit finding: lookback_period was constrained only by the +// advertised MCP jsonschema enum (Register's BuildInputSchema override), not +// re-validated in the handler, so a direct MCP call bypassing schema +// enforcement could pass an unsupported value through to the provider. +func TestSearchRecommendationsInvalidLookbackPeriod(t *testing.T) { + t.Parallel() + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + LookbackPeriod: "90d", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid lookback_period") +} + +func TestSearchRecommendationsInvalidSPType(t *testing.T) { + t.Parallel() + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceSavingsPlansAll}} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: string(common.ServiceSavingsPlansAll), + IncludeSPTypes: []string{"NotARealType"}, + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "include_sp_types") +} + +func TestSearchRecommendationsProviderErrorSurfaced(t *testing.T) { + t.Parallel() + tool := &searchRecommendationsTool{ + createProvider: func(_ string, _ *provider.ProviderConfig) (provider.Provider, error) { + return nil, errors.New("no AWS credentials found") + }, + } + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "no AWS credentials found") +} + +func TestSearchRecommendationsClientErrorSurfaced(t *testing.T) { + t.Parallel() + client := &fakeRecommendationsClient{err: errors.New("Cost Explorer API throttled")} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}, recClient: client} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "Cost Explorer API throttled") +} + +// TestSearchRecommendationsSPDefaultsAppliedWhenOmitted is the regression +// guard for issue #1506: AWS's GetSavingsPlansPurchaseRecommendation +// requires term/payment_option/lookback_period on every call, unlike +// GetReservationPurchaseRecommendation (EC2/RDS/etc), which defaults them +// server-side. Before the fix, an SP search omitting these three fields +// forwarded them blank and only succeeds after AWS rejects the request (in +// production, that means a caller must retry with each field filled in, one +// at a time). This proves a single search call now builds the +// no-upfront/1yr/30d defaults and reaches the provider exactly once. +func TestSearchRecommendationsSPDefaultsAppliedWhenOmitted(t *testing.T) { + t.Parallel() + client := &fakeRecommendationsClient{} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceSavingsPlansAll}, recClient: client} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: string(common.ServiceSavingsPlansAll), + }) + + require.NoError(t, err) + assert.Equal(t, 1, client.calls, "must reach the provider exactly once") + require.NotNil(t, client.lastParams) + assert.Equal(t, "1yr", client.lastParams.Term) + assert.Equal(t, "no-upfront", client.lastParams.PaymentOption) + assert.Equal(t, "30d", client.lastParams.LookbackPeriod) +} + +// TestSearchRecommendationsSPExplicitTermYearsPreserved proves the SP +// defaults never override an explicitly supplied value: term_years=3 must +// reach the provider as "3yr", not the 1yr default. +func TestSearchRecommendationsSPExplicitTermYearsPreserved(t *testing.T) { + t.Parallel() + client := &fakeRecommendationsClient{} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceSavingsPlansAll}, recClient: client} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: string(common.ServiceSavingsPlansAll), + TermYears: 3, + }) + + require.NoError(t, err) + require.NotNil(t, client.lastParams) + assert.Equal(t, "3yr", client.lastParams.Term) + assert.Equal(t, "no-upfront", client.lastParams.PaymentOption, "payment_option still defaults when omitted") + assert.Equal(t, "30d", client.lastParams.LookbackPeriod, "lookback_period still defaults when omitted") +} + +// TestSearchRecommendationsEC2NoDefaultsInjected proves the SP defaults are +// scoped to AWS Savings Plans searches only: an EC2 (reservation) search +// omitting term/payment/lookback must reach the provider with those fields +// still blank, since GetReservationPurchaseRecommendation already searches +// all terms/payment options/lookback windows when they're omitted -- +// defaulting them here would wrongly narrow those results. +func TestSearchRecommendationsEC2NoDefaultsInjected(t *testing.T) { + t.Parallel() + client := &fakeRecommendationsClient{} + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}, recClient: client} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + }) + + require.NoError(t, err) + require.NotNil(t, client.lastParams) + assert.Empty(t, client.lastParams.Term) + assert.Empty(t, client.lastParams.PaymentOption) + assert.Empty(t, client.lastParams.LookbackPeriod) +} + +// TestSearchRecommendationsMultipleInvalidFieldsJoinedInOneError is the +// validate-all-at-once safety net (issue #1506 change 2): a request with +// several invalid fields must name all of them in one returned error instead +// of surfacing only the first one found. +func TestSearchRecommendationsMultipleInvalidFieldsJoinedInOneError(t *testing.T) { + t.Parallel() + fp := &fakeProvider{name: "aws", services: []common.ServiceType{common.ServiceEC2}} + tool := newTestSearchTool(fp) + + _, _, err := tool.handle(context.Background(), nil, searchRecommendationsArgs{ + Provider: "aws", + Service: "ec2", + PaymentOption: "not-a-real-option", + LookbackPeriod: "90d", + IncludeSPTypes: []string{"NotARealType"}, + }) + + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid payment_option") + assert.Contains(t, err.Error(), "invalid lookback_period") + assert.Contains(t, err.Error(), "include_sp_types") +} diff --git a/pkg/common/types.go b/pkg/common/types.go index d5507618f..56d4e9b12 100644 --- a/pkg/common/types.go +++ b/pkg/common/types.go @@ -312,6 +312,7 @@ type PurchaseResult struct { const ( PurchaseSourceCLI = "cudly-cli" PurchaseSourceWeb = "cudly-web" + PurchaseSourceMCP = "cudly-mcp" ) // PurchaseTagKey is the tag/label key every CUDly-purchased commitment carries @@ -368,12 +369,12 @@ type PurchaseOptions struct { func NormalizeSource(s string) (string, error) { lower := strings.ToLower(strings.TrimSpace(s)) switch lower { - case PurchaseSourceCLI, PurchaseSourceWeb: + case PurchaseSourceCLI, PurchaseSourceWeb, PurchaseSourceMCP: return lower, nil case "": return "", fmt.Errorf("purchase source is required") default: - return "", fmt.Errorf("invalid purchase source %q (allowed: %s, %s)", s, PurchaseSourceCLI, PurchaseSourceWeb) + return "", fmt.Errorf("invalid purchase source %q (allowed: %s, %s, %s)", s, PurchaseSourceCLI, PurchaseSourceWeb, PurchaseSourceMCP) } } diff --git a/pkg/common/types_test.go b/pkg/common/types_test.go index 1949a8b61..d613dbe92 100644 --- a/pkg/common/types_test.go +++ b/pkg/common/types_test.go @@ -480,8 +480,10 @@ func TestNormalizeSource(t *testing.T) { }{ {"cli lowercase", "cudly-cli", "cudly-cli", false}, {"web lowercase", "cudly-web", "cudly-web", false}, + {"mcp lowercase", "cudly-mcp", "cudly-mcp", false}, {"cli mixed case", "CUDly-CLI", "cudly-cli", false}, {"web mixed case", "CUDly-Web", "cudly-web", false}, + {"mcp mixed case", "CUDly-MCP", "cudly-mcp", false}, {"cli with whitespace", " cudly-cli\n", "cudly-cli", false}, {"empty string", "", "", true}, {"whitespace only", " ", "", true}, diff --git a/providers/aws/service_client.go b/providers/aws/service_client.go index 8e5aa1f8d..d8765eb5e 100644 --- a/providers/aws/service_client.go +++ b/providers/aws/service_client.go @@ -85,14 +85,27 @@ func (r *RecommendationsClientAdapter) GetRecommendations(ctx context.Context, p return recs, nil } -// applyRecommendationFilters applies account and region filters to recommendations +// applyRecommendationFilters applies account and region filters to recommendations. +// +// GetReservationPurchaseRecommendation and GetSavingsPlansPurchaseRecommendation +// are both account-level Cost Explorer APIs: neither request carries a region +// parameter, so AWS returns recommendations across every region the account +// has usage in regardless of what params.Region asks for (issue #1506) -- +// this is the only region enforcement a single-region search gets. Region is +// folded into the include-region set here (rather than requiring the caller +// to pass include_regions) so cudly_search_recommendations' region="us-east-1" +// argument -- documented as filtering the search -- actually does. func applyRecommendationFilters(recs []common.Recommendation, params common.RecommendationParams) []common.Recommendation { if len(params.AccountFilter) > 0 { recs = filterByAccounts(recs, params.AccountFilter) } - if len(params.IncludeRegions) > 0 { - recs = filterByIncludedRegions(recs, params.IncludeRegions) + includeRegions := params.IncludeRegions + if params.Region != "" { + includeRegions = append(append([]string{}, includeRegions...), params.Region) + } + if len(includeRegions) > 0 { + recs = filterByIncludedRegions(recs, includeRegions) } if len(params.ExcludeRegions) > 0 { diff --git a/providers/aws/service_client_test.go b/providers/aws/service_client_test.go index ab0f11191..87481f423 100644 --- a/providers/aws/service_client_test.go +++ b/providers/aws/service_client_test.go @@ -290,3 +290,51 @@ func TestRecommendationsClientAdapter_GetAllRecommendations(t *testing.T) { assert.Error(t, err) }) } + +// TestApplyRecommendationFilters_Region is the regression guard for issue +// #1506's Change 3: GetReservationPurchaseRecommendation and +// GetSavingsPlansPurchaseRecommendation are account-level Cost Explorer +// calls with no region parameter, so AWS returns recommendations from every +// region the account has usage in regardless of params.Region. Before the +// fix, only params.IncludeRegions/ExcludeRegions were honored here, so a +// caller passing region alone (as cudly_search_recommendations documents) +// got no region filtering at all -- an eu-west-1 recommendation could +// surface from a us-east-1 search. +func TestApplyRecommendationFilters_Region(t *testing.T) { + recs := []common.Recommendation{ + {Account: "111", Region: "us-east-1"}, + {Account: "222", Region: "eu-west-1"}, + } + + t.Run("region alone filters out other regions", func(t *testing.T) { + got := applyRecommendationFilters(recs, common.RecommendationParams{Region: "us-east-1"}) + require.Len(t, got, 1) + assert.Equal(t, "us-east-1", got[0].Region) + }) + + t.Run("no region constraint returns everything", func(t *testing.T) { + got := applyRecommendationFilters(recs, common.RecommendationParams{}) + assert.Len(t, got, 2) + }) + + t.Run("region and include_regions are additive", func(t *testing.T) { + threeRegionRecs := append(append([]common.Recommendation{}, recs...), common.Recommendation{Account: "333", Region: "ap-southeast-1"}) + got := applyRecommendationFilters(threeRegionRecs, common.RecommendationParams{ + Region: "us-east-1", + IncludeRegions: []string{"eu-west-1"}, + }) + gotRegions := make([]string, len(got)) + for i, r := range got { + gotRegions[i] = r.Region + } + assert.ElementsMatch(t, []string{"us-east-1", "eu-west-1"}, gotRegions) + }) + + t.Run("exclude_regions still applies on top of region", func(t *testing.T) { + got := applyRecommendationFilters(recs, common.RecommendationParams{ + Region: "us-east-1", + ExcludeRegions: []string{"us-east-1"}, + }) + assert.Empty(t, got) + }) +} diff --git a/providers/azure/services/compute/client.go b/providers/azure/services/compute/client.go index fd447817b..6e1ded048 100644 --- a/providers/azure/services/compute/client.go +++ b/providers/azure/services/compute/client.go @@ -407,11 +407,14 @@ func (c *ComputeClient) triggerCapacityProviderRegistration(ctx context.Context, // buildReservationBody builds the JSON body for a reservation purchase request. // The same body is sent to both calculatePrice and purchase endpoints (issue #677). +// billingPlan is resolved by the caller (PurchaseCommitment) via +// reservations.BillingPlanForPaymentOption before any side-effecting call is +// made, so it is threaded in here rather than re-derived from rec.PaymentOption. // The purchase-automation and cudly-idempotency-token tags are attached via // reservations.ApplyPurchaseTags so the resulting reservation is identifiable // in the portal AND a re-driven purchase can find it via tag lookup before // buying a duplicate (issue #721). -func (c *ComputeClient) buildReservationBody(rec common.Recommendation, source, idempotencyToken string) ([]byte, error) { +func (c *ComputeClient) buildReservationBody(rec common.Recommendation, billingPlan armreservations.ReservationBillingPlan, source, idempotencyToken string) ([]byte, error) { termYears, err := reservations.ParseTermYears(rec.Term) if err != nil { return nil, err @@ -422,6 +425,7 @@ func (c *ComputeClient) buildReservationBody(rec common.Recommendation, source, "properties": map[string]interface{}{ "reservedResourceType": string(armreservations.ReservedResourceTypeVirtualMachines), "billingScopeId": fmt.Sprintf("/subscriptions/%s", c.subscriptionID), + "billingPlan": string(billingPlan), "term": fmt.Sprintf("P%dY", termYears), "quantity": rec.Count, "displayName": reservations.BuildDisplayName(reservations.DisplayNameFields{ @@ -478,15 +482,27 @@ func (c *ComputeClient) PurchaseCommitment(ctx context.Context, rec common.Recom return result, result.Error } - // Ensure Microsoft.Capacity provider is registered (cached after first call). - c.ensureCapacityProviderRegistered(ctx) + // Validate the payment option and build the request body BEFORE any + // side-effecting call. Microsoft.Capacity provider registration below is + // a real ARM operation (a GET, and a POST to register when unregistered); + // every fallible local parse -- payment option AND term (buildReservationBody + // calls reservations.ParseTermYears) -- must be rejected here first so a + // doomed purchase never triggers it. + billingPlan, err := reservations.BillingPlanForPaymentOption(rec.PaymentOption) + if err != nil { + result.Error = err + return result, result.Error + } - bodyBytes, err := c.buildReservationBody(rec, opts.Source, opts.IdempotencyToken) + bodyBytes, err := c.buildReservationBody(rec, billingPlan, opts.Source, opts.IdempotencyToken) if err != nil { result.Error = fmt.Errorf("failed to marshal request: %w", err) return result, result.Error } + // Ensure Microsoft.Capacity provider is registered (cached after first call). + c.ensureCapacityProviderRegistered(ctx) + token, err := c.cred.GetToken(ctx, policy.TokenRequestOptions{ Scopes: []string{"https://management.azure.com/.default"}, }) diff --git a/providers/azure/services/compute/client_test.go b/providers/azure/services/compute/client_test.go index 97862ea94..8c551fd29 100644 --- a/providers/azure/services/compute/client_test.go +++ b/providers/azure/services/compute/client_test.go @@ -22,6 +22,7 @@ import ( "github.com/LeanerCloud/CUDly/pkg/common" "github.com/LeanerCloud/CUDly/providers/azure/mocks" + "github.com/LeanerCloud/CUDly/providers/azure/services/internal/reservations" ) func TestNewClient(t *testing.T) { @@ -633,6 +634,7 @@ func TestComputeClient_PurchaseCommitment_Success(t *testing.T) { Term: "1yr", Count: 1, CommitmentCost: 2000.0, + PaymentOption: "no-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -662,6 +664,7 @@ func TestComputeClient_PurchaseCommitment_3YearTerm(t *testing.T) { Term: "3yr", Count: 1, CommitmentCost: 5000.0, + PaymentOption: "all-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -690,6 +693,7 @@ func TestComputeClient_PurchaseCommitment_Accepted(t *testing.T) { Term: "1yr", Count: 1, CommitmentCost: 2000.0, + PaymentOption: "no-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -705,9 +709,10 @@ func TestComputeClient_PurchaseCommitment_TokenError(t *testing.T) { client := NewClientWithHTTP(mockCred, "test-subscription", "eastus", mockHTTP) rec := common.Recommendation{ - ResourceType: "Standard_D2s_v3", - Term: "1yr", - Count: 1, + ResourceType: "Standard_D2s_v3", + Term: "1yr", + Count: 1, + PaymentOption: "no-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -729,9 +734,10 @@ func TestComputeClient_PurchaseCommitment_HTTPError(t *testing.T) { })).Return(nil, errors.New("network error")).Once() rec := common.Recommendation{ - ResourceType: "Standard_D2s_v3", - Term: "1yr", - Count: 1, + ResourceType: "Standard_D2s_v3", + Term: "1yr", + Count: 1, + PaymentOption: "no-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -759,9 +765,10 @@ func TestComputeClient_PurchaseCommitment_BadStatus(t *testing.T) { ).Once() rec := common.Recommendation{ - ResourceType: "Standard_D2s_v3", - Term: "1yr", - Count: 1, + ResourceType: "Standard_D2s_v3", + Term: "1yr", + Count: 1, + PaymentOption: "no-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -800,6 +807,7 @@ func TestComputeClient_PurchaseCommitment_TwoStepFlow(t *testing.T) { Term: "1yr", Count: 1, CommitmentCost: 500.0, + PaymentOption: "all-upfront", } result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) @@ -842,7 +850,7 @@ func TestComputeClient_PurchaseCommitment_SessionTimeoutRetry(t *testing.T) { return r.URL.Path == "/providers/Microsoft.Capacity/reservationOrders/order-second/purchase" })).Return(mocks.CreateMockHTTPResponse(http.StatusOK, `{}`), nil).Once() - rec := common.Recommendation{ResourceType: "Standard_B2ats_v2", Term: "1yr", Count: 1} + rec := common.Recommendation{ResourceType: "Standard_B2ats_v2", Term: "1yr", Count: 1, PaymentOption: "no-upfront"} result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) require.NoError(t, err) assert.True(t, result.Success) @@ -881,7 +889,7 @@ func TestComputeClient_PurchaseCommitment_TagInjection(t *testing.T) { r.URL.Path == "/providers/Microsoft.Capacity/reservationOrders/"+orderID+"/purchase" })).Return(mocks.CreateMockHTTPResponse(http.StatusOK, `{}`), nil).Once() - rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Term: "1yr", Count: 1, CommitmentCost: 2000.0} + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Term: "1yr", Count: 1, CommitmentCost: 2000.0, PaymentOption: "monthly"} result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: source}) require.NoError(t, err) assert.True(t, result.Success) @@ -915,6 +923,53 @@ func TestComputeClient_PurchaseCommitment_RequiresSource(t *testing.T) { mockHTTP.AssertNotCalled(t, "Do", mock.Anything) } +// TestComputeClient_PurchaseCommitment_RejectsInvalidPaymentOptionBeforeSideEffects +// pins the fix ordering: an invalid/empty PaymentOption must be rejected +// BEFORE PurchaseCommitment ever calls ensureCapacityProviderRegistered, which +// issues a real ARM GET (and potentially a POST to register) against +// Microsoft.Capacity. Before this fix, the provider-registration check ran +// first, so a doomed purchase (bad payment option) still triggered that +// side-effecting ARM call. mockHTTP.AssertNotCalled proves zero HTTP calls +// were made -- this test fails pre-fix, because the old code issued the +// Microsoft.Capacity GET before buildReservationBody ever validated +// PaymentOption. +func TestComputeClient_PurchaseCommitment_RejectsInvalidPaymentOptionBeforeSideEffects(t *testing.T) { + ctx := context.Background() + mockHTTP := &mocks.MockHTTPClient{} + mockCred := &MockTokenCredential{token: "test-token"} + client := NewClientWithHTTP(mockCred, "test-subscription", "eastus", mockHTTP) + + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Term: "1yr", Count: 1, PaymentOption: "partial-upfront"} + result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) + require.Error(t, err) + assert.False(t, result.Success) + assert.Contains(t, err.Error(), "partial-upfront has no azure equivalent") + mockHTTP.AssertNotCalled(t, "Do", mock.Anything) +} + +// TestComputeClient_PurchaseCommitment_RejectsInvalidTermBeforeSideEffects +// pins the same guarantee as the invalid-payment-option test above but for +// rec.Term: buildReservationBody parses the term via +// reservations.ParseTermYears, and that parse must be rejected before +// ensureCapacityProviderRegistered's real ARM GET/POST ever fires. Regression +// test for a prior fix that reordered PaymentOption validation ahead of +// registration but left Term parsing (inside buildReservationBody) running +// after it, so an invalid term still triggered the provider-registration +// side effect before the purchase was rejected. +func TestComputeClient_PurchaseCommitment_RejectsInvalidTermBeforeSideEffects(t *testing.T) { + ctx := context.Background() + mockHTTP := &mocks.MockHTTPClient{} + mockCred := &MockTokenCredential{token: "test-token"} + client := NewClientWithHTTP(mockCred, "test-subscription", "eastus", mockHTTP) + + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Term: "5yr", Count: 1, PaymentOption: "no-upfront"} + result, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) + require.Error(t, err) + assert.False(t, result.Success) + assert.Contains(t, err.Error(), "unsupported reservation term") + mockHTTP.AssertNotCalled(t, "Do", mock.Anything) +} + // TestComputeClient_ConvertAzureVMRecommendation_NilGuards pins the new // contract: unusable SDK payloads (nil, wrong concrete type, nil Properties) // produce a nil *Recommendation so the caller can filter it out. Before @@ -990,11 +1045,67 @@ func TestFetchAzurePricing_WrapperSmokeTest(t *testing.T) { assert.Equal(t, "Standard_D2s_v3", result.Items[0].ArmSKUName) } +// TestBuildReservationBody_BillingPlan pins the billingPlan wiring: Azure +// reservations support exactly two billing plans (Upfront, Monthly -- no +// partial-upfront), and rec.PaymentOption must map onto the correct one +// regardless of which vocabulary populated it (the converter's +// "upfront"/"monthly" or the CLI/MCP's "all-upfront"/"no-upfront"). An empty +// or unrecognized value (including "partial-upfront", which Azure cannot +// express) must fail loud rather than silently defaulting to Upfront +// (feedback_no_silent_fallbacks). +// +// buildReservationBody itself no longer resolves the billing plan: the +// caller (PurchaseCommitment) validates rec.PaymentOption via +// reservations.BillingPlanForPaymentOption before any side-effecting call +// (Microsoft.Capacity provider registration) and threads the result in, so +// this test mirrors that same call order rather than duplicating validation +// inside buildReservationBody. +func TestBuildReservationBody_BillingPlan(t *testing.T) { + cases := []struct { + name string + paymentOption string + wantPlan string + wantErrSub string + }{ + {name: "all-upfront maps to Upfront", paymentOption: "all-upfront", wantPlan: "Upfront"}, + {name: "upfront maps to Upfront", paymentOption: "upfront", wantPlan: "Upfront"}, + {name: "no-upfront maps to Monthly", paymentOption: "no-upfront", wantPlan: "Monthly"}, + {name: "monthly maps to Monthly", paymentOption: "monthly", wantPlan: "Monthly"}, + {name: "partial-upfront is rejected", paymentOption: "partial-upfront", wantErrSub: "partial-upfront has no azure equivalent"}, + {name: "empty payment option is rejected", paymentOption: "", wantErrSub: "azure reservations support only upfront or monthly billing"}, + {name: "unrecognized payment option is rejected", paymentOption: "bogus", wantErrSub: "azure reservations support only upfront or monthly billing"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + billingPlan, err := reservations.BillingPlanForPaymentOption(tc.paymentOption) + if tc.wantErrSub != "" { + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErrSub) + return + } + require.NoError(t, err) + + c := &ComputeClient{region: "eastus", subscriptionID: "sub-abc"} + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr", PaymentOption: tc.paymentOption} + + body, err := c.buildReservationBody(rec, billingPlan, common.PurchaseSourceWeb, "") + require.NoError(t, err) + var got map[string]interface{} + require.NoError(t, json.Unmarshal(body, &got)) + props, ok := got["properties"].(map[string]interface{}) + require.True(t, ok, "properties map missing from reservation body") + assert.Equal(t, tc.wantPlan, props["billingPlan"]) + }) + } +} + func TestBuildReservationBody_IncludesPurchaseAutomationTag(t *testing.T) { c := &ComputeClient{region: "eastus", subscriptionID: "sub-abc"} - rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr"} + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr", PaymentOption: "no-upfront"} + billingPlan, err := reservations.BillingPlanForPaymentOption(rec.PaymentOption) + require.NoError(t, err) - body, err := c.buildReservationBody(rec, common.PurchaseSourceWeb, "") + body, err := c.buildReservationBody(rec, billingPlan, common.PurchaseSourceWeb, "") require.NoError(t, err) var got map[string]interface{} @@ -1006,9 +1117,11 @@ func TestBuildReservationBody_IncludesPurchaseAutomationTag(t *testing.T) { func TestBuildReservationBody_OmitsTagsWhenSourceAndTokenEmpty(t *testing.T) { c := &ComputeClient{region: "eastus", subscriptionID: "sub-abc"} - rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr"} + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr", PaymentOption: "no-upfront"} + billingPlan, err := reservations.BillingPlanForPaymentOption(rec.PaymentOption) + require.NoError(t, err) - body, err := c.buildReservationBody(rec, "", "") + body, err := c.buildReservationBody(rec, billingPlan, "", "") require.NoError(t, err) var got map[string]interface{} @@ -1024,10 +1137,12 @@ func TestBuildReservationBody_OmitsTagsWhenSourceAndTokenEmpty(t *testing.T) { // and skip the duplicate buy. func TestBuildReservationBody_IncludesIdempotencyTokenTag(t *testing.T) { c := &ComputeClient{region: "eastus", subscriptionID: "sub-abc"} - rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr"} + rec := common.Recommendation{ResourceType: "Standard_D2s_v3", Count: 1, Term: "1yr", PaymentOption: "no-upfront"} token := common.DeriveIdempotencyToken("exec-721-compute", 0) + billingPlan, err := reservations.BillingPlanForPaymentOption(rec.PaymentOption) + require.NoError(t, err) - body, err := c.buildReservationBody(rec, common.PurchaseSourceWeb, token) + body, err := c.buildReservationBody(rec, billingPlan, common.PurchaseSourceWeb, token) require.NoError(t, err) var got map[string]interface{} @@ -1280,6 +1395,7 @@ func TestComputeClient_PurchaseCommitment_DisplayNameConformsToAzureAllowlist(t Term: "1yr", Count: 1, CommitmentCost: 2000.0, + PaymentOption: "all-upfront", } _, err := client.PurchaseCommitment(ctx, rec, common.PurchaseOptions{Source: common.PurchaseSourceCLI}) require.NoError(t, err) diff --git a/providers/azure/services/internal/reservations/purchase.go b/providers/azure/services/internal/reservations/purchase.go index 7bf666e69..c868761ce 100644 --- a/providers/azure/services/internal/reservations/purchase.go +++ b/providers/azure/services/internal/reservations/purchase.go @@ -49,6 +49,8 @@ import ( "strings" "time" + "github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/reservations/armreservations" + "github.com/LeanerCloud/CUDly/pkg/common" ) @@ -70,6 +72,41 @@ func ParseTermYears(term string) (int, error) { } } +// BillingPlanForPaymentOption maps a recommendation's payment-option string +// to the armreservations.ReservationBillingPlan Azure's purchase API expects +// in properties.billingPlan (confirmed against +// armreservations.PurchaseRequestProperties.BillingPlan and the two-member +// enum in constants.go: ReservationBillingPlanUpfront = "Upfront", +// ReservationBillingPlanMonthly = "Monthly"). Azure reservations support +// exactly those two billing plans -- there is no partial-upfront -- and +// Monthly costs the same total as Upfront (no premium for spreading +// payments), so "no-upfront"/"monthly" is a safe, cost-neutral default at +// the caller layer. +// +// CUDly's own recommendation converter (providers/azure/internal/ +// recommendations/converter.go) emits PaymentOption "upfront"/"monthly"; +// the MCP tool boundary and the CLI --payment flag use "all-upfront"/ +// "no-upfront"/"partial-upfront". Both vocabularies are accepted here so +// this function is the single mapping point regardless of which caller +// populated rec.PaymentOption. +// +// An empty or unrecognized value (including "partial-upfront", which Azure +// cannot express at all) is a hard error rather than a silent default: the +// default belongs at the caller (CLI/MCP) layer, never silently applied on +// this money-affecting path (feedback_no_silent_fallbacks). +func BillingPlanForPaymentOption(paymentOption string) (armreservations.ReservationBillingPlan, error) { + switch strings.ToLower(strings.TrimSpace(paymentOption)) { + case "all-upfront", "upfront": + return armreservations.ReservationBillingPlanUpfront, nil + case "no-upfront", "monthly": + return armreservations.ReservationBillingPlanMonthly, nil + default: + return "", fmt.Errorf( + "azure reservations support only upfront or monthly billing; %q is not available (partial-upfront has no azure equivalent)", + paymentOption) + } +} + // apiVersion is the GA api-version for the Microsoft.Capacity Reservations API. // Pinned to 2022-11-01 — the last stable version before Azure introduced the // calculatePrice requirement for new SKU families. diff --git a/providers/azure/services/internal/reservations/purchase_test.go b/providers/azure/services/internal/reservations/purchase_test.go index e1a112148..e62abd98c 100644 --- a/providers/azure/services/internal/reservations/purchase_test.go +++ b/providers/azure/services/internal/reservations/purchase_test.go @@ -8,6 +8,7 @@ import ( "net/http" "testing" + "github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/reservations/armreservations" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" @@ -710,3 +711,40 @@ func TestParseTermYears(t *testing.T) { } } } + +// TestBillingPlanForPaymentOption pins the two-billing-plan Azure contract: +// Upfront and Monthly are the only members of armreservations. +// ReservationBillingPlan (constants.go), and Monthly costs the same total +// as Upfront (no partial-upfront exists). Both the converter's +// "upfront"/"monthly" vocabulary and the CLI/MCP's "all-upfront"/ +// "no-upfront" vocabulary must map onto the same two SDK enum values; every +// other input (including "partial-upfront", empty, and unrecognized +// strings) must be a hard error, never a silent default +// (feedback_no_silent_fallbacks). +func TestBillingPlanForPaymentOption(t *testing.T) { + tests := []struct { + paymentOption string + want armreservations.ReservationBillingPlan + wantErr bool + }{ + {"all-upfront", armreservations.ReservationBillingPlanUpfront, false}, + {"upfront", armreservations.ReservationBillingPlanUpfront, false}, + {"ALL-UPFRONT", armreservations.ReservationBillingPlanUpfront, false}, // case-insensitive + {" upfront ", armreservations.ReservationBillingPlanUpfront, false}, // whitespace-tolerant + {"no-upfront", armreservations.ReservationBillingPlanMonthly, false}, + {"monthly", armreservations.ReservationBillingPlanMonthly, false}, + {"partial-upfront", "", true}, // Azure has no partial-upfront equivalent + {"", "", true}, // empty must error, never silently default to Upfront + {"bogus", "", true}, + } + for _, tc := range tests { + got, err := BillingPlanForPaymentOption(tc.paymentOption) + if tc.wantErr { + assert.Error(t, err, "payment_option=%q should be an error", tc.paymentOption) + assert.Empty(t, got, "payment_option=%q error return should be empty", tc.paymentOption) + } else { + require.NoError(t, err, "payment_option=%q should not error", tc.paymentOption) + assert.Equal(t, tc.want, got, "payment_option=%q", tc.paymentOption) + } + } +}