Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 40 additions & 12 deletions internal/api/antigravity_cli.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package api

import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"io"
Expand Down Expand Up @@ -40,6 +42,9 @@ type agySession struct {
pty pty.Pty
cmd *pty.Cmd
conn *AntigravityConnection
// csrfToken is the token agy was launched with. Its language server
// rejects Connect-RPC calls without it (missing CSRF token, HTTP 401).
csrfToken string
}

// AntigravityCLIRunner manages a bounded warm agy process so onWatch can read
Expand Down Expand Up @@ -134,8 +139,8 @@ func (r *AntigravityCLIRunner) Fetch(ctx context.Context) (*AntigravitySnapshot,
return nil, err
}

base := r.sess.conn.BaseURL
summary, status, err := r.post(ctx, base, agyQuotaSummaryRPC)
conn := r.sess.conn
summary, status, err := r.post(ctx, conn, agyQuotaSummaryRPC)
if err != nil || status != http.StatusOK {
r.recordFailureLocked()
if err == nil {
Expand All @@ -158,7 +163,7 @@ func (r *AntigravityCLIRunner) Fetch(ctx context.Context) (*AntigravitySnapshot,
}

// Identity/plan are best-effort; the quota summary already succeeded.
if statusBody, st, serr := r.post(ctx, base, agyUserStatusRPC); serr == nil && st == http.StatusOK {
if statusBody, st, serr := r.post(ctx, conn, agyUserStatusRPC); serr == nil && st == http.StatusOK {
if us, perr := ParseAntigravityResponse(statusBody); perr == nil && us.UserStatus != nil {
snap.Email = us.UserStatus.Email
if us.UserStatus.PlanStatus != nil {
Expand Down Expand Up @@ -206,13 +211,33 @@ func (r *AntigravityCLIRunner) ensureLocked(ctx context.Context) error {
return nil
}

// newAgyCSRFToken returns a fresh random token for one managed agy launch.
func newAgyCSRFToken() (string, error) {
b := make([]byte, 16)
if _, err := rand.Read(b); err != nil {
return "", err
}
return hex.EncodeToString(b), nil
}

// agyLaunchArgs sets the CSRF token agy's language server will require. Without
// it, agy on Windows rejects every call as "missing CSRF token" (HTTP 401), so
// the session never becomes ready (issue #140).
func agyLaunchArgs(csrfToken string) []string {
return []string{"--csrf_token", csrfToken}
}

// launch starts agy inside a pseudo-terminal and drains its output.
func (r *AntigravityCLIRunner) launch(binPath string) (*agySession, error) {
token, err := newAgyCSRFToken()
if err != nil {
return nil, fmt.Errorf("antigravity cli: csrf token: %w", err)
}
p, err := pty.New()
if err != nil {
return nil, fmt.Errorf("antigravity cli: open pty: %w", err)
}
cmd := p.CommandContext(r.rootCtx, binPath)
cmd := p.CommandContext(r.rootCtx, binPath, agyLaunchArgs(token)...)
cmd.Env = append(os.Environ(), "TERM=xterm-256color")
if err := cmd.Start(); err != nil {
_ = p.Close()
Expand All @@ -221,7 +246,7 @@ func (r *AntigravityCLIRunner) launch(binPath string) (*agySession, error) {
// Drain PTY output so the process is not blocked on a full buffer.
go func() { _, _ = io.Copy(io.Discard, p) }()
r.logger.Debug("launched managed agy", "pid", cmd.Process.Pid, "path", binPath)
return &agySession{pty: p, cmd: cmd}, nil
return &agySession{pty: p, cmd: cmd, csrfToken: token}, nil
}

// awaitReady polls until the quota endpoint parses, since a fresh agy can bind
Expand All @@ -235,8 +260,8 @@ func (r *AntigravityCLIRunner) awaitReady(ctx context.Context, sess *agySession)
}
ports, err := r.client.discoverPorts(ctx, pid)
if err == nil && len(ports) > 0 {
if conn, _ := r.client.probeForConnectAPI(ctx, ports, ""); conn != nil {
if _, status, perr := r.post(ctx, conn.BaseURL, agyQuotaSummaryRPC); perr == nil && status == http.StatusOK {
if conn, _ := r.client.probeForConnectAPI(ctx, ports, sess.csrfToken); conn != nil {
if _, status, perr := r.post(ctx, conn, agyQuotaSummaryRPC); perr == nil && status == http.StatusOK {
return conn, nil
}
}
Expand All @@ -257,19 +282,22 @@ func (r *AntigravityCLIRunner) sessionHealthy(ctx context.Context) bool {
}
checkCtx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
_, status, err := r.post(checkCtx, r.sess.conn.BaseURL, agyQuotaSummaryRPC)
_, status, err := r.post(checkCtx, r.sess.conn, agyQuotaSummaryRPC)
return err == nil && status == http.StatusOK
}

// post issues a Connect-RPC POST against the agy language server. No CSRF token
// is required for the CLI's server.
func (r *AntigravityCLIRunner) post(ctx context.Context, baseURL, rpcPath string) ([]byte, int, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodPost, baseURL+rpcPath, strings.NewReader(agyMetadataBody))
// post issues a Connect-RPC POST against the agy language server, with the
// CSRF token the managed agy was launched with.
func (r *AntigravityCLIRunner) post(ctx context.Context, conn *AntigravityConnection, rpcPath string) ([]byte, int, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodPost, conn.BaseURL+rpcPath, strings.NewReader(agyMetadataBody))
if err != nil {
return nil, 0, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Connect-Protocol-Version", "1")
if conn.CSRFToken != "" {
req.Header.Set("X-Codeium-Csrf-Token", conn.CSRFToken)
}

resp, err := r.client.httpClient.Do(req)
if err != nil {
Expand Down
40 changes: 40 additions & 0 deletions internal/api/antigravity_cli_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package api

import (
"context"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
Expand Down Expand Up @@ -159,3 +161,41 @@ func TestResolveAgyPath_EnvMissingFileErrors(t *testing.T) {
t.Error("expected error when ANTIGRAVITY_CLI_PATH points at a missing file")
}
}

// agy's language server rejects Connect-RPC calls that lack the CSRF token it
// was launched with (issue #140). The runner must launch agy with a token and
// send the same token on every call.
func TestAgyRunner_SendsCSRFToken(t *testing.T) {
const token = "test-csrf-token"
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("X-Codeium-Csrf-Token") != token {
w.WriteHeader(http.StatusUnauthorized)
_, _ = w.Write([]byte(`{"code":"unauthenticated","message":"missing CSRF token"}`))
return
}
_, _ = w.Write([]byte(`{}`))
}))
defer srv.Close()

r := NewAntigravityCLIRunner(nil)
defer r.Stop()
_, status, err := r.post(context.Background(), &AntigravityConnection{BaseURL: srv.URL, CSRFToken: token}, agyQuotaSummaryRPC)
if err != nil || status != http.StatusOK {
t.Fatalf("post with token: status=%d err=%v", status, err)
}
}

func TestAgyLaunchArgs_PassCSRFToken(t *testing.T) {
a, err := newAgyCSRFToken()
if err != nil {
t.Fatal(err)
}
b, _ := newAgyCSRFToken()
if len(a) < 32 || a == b {
t.Fatalf("tokens must be long and unique per launch: %q %q", a, b)
}
args := agyLaunchArgs(a)
if len(args) != 2 || args[0] != "--csrf_token" || args[1] != a {
t.Fatalf("agyLaunchArgs(%q) = %q", a, args)
}
}
Loading