diff --git a/SECURITY.md b/SECURITY.md index 0afdd54..83f3fe1 100644 --- a/SECURITY.md +++ b/SECURITY.md @@ -44,7 +44,14 @@ What the server does contain: (`internal/tool/runner/env.go`). - Git credentials are supplied through short-lived `GIT_ASKPASS` scripts that are removed before the CLI starts. Tokens never reach the URL or - `.git/config`. + `.git/config`, and the script answers only for the host the session was + cloned from. +- Git commands the server runs in a workspace treat the workspace as untrusted: + `.git/config` is rebuilt from an allowlist first, settings that run programs + (fsmonitor, hooks, credential helpers, submodule recursion) are pinned on the + command line, the global git config is ignored, and the environment is + allowlisted. Fetches and pushes go to the URL the session was cloned from + (`internal/tool/git/safe.go`). - Webhook-triggered work from authors without write access is refused by default (`code_review.allow_untrusted_authors`). - Opening a pull request always requires an explicit action. This is a diff --git a/cmd/codeforge/main.go b/cmd/codeforge/main.go index cc9b716..44a4c30 100644 --- a/cmd/codeforge/main.go +++ b/cmd/codeforge/main.go @@ -270,6 +270,7 @@ func run() error { var webhookReceiverHandler *handlers.WebhookReceiverHandler if cfg.CodeReview.WebhookSecrets.GitHub != "" || cfg.CodeReview.WebhookSecrets.GitLab != "" { webhookReceiverHandler = handlers.NewWebhookReceiverHandler(sessionService, rdb, cfg.CodeReview, settings.NewStore(rdb)) + webhookReceiverHandler.SetGitLabMemberLookup(handlers.NewGitLabMemberLookup(keyResolver, cfg.CodeReview.DefaultKeyName)) } // Initialize tenant service and handler diff --git a/docs/configuration.md b/docs/configuration.md index 12c46de..4ec5400 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -113,6 +113,21 @@ instances and custom ports (e.g. `http://gitlab.example.com:8080`) work. |----------|---------|-------------| | `CODEFORGE_SUBSCRIPTION__ENABLED` | `false` | Enable the tenant subscription model. When disabled, only the static operator Bearer token is accepted and the per-session API-key (BYOK) flow is unchanged. When enabled, tenant API tokens (`cfk_...`) are also accepted and resolve to managed keys from the key pool. | +A tenant session runs with the credentials the tenant brings. The operator's +credentials never fill in for it: + +- Git access uses only the request's `access_token`. There is no fallback to + registered keys or `GITHUB_TOKEN`/`GITLAB_TOKEN`, so without a token only + public repositories can be cloned, and creating or updating a PR needs one. +- `provider_key` is rejected (`403`), and `repo_url` must be an `http(s)` URL. +- `config.tools` get only the config the tenant supplies; nothing is + auto-filled from registered keys. +- The operator's registered MCP servers are not added; only the session's own + `config.mcp_servers` are. +- `config.workspace_session_id` must name one of the tenant's own sessions. +- Session routes (`/sessions/{id}/...`) answer `404` for another tenant's + session. + ### Notifications Chat notifications for terminal session events. Disabled unless at least one webhook URL is set. @@ -169,7 +184,11 @@ By default CodeForge therefore requires write access from the author: point in the system. - **GitLab** payloads carry no equivalent of `author_association`, so fork MRs (source project ≠ target project) are skipped outright and commands are - refused on them. + refused on them. On other MRs a command runs only when its author has + Developer access (30) or higher on the project. CodeForge looks that up + through the GitLab members API with the `default_key_name` token, so the token + must be able to read project members; if the lookup fails, the command is + refused. Setting `allow_untrusted_authors: true` disables all of the above. Only do that where each session is genuinely isolated — see diff --git a/docs/deployment.md b/docs/deployment.md index 4cd8f19..c01a804 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -193,7 +193,14 @@ What CodeForge does to contain this today: encryption key, operator token, webhook secrets, and Redis URL never cross into the session (`internal/tool/runner/env.go`). - Git credentials are used via short-lived `GIT_ASKPASS` scripts that are - removed before the CLI starts; tokens are never in the URL or `.git/config`. + removed before the CLI starts; tokens are never in the URL or `.git/config`, + and the script answers only for the host the session was cloned from. +- Git commands the server runs in a workspace rebuild `.git/config` from an + allowlist, pin the settings that run programs on the command line, get an + allowlisted environment, and ignore the global git config (`~/.gitconfig`), + which the CLI could otherwise write. Configure TLS trust for a self-hosted + instance with `GIT_SSL_CAINFO` / `SSL_CERT_FILE` or the root-owned + `/etc/gitconfig`, not `~/.gitconfig`. - Webhook-triggered work from authors without write access is **off by default** (`code_review.allow_untrusted_authors`) — see [Configuration](configuration.md#who-can-trigger-a-webhook-review). diff --git a/internal/server/handlers/sessions.go b/internal/server/handlers/sessions.go index 2bf0048..b6faf26 100644 --- a/internal/server/handlers/sessions.go +++ b/internal/server/handlers/sessions.go @@ -7,6 +7,7 @@ import ( "fmt" "log/slog" "net/http" + "net/url" "os" "strconv" "strings" @@ -39,6 +40,12 @@ type tenantSessionCounter interface { CountActiveByTenant(ctx context.Context, tenantID string) (int, error) } +// sessionGetter loads a session by ID. Implemented by *session.Service; kept as +// an interface so handler tests can fake it. +type sessionGetter interface { + Get(ctx context.Context, sessionID string) (*session.Session, error) +} + // workspacePathResolver resolves a session's live workspace directory and // records workspace activity (Touch extends the TTL window on access). // Implemented by *workspace.Manager; kept as an interface so handler tests can fake it. @@ -57,6 +64,7 @@ type SessionHandler struct { domains gitpkg.DomainsSource // optional, nil = standard github.com/gitlab.com detection only tenantService *tenant.Service // optional, nil = subscription disabled sessionCounter tenantSessionCounter // optional, nil = concurrency limit not enforced + sessions sessionGetter // nil = tenants cannot reference other sessions workspaces workspacePathResolver // optional, nil = diff endpoint reports workspace missing } @@ -65,6 +73,7 @@ func NewSessionHandler(service *session.Service, prService *session.PRService, c h := &SessionHandler{service: service, prService: prService, canceller: canceller, cliRegistry: cliRegistry, keyRegistry: keyRegistry, domains: domains, tenantService: tenantService, workspaces: workspaces} if service != nil { h.sessionCounter = service + h.sessions = service } return h } @@ -204,25 +213,38 @@ func (h *SessionHandler) createSession(w http.ResponseWriter, r *http.Request, r // tenant_id); operator/no-tenant requests pass unconditionally. A mismatch returns // 404 (not 403) so a tenant cannot probe other tenants' session IDs. Routes without // a sessionID (List, Create) pass through and enforce their own scoping. +// +// The middleware reads {sessionID} with chi.URLParam, so it must be attached +// where chi has already matched that param (r.With on the route, or r.Use inside +// an r.Route("/{sessionID}", ...) subrouter). Attached with r.Use above the +// pattern, the param is still empty and every request passes unchecked. func (h *SessionHandler) OwnershipMiddleware(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - tnt := middleware.TenantFromContext(r.Context()) - sessionID := chi.URLParam(r, "sessionID") - if tnt == nil || sessionID == "" { + return SessionOwnership(h.service.Get)(next) +} + +// SessionOwnership builds the ownership check behind OwnershipMiddleware around a +// session lookup, so the check can be exercised without Redis. +func SessionOwnership(lookup func(ctx context.Context, sessionID string) (*session.Session, error)) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + tnt := middleware.TenantFromContext(r.Context()) + sessionID := chi.URLParam(r, "sessionID") + if tnt == nil || sessionID == "" { + next.ServeHTTP(w, r) + return + } + t, err := lookup(r.Context(), sessionID) + if err != nil { + writeAppError(w, err) + return + } + if t.TenantID != tnt.ID { + writeError(w, http.StatusNotFound, "session not found") + return + } next.ServeHTTP(w, r) - return - } - t, err := h.service.Get(r.Context(), sessionID) - if err != nil { - writeAppError(w, err) - return - } - if t.TenantID != tnt.ID { - writeError(w, http.StatusNotFound, "session not found") - return - } - next.ServeHTTP(w, r) - }) + }) + } } // applyTenant enforces a subscription tenant's tier limits and assigns a managed @@ -233,6 +255,10 @@ func (h *SessionHandler) applyTenant(ctx context.Context, req *session.CreateSes return 0, "" } + if status, msg := h.checkTenantSources(ctx, req, tnt); status != 0 { + return status, msg + } + cli := h.cliRegistry.DefaultCLI() if req.Config != nil && req.Config.CLI != "" { cli = req.Config.CLI @@ -308,6 +334,31 @@ func (h *SessionHandler) applyTenant(ctx context.Context, req *session.CreateSes return 0, "" } +// checkTenantSources rejects request fields that would let a tenant session run +// with access the tenant did not bring: one of the operator's registered keys +// (provider_key), a repository reached through the server's filesystem +// (file:// and other non-HTTP URLs), or another tenant's workspace. The +// executor separately keeps tenant sessions off the operator's fallback +// credentials (see session.Session.UsesOperatorCredentials). +func (h *SessionHandler) checkTenantSources(ctx context.Context, req *session.CreateSessionRequest, tnt *tenant.Tenant) (int, string) { + if req.ProviderKey != "" { + return http.StatusForbidden, "provider_key is not available to subscription tenants; pass access_token instead" + } + if u, err := url.Parse(req.RepoURL); err != nil || (u.Scheme != "https" && u.Scheme != "http") || u.Host == "" { + return http.StatusBadRequest, "repo_url must be an http(s) URL" + } + if req.Config != nil && req.Config.WorkspaceSessionID != "" { + if h.sessions == nil { + return http.StatusNotFound, "workspace session not found" + } + ref, err := h.sessions.Get(ctx, req.Config.WorkspaceSessionID) + if err != nil || ref.TenantID != tnt.ID { + return http.StatusNotFound, "workspace session not found" + } + } + return 0, "" +} + // stringInJSONList reports whether target is allowed by a JSON array allow-list like // `["claude-code","codex"]`. An empty/whitespace list means "no restriction" (allow). // A NON-empty but malformed list fails CLOSED (deny) — a corrupt restriction must not diff --git a/internal/server/handlers/subscription_test.go b/internal/server/handlers/subscription_test.go index 84519b8..0f56adc 100644 --- a/internal/server/handlers/subscription_test.go +++ b/internal/server/handlers/subscription_test.go @@ -8,6 +8,7 @@ import ( _ "modernc.org/sqlite" + "github.com/freema/codeforge/internal/apperror" "github.com/freema/codeforge/internal/crypto" "github.com/freema/codeforge/internal/database" "github.com/freema/codeforge/internal/session" @@ -60,6 +61,10 @@ func testCLIRegistry() *runner.Registry { return reg } +// tenantRepo is a repository URL a tenant may use (applyTenant runs after the +// request passed validation, so it always has one). +const tenantRepo = "https://github.com/acme/repo.git" + type fakeCounter struct{ active int } func (f fakeCounter) CountActiveByTenant(_ context.Context, _ string) (int, error) { @@ -75,13 +80,13 @@ func TestApplyTenant_ConcurrencyLimit(t *testing.T) { h := NewSessionHandler(nil, nil, nil, testCLIRegistry(), nil, nil, svc, nil) h.sessionCounter = fakeCounter{active: tnt.MaxConcurrentSessions} - if status, _ := h.applyTenant(ctx, &session.CreateSessionRequest{}, tnt); status != 429 { + if status, _ := h.applyTenant(ctx, &session.CreateSessionRequest{RepoURL: tenantRepo}, tnt); status != 429 { t.Fatalf("at concurrency limit: status = %d, want 429", status) } // Under the limit, a BYOK request passes (no pool needed). h.sessionCounter = fakeCounter{active: tnt.MaxConcurrentSessions - 1} - req := &session.CreateSessionRequest{Config: &session.Config{AIApiKey: "byok"}} + req := &session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{AIApiKey: "byok"}} if status, msg := h.applyTenant(ctx, req, tnt); status != 0 { t.Fatalf("under concurrency limit: status = %d (%s), want 0", status, msg) } @@ -111,7 +116,7 @@ func TestApplyTenant(t *testing.T) { h := NewSessionHandler(nil, nil, nil, testCLIRegistry(), nil, nil, svc, nil) t.Run("disallowed CLI -> 403", func(t *testing.T) { - req := &session.CreateSessionRequest{Config: &session.Config{CLI: "cursor"}} + req := &session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{CLI: "cursor"}} status, _ := h.applyTenant(ctx, req, tnt) if status != 403 { t.Fatalf("status = %d, want 403", status) @@ -119,7 +124,7 @@ func TestApplyTenant(t *testing.T) { }) t.Run("allowed CLI, no BYOK -> pool key assigned + tenant_id stamped + budget capped", func(t *testing.T) { - req := &session.CreateSessionRequest{} + req := &session.CreateSessionRequest{RepoURL: tenantRepo} status, msg := h.applyTenant(ctx, req, tnt) if status != 0 { t.Fatalf("status = %d (%s), want 0", status, msg) @@ -136,7 +141,7 @@ func TestApplyTenant(t *testing.T) { }) t.Run("BYOK key preserved, pool not consulted", func(t *testing.T) { - req := &session.CreateSessionRequest{Config: &session.Config{AIApiKey: "my-own-key"}} + req := &session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{AIApiKey: "my-own-key"}} status, _ := h.applyTenant(ctx, req, tnt) if status != 0 { t.Fatalf("status = %d, want 0", status) @@ -152,10 +157,65 @@ func TestApplyTenant(t *testing.T) { for i := 0; i < lt.MaxSessionsPerDay; i++ { _ = store.LogUsage(ctx, &tenant.UsageLog{TenantID: lt.ID, SessionID: string(rune('a' + i)), CLI: "claude-code"}) } - req := &session.CreateSessionRequest{} + req := &session.CreateSessionRequest{RepoURL: tenantRepo} status, _ := h.applyTenant(ctx, req, lt) if status != 429 { t.Fatalf("status = %d, want 429 after hitting daily limit", status) } }) } + +type fakeSessions map[string]*session.Session + +func (f fakeSessions) Get(_ context.Context, id string) (*session.Session, error) { + if t, ok := f[id]; ok { + return t, nil + } + return nil, apperror.NotFound("session %s not found", id) +} + +// TestApplyTenant_RejectsForeignCredentials covers request fields that would +// let a tenant session run with the operator's or another tenant's access. +func TestApplyTenant_RejectsForeignCredentials(t *testing.T) { + ctx := context.Background() + svc, store, _ := newTenantService(t) + res, err := svc.CreateTenant(ctx, "acme", "acme", tenant.TierFree) + if err != nil { + t.Fatalf("create tenant: %v", err) + } + tnt, err := store.GetTenant(ctx, res.Tenant.ID) + if err != nil { + t.Fatalf("get tenant: %v", err) + } + + h := NewSessionHandler(nil, nil, nil, testCLIRegistry(), nil, nil, svc, nil) + h.sessions = fakeSessions{ + "own": {ID: "own", TenantID: tnt.ID}, + "other": {ID: "other", TenantID: "another-tenant"}, + "operator": {ID: "operator"}, + } + + tests := []struct { + name string + req session.CreateSessionRequest + wantStatus int + }{ + {"operator provider key", session.CreateSessionRequest{RepoURL: tenantRepo, ProviderKey: "operator-github"}, 403}, + {"file URL", session.CreateSessionRequest{RepoURL: "file:///data/workspaces/other/"}, 400}, + {"ssh URL", session.CreateSessionRequest{RepoURL: "ssh://git@github.com/acme/repo.git"}, 400}, + {"another tenant's workspace", session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{WorkspaceSessionID: "other"}}, 404}, + {"an operator session's workspace", session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{WorkspaceSessionID: "operator"}}, 404}, + {"unknown workspace", session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{WorkspaceSessionID: "missing"}}, 404}, + // BYOK keeps the pool out of these so only the checks above decide. + {"own workspace", session.CreateSessionRequest{RepoURL: tenantRepo, Config: &session.Config{WorkspaceSessionID: "own", AIApiKey: "byok"}}, 0}, + {"own access token", session.CreateSessionRequest{RepoURL: tenantRepo, AccessToken: "ghp_tenant", Config: &session.Config{AIApiKey: "byok"}}, 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := tt.req + if status, msg := h.applyTenant(ctx, &req, tnt); status != tt.wantStatus { + t.Fatalf("status = %d (%s), want %d", status, msg, tt.wantStatus) + } + }) + } +} diff --git a/internal/server/handlers/webhook_receiver.go b/internal/server/handlers/webhook_receiver.go index c0edb1f..9fae814 100644 --- a/internal/server/handlers/webhook_receiver.go +++ b/internal/server/handlers/webhook_receiver.go @@ -7,6 +7,7 @@ import ( "crypto/subtle" "encoding/hex" "encoding/json" + "errors" "fmt" "io" "log/slog" @@ -18,6 +19,7 @@ import ( "github.com/freema/codeforge/internal/redisclient" "github.com/freema/codeforge/internal/session" "github.com/freema/codeforge/internal/settings" + gitpkg "github.com/freema/codeforge/internal/tool/git" ) // defaultReviewCLI is the CLI used for webhook-triggered reviews when no @@ -29,7 +31,41 @@ type WebhookReceiverHandler struct { sessionService *session.Service redis *redisclient.Client cfg config.CodeReviewConfig - settings *settings.Store // optional, nil = runtime overrides disabled + settings *settings.Store // optional, nil = runtime overrides disabled + gitlabMembers GitLabMemberLookup // nil = GitLab note commands are refused +} + +// GitLabMemberLookup returns a GitLab user's effective access level on a +// project, or 0 when the user is not a member. +type GitLabMemberLookup func(ctx context.Context, repoURL string, projectID, userID int) (int, error) + +// webhookTokenResolver is the part of keys.Resolver the member lookup needs. +type webhookTokenResolver interface { + ResolveToken(ctx context.Context, repoURL, accessToken, providerKey string) (string, error) +} + +// NewGitLabMemberLookup asks the GitLab instance hosting repoURL for the user's +// access level, authenticating with the token resolved for keyName (the key +// webhook-triggered sessions run with). +func NewGitLabMemberLookup(tokens webhookTokenResolver, keyName string) GitLabMemberLookup { + return func(ctx context.Context, repoURL string, projectID, userID int) (int, error) { + repo, err := gitpkg.ParseRepoURL(repoURL, nil) + if err != nil { + return 0, err + } + token, err := tokens.ResolveToken(ctx, repoURL, "", keyName) + if err != nil { + return 0, err + } + return gitpkg.GitLabAccessLevel(ctx, repo.BaseURL(), token, projectID, userID) + } +} + +// SetGitLabMemberLookup enables GitLab note commands. Without a lookup the +// commenter's access cannot be verified and every command is refused (unless +// code_review.allow_untrusted_authors is set). +func (h *WebhookReceiverHandler) SetGitLabMemberLookup(fn GitLabMemberLookup) { + h.gitlabMembers = fn } // NewWebhookReceiverHandler creates a new webhook receiver handler. @@ -169,6 +205,7 @@ type gitlabNoteEvent struct { Note string `json:"note"` NoteableType string `json:"noteable_type"` //nolint:misspell // GitLab API uses "noteable_type" System bool `json:"system"` + ProjectID int `json:"project_id"` } `json:"object_attributes"` MergeRequest *struct { IID int `json:"iid"` @@ -178,14 +215,32 @@ type gitlabNoteEvent struct { TargetProjectID int `json:"target_project_id"` } `json:"merge_request"` User struct { + ID int `json:"id"` Username string `json:"username"` } `json:"user"` Project struct { + ID int `json:"id"` PathWithNamespace string `json:"path_with_namespace"` HTTPURLToRepo string `json:"http_url_to_repo"` } `json:"project"` } +// gitlabNoteAuthorAccess looks up the note author's access level on the project the +// note was posted in. +func (h *WebhookReceiverHandler) gitlabNoteAuthorAccess(ctx context.Context, e *gitlabNoteEvent) (int, error) { + if h.gitlabMembers == nil { + return 0, errors.New("member lookup not configured") + } + projectID := e.Project.ID + if projectID == 0 { + projectID = e.ObjectAttributes.ProjectID + } + if projectID == 0 || e.User.ID == 0 { + return 0, errors.New("note event carries no project or user id") + } + return h.gitlabMembers(ctx, e.Project.HTTPURLToRepo, projectID, e.User.ID) +} + // --- GitLab webhook types --- type gitlabMREvent struct { @@ -676,9 +731,8 @@ func (h *WebhookReceiverHandler) handleGitLabNote(w http.ResponseWriter, r *http } // Same reasoning as the GitHub comment path: /fix turns the note body into - // the prompt of a code-writing session. GitLab does not report the - // commenter's access level here, so the only provenance signal available - // is whether the MR itself comes from a fork. + // the prompt of a code-writing session. Commands on fork MRs are refused + // outright, and on any MR the author must be able to push to the project. if isGitLabFork(event.MergeRequest.SourceProjectID, event.MergeRequest.TargetProjectID) && !h.cfg.AllowUntrustedAuthors { log.Warn("gitlab webhook: ignoring forge command on fork MR", "command", cmd, @@ -693,6 +747,29 @@ func (h *WebhookReceiverHandler) handleGitLabNote(w http.ResponseWriter, r *http return } + // Note hooks do not carry the author's role, so it is looked up through the + // API: Developer (30) or above, the GitLab counterpart of GitHub's + // OWNER/MEMBER/COLLABORATOR. Anything that prevents confirming it — no + // lookup, missing IDs, an API error — drops the command. + if !h.cfg.AllowUntrustedAuthors { + level, err := h.gitlabNoteAuthorAccess(r.Context(), &event) + if err != nil || level < gitpkg.GitLabDeveloperAccess { + log.Warn("gitlab webhook: ignoring forge command from author without push access", + "command", cmd, + "mr", event.MergeRequest.IID, + "repo", event.Project.PathWithNamespace, + "user", event.User.Username, + "access_level", level, + "error", err, + ) + writeJSON(w, http.StatusOK, map[string]string{ + "status": "ignored", + "reason": "commenter has no developer access (set code_review.allow_untrusted_authors to override)", + }) + return + } + } + repoURL := event.Project.HTTPURLToRepo mrIID := event.MergeRequest.IID diff --git a/internal/server/handlers/webhook_receiver_test.go b/internal/server/handlers/webhook_receiver_test.go index 492b675..ae66e88 100644 --- a/internal/server/handlers/webhook_receiver_test.go +++ b/internal/server/handlers/webhook_receiver_test.go @@ -1,10 +1,12 @@ package handlers import ( + "context" "crypto/hmac" "crypto/sha256" "encoding/hex" "encoding/json" + "errors" "net/http" "net/http/httptest" "strings" @@ -489,3 +491,85 @@ func TestWebhookUntrustedAuthorGating(t *testing.T) { }) } } + +type staticTokenResolver struct{ token string } + +func (r staticTokenResolver) ResolveToken(context.Context, string, string, string) (string, error) { + if r.token == "" { + return "", errors.New("no token") + } + return r.token, nil +} + +// TestGitLabNoteCommandRequiresDeveloperAccess checks that MR note commands are +// dispatched only for authors with Developer access or above, as reported by +// a local stand-in for the GitLab members API. +func TestGitLabNoteCommandRequiresDeveloperAccess(t *testing.T) { + const secret = "gl-secret" + + tests := []struct { + name string + apiStatus int + apiBody string + noLookup bool + noUserID bool + allowUntrusted bool + wantStatus int + wantBodySubstr string + }{ + {name: "reporter is refused", apiStatus: http.StatusOK, apiBody: `{"access_level":20,"state":"active"}`, wantStatus: http.StatusOK, wantBodySubstr: "no developer access"}, + {name: "guest is refused", apiStatus: http.StatusOK, apiBody: `{"access_level":10,"state":"active"}`, wantStatus: http.StatusOK, wantBodySubstr: "no developer access"}, + {name: "non-member is refused", apiStatus: http.StatusNotFound, apiBody: `{"message":"404 Not found"}`, wantStatus: http.StatusOK, wantBodySubstr: "no developer access"}, + {name: "API failure is refused", apiStatus: http.StatusInternalServerError, apiBody: `{}`, wantStatus: http.StatusOK, wantBodySubstr: "no developer access"}, + {name: "missing lookup is refused", noLookup: true, wantStatus: http.StatusOK, wantBodySubstr: "no developer access"}, + {name: "missing user id is refused", noUserID: true, apiStatus: http.StatusOK, apiBody: `{"access_level":40,"state":"active"}`, wantStatus: http.StatusOK, wantBodySubstr: "no developer access"}, + // Past the gate the handler stops at the unset default key, which + // proves the command was let through without needing a session store. + {name: "developer passes", apiStatus: http.StatusOK, apiBody: `{"access_level":30,"state":"active"}`, wantStatus: http.StatusBadRequest, wantBodySubstr: "default_key_name"}, + {name: "maintainer passes", apiStatus: http.StatusOK, apiBody: `{"access_level":40,"state":"active"}`, wantStatus: http.StatusBadRequest, wantBodySubstr: "default_key_name"}, + {name: "allow_untrusted_authors skips the lookup", noLookup: true, allowUntrusted: true, wantStatus: http.StatusBadRequest, wantBodySubstr: "default_key_name"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gitlab := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/v4/projects/42/members/all/7" { + t.Errorf("member lookup path = %s", r.URL.Path) + } + w.WriteHeader(tt.apiStatus) + _, _ = w.Write([]byte(tt.apiBody)) + })) + defer gitlab.Close() + + userID := `"id":7,` + if tt.noUserID { + userID = "" + } + body := `{"object_kind":"note","object_attributes":{"note":"/fix do something","noteable_type":"MergeRequest","system":false,"project_id":42},` + //nolint:misspell // GitLab API uses "noteable_type" + `"merge_request":{"iid":5,"source_branch":"feat","target_branch":"main","source_project_id":42,"target_project_id":42},` + + `"user":{` + userID + `"username":"someone"},` + + `"project":{"id":42,"path_with_namespace":"group/repo","http_url_to_repo":"` + gitlab.URL + `/group/repo.git"}}` + + h := &WebhookReceiverHandler{cfg: config.CodeReviewConfig{ + WebhookSecrets: config.WebhookSecretsConfig{GitLab: secret}, + AllowUntrustedAuthors: tt.allowUntrusted, + }} + if !tt.noLookup { + h.SetGitLabMemberLookup(NewGitLabMemberLookup(staticTokenResolver{token: "tok"}, "my-key")) + } + + req := httptest.NewRequest(http.MethodPost, "/api/v1/webhooks/gitlab", strings.NewReader(body)) + req.Header.Set("X-Gitlab-Token", secret) + req.Header.Set("X-Gitlab-Event", "Note Hook") + w := httptest.NewRecorder() + h.GitLabWebhook(w, req) + + if w.Code != tt.wantStatus { + t.Errorf("status = %d, want %d (body %q)", w.Code, tt.wantStatus, w.Body.String()) + } + if got := w.Body.String(); !strings.Contains(got, tt.wantBodySubstr) { + t.Errorf("body = %q, want substring %q", got, tt.wantBodySubstr) + } + }) + } +} diff --git a/internal/server/routes.go b/internal/server/routes.go new file mode 100644 index 0000000..5ee3ffa --- /dev/null +++ b/internal/server/routes.go @@ -0,0 +1,50 @@ +package server + +import ( + "net/http" + + "github.com/go-chi/chi/v5" +) + +// sessionRoutes is the part of *handlers.SessionHandler served under /sessions. +type sessionRoutes interface { + List(http.ResponseWriter, *http.Request) + Create(http.ResponseWriter, *http.Request) + Get(http.ResponseWriter, *http.Request) + Instruct(http.ResponseWriter, *http.Request) + Cancel(http.ResponseWriter, *http.Request) + Review(http.ResponseWriter, *http.Request) + PostReviewComments(http.ResponseWriter, *http.Request) + CreatePR(http.ResponseWriter, *http.Request) + PushToPR(http.ResponseWriter, *http.Request) + GetPRStatus(http.ResponseWriter, *http.Request) + Diff(http.ResponseWriter, *http.Request) +} + +// mountSessionRoutes registers the /sessions routes. The ownership middleware is +// attached inside the /{sessionID} subrouter: chi matches URL params while +// routing, so middleware registered with r.Use above the pattern would see an +// empty {sessionID} and let every tenant through. +func mountSessionRoutes(r chi.Router, h sessionRoutes, ownership, rateLimitMw func(http.Handler) http.Handler) { + r.Route("/sessions", func(r chi.Router) { + r.Get("/", h.List) + if rateLimitMw != nil { + r.With(rateLimitMw).Post("/", h.Create) + } else { + r.Post("/", h.Create) + } + + r.Route("/{sessionID}", func(r chi.Router) { + r.Use(ownership) // tenant may touch only its own sessions + r.Get("/", h.Get) + r.Post("/instruct", h.Instruct) + r.Post("/cancel", h.Cancel) + r.Post("/review", h.Review) + r.Post("/post-review", h.PostReviewComments) + r.Post("/create-pr", h.CreatePR) + r.Post("/push", h.PushToPR) + r.Get("/pr-status", h.GetPRStatus) + r.Get("/diff", h.Diff) + }) + }) +} diff --git a/internal/server/routes_test.go b/internal/server/routes_test.go new file mode 100644 index 0000000..931919d --- /dev/null +++ b/internal/server/routes_test.go @@ -0,0 +1,124 @@ +package server + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/go-chi/chi/v5" + + "github.com/freema/codeforge/internal/apperror" + "github.com/freema/codeforge/internal/server/handlers" + "github.com/freema/codeforge/internal/server/middleware" + "github.com/freema/codeforge/internal/session" + "github.com/freema/codeforge/internal/tenant" +) + +// recordingSessionRoutes answers every route with 200 and remembers that a +// handler was reached. +type recordingSessionRoutes struct{ reached bool } + +func (f *recordingSessionRoutes) hit(w http.ResponseWriter, _ *http.Request) { + f.reached = true + w.WriteHeader(http.StatusOK) +} + +func (f *recordingSessionRoutes) List(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) Create(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) Get(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) Instruct(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) Cancel(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) Review(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) PostReviewComments(w http.ResponseWriter, r *http.Request) { + f.hit(w, r) +} +func (f *recordingSessionRoutes) CreatePR(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) PushToPR(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) GetPRStatus(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } +func (f *recordingSessionRoutes) Diff(w http.ResponseWriter, r *http.Request) { f.hit(w, r) } + +var sessionIDRoutes = []struct{ method, path string }{ + {http.MethodGet, "/sessions/sess-a"}, + {http.MethodPost, "/sessions/sess-a/instruct"}, + {http.MethodPost, "/sessions/sess-a/cancel"}, + {http.MethodPost, "/sessions/sess-a/review"}, + {http.MethodPost, "/sessions/sess-a/post-review"}, + {http.MethodPost, "/sessions/sess-a/create-pr"}, + {http.MethodPost, "/sessions/sess-a/push"}, + {http.MethodGet, "/sessions/sess-a/pr-status"}, + {http.MethodGet, "/sessions/sess-a/diff"}, +} + +// newSessionRouter mounts the session routes the way the server does, with the +// real ownership check over an in-memory lookup that knows one session owned by +// tenant-a. caller is put in the request context as the authenticated tenant +// (nil means operator). +func newSessionRouter(caller *tenant.Tenant, h sessionRoutes) http.Handler { + lookup := func(_ context.Context, id string) (*session.Session, error) { + if id == "sess-a" { + return &session.Session{ID: id, TenantID: "tenant-a"}, nil + } + return nil, apperror.NotFound("session %s not found", id) + } + r := chi.NewRouter() + r.Use(func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + if caller != nil { + req = req.WithContext(middleware.ContextWithTenant(req.Context(), caller)) + } + next.ServeHTTP(w, req) + }) + }) + mountSessionRoutes(r, h, handlers.SessionOwnership(lookup), nil) + return r +} + +func TestSessionRoutes_TenantCannotReachForeignSession(t *testing.T) { + for _, rt := range sessionIDRoutes { + t.Run(rt.method+" "+rt.path, func(t *testing.T) { + h := &recordingSessionRoutes{} + rec := httptest.NewRecorder() + newSessionRouter(&tenant.Tenant{ID: "tenant-b"}, h).ServeHTTP(rec, httptest.NewRequest(rt.method, rt.path, nil)) + + if rec.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", rec.Code) + } + if h.reached { + t.Error("handler ran for another tenant's session") + } + }) + } +} + +func TestSessionRoutes_OwnerAndOperatorReachSession(t *testing.T) { + callers := map[string]*tenant.Tenant{ + "owner": {ID: "tenant-a"}, + "operator": nil, + } + for name, caller := range callers { + for _, rt := range sessionIDRoutes { + t.Run(name+" "+rt.method+" "+rt.path, func(t *testing.T) { + h := &recordingSessionRoutes{} + rec := httptest.NewRecorder() + newSessionRouter(caller, h).ServeHTTP(rec, httptest.NewRequest(rt.method, rt.path, nil)) + + if rec.Code != http.StatusOK || !h.reached { + t.Errorf("status = %d, reached = %v; want 200 and handler reached", rec.Code, h.reached) + } + }) + } + } +} + +func TestSessionRoutes_ListAndCreateSkipOwnership(t *testing.T) { + for _, method := range []string{http.MethodGet, http.MethodPost} { + h := &recordingSessionRoutes{} + rec := httptest.NewRecorder() + newSessionRouter(&tenant.Tenant{ID: "tenant-b"}, h).ServeHTTP(rec, httptest.NewRequest(method, "/sessions", nil)) + + if rec.Code != http.StatusOK || !h.reached { + t.Errorf("%s /sessions: status = %d, reached = %v", method, rec.Code, h.reached) + } + } +} diff --git a/internal/server/server.go b/internal/server/server.go index fc9dd20..4760b1c 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -110,24 +110,7 @@ func New(cfg *config.Config, redis *redisclient.Client, sqliteDB *database.DB, s r.Group(func(r chi.Router) { r.Use(chimw.Timeout(60 * time.Second)) - r.Route("/sessions", func(r chi.Router) { - r.Use(sessionHandler.OwnershipMiddleware) // tenant may touch only its own {sessionID} routes - r.Get("/", sessionHandler.List) - if rateLimitMw != nil { - r.With(rateLimitMw).Post("/", sessionHandler.Create) - } else { - r.Post("/", sessionHandler.Create) - } - r.Get("/{sessionID}", sessionHandler.Get) - r.Post("/{sessionID}/instruct", sessionHandler.Instruct) - r.Post("/{sessionID}/cancel", sessionHandler.Cancel) - r.Post("/{sessionID}/review", sessionHandler.Review) - r.Post("/{sessionID}/post-review", sessionHandler.PostReviewComments) - r.Post("/{sessionID}/create-pr", sessionHandler.CreatePR) - r.Post("/{sessionID}/push", sessionHandler.PushToPR) - r.Get("/{sessionID}/pr-status", sessionHandler.GetPRStatus) - r.Get("/{sessionID}/diff", sessionHandler.Diff) - }) + mountSessionRoutes(r, sessionHandler, sessionHandler.OwnershipMiddleware, rateLimitMw) r.Get("/session-types", sessionHandler.ListSessionTypes) diff --git a/internal/session/model.go b/internal/session/model.go index 1d0c3d6..4d6476d 100644 --- a/internal/session/model.go +++ b/internal/session/model.go @@ -75,6 +75,7 @@ type Session struct { // Subscription tenant that owns this session (empty = operator/BYOK). // Set server-side from the authenticated tenant, never from client input. + // See UsesOperatorCredentials. TenantID string `json:"tenant_id,omitempty"` // Observability @@ -242,3 +243,12 @@ func UnmarshalUsageInfo(data string) *UsageInfo { } return &u } + +// UsesOperatorCredentials reports whether the server may fill in credentials +// the session did not bring itself: registered provider keys, the +// GITHUB_TOKEN/GITLAB_TOKEN fallback, tool config auto-filled from provider +// keys, and the operator's registered MCP servers. Sessions of a subscription +// tenant may not — they run with what the tenant supplied. +func (s *Session) UsesOperatorCredentials() bool { + return s.TenantID == "" +} diff --git a/internal/session/pr_service.go b/internal/session/pr_service.go index e3b2830..324ee93 100644 --- a/internal/session/pr_service.go +++ b/internal/session/pr_service.go @@ -2,6 +2,7 @@ package session import ( "context" + "errors" "fmt" "log/slog" "path/filepath" @@ -56,6 +57,28 @@ func NewPRService(sessionService *Service, analyzer *runner.Analyzer, workspaceR return svc } +// errTenantNeedsAccessToken is returned when a tenant session without its own +// access token reaches an operation that needs the provider API. +var errTenantNeedsAccessToken = errors.New("the session has no access_token, and subscription sessions do not use the operator's provider credentials") + +// resolveAccessToken fills t.AccessToken from the key registry or the +// GITHUB_TOKEN/GITLAB_TOKEN fallback when the session brought none. Tenant +// sessions never fall back to the operator's credentials. +func (s *PRService) resolveAccessToken(ctx context.Context, t *Session) error { + if t.AccessToken != "" || s.tokenResolver == nil { + return nil + } + if !t.UsesOperatorCredentials() { + return errTenantNeedsAccessToken + } + token, err := s.tokenResolver.ResolveToken(ctx, t.RepoURL, "", t.ProviderKey) + if err != nil { + return err + } + t.AccessToken = token + return nil +} + // CreatePRRequest is the request body for POST /sessions/:id/create-pr. type CreatePRRequest struct { Title string `json:"title,omitempty"` @@ -103,12 +126,8 @@ func (s *PRService) CreatePR(ctx context.Context, sessionID string, req CreatePR } // Resolve access token (inline → registry → env) if not already set. - if s.tokenResolver != nil && t.AccessToken == "" { - token, err := s.tokenResolver.ResolveToken(ctx, t.RepoURL, t.AccessToken, t.ProviderKey) - if err != nil { - return nil, fmt.Errorf("resolving access token for PR: %w", err) - } - t.AccessToken = token + if err := s.resolveAccessToken(ctx, t); err != nil { + return nil, fmt.Errorf("resolving access token for PR: %w", err) } // Remember previous status so we can revert on non-fatal errors @@ -180,6 +199,7 @@ func (s *PRService) CreatePR(ctx context.Context, sessionID string, req CreatePR // Create branch, commit, push err = gitpkg.CreateBranchAndPush(ctx, gitpkg.BranchOptions{ WorkDir: workDir, + RepoURL: t.RepoURL, BranchName: branchName, BaseBranch: baseBranch, CommitMsg: commitMsg, @@ -266,12 +286,8 @@ func (s *PRService) PushToPR(ctx context.Context, sessionID string) (*PushToPRRe } // Resolve access token if not already set - if s.tokenResolver != nil && t.AccessToken == "" { - token, err := s.tokenResolver.ResolveToken(ctx, t.RepoURL, t.AccessToken, t.ProviderKey) - if err != nil { - return nil, fmt.Errorf("resolving access token for push: %w", err) - } - t.AccessToken = token + if err := s.resolveAccessToken(ctx, t); err != nil { + return nil, fmt.Errorf("resolving access token for push: %w", err) } // Generate commit message — try AI, fall back to generic @@ -287,6 +303,7 @@ func (s *PRService) PushToPR(ctx context.Context, sessionID string) (*PushToPRRe // Stage, commit, and push to existing branch if err := gitpkg.CommitAndPushToExisting(ctx, gitpkg.PushExistingOptions{ WorkDir: workDir, + RepoURL: t.RepoURL, BranchName: t.Branch, CommitMsg: commitMsg, AuthorName: s.cfg.CommitAuthor, @@ -336,12 +353,8 @@ func (s *PRService) GetPRStatus(ctx context.Context, sessionID string) (*gitpkg. } // Resolve token - if s.tokenResolver != nil && t.AccessToken == "" { - token, resolveErr := s.tokenResolver.ResolveToken(ctx, t.RepoURL, t.AccessToken, t.ProviderKey) - if resolveErr != nil { - return nil, fmt.Errorf("resolving token: %w", resolveErr) - } - t.AccessToken = token + if err := s.resolveAccessToken(ctx, t); err != nil { + return nil, fmt.Errorf("resolving token: %w", err) } status, err := gitpkg.GetPRStatus(ctx, repoInfo, t.AccessToken, t.PRNumber) diff --git a/internal/session/pr_service_test.go b/internal/session/pr_service_test.go new file mode 100644 index 0000000..9da7029 --- /dev/null +++ b/internal/session/pr_service_test.go @@ -0,0 +1,56 @@ +package session + +import ( + "context" + "errors" + "testing" +) + +// recordingTokenResolver stands in for the operator's key registry and +// GITHUB_TOKEN/GITLAB_TOKEN fallback. +type recordingTokenResolver struct{ calls int } + +func (r *recordingTokenResolver) ResolveToken(context.Context, string, string, string) (string, error) { + r.calls++ + return "operator-token", nil +} + +func TestResolveAccessToken(t *testing.T) { + ctx := context.Background() + + t.Run("tenant session without a token gets no operator token", func(t *testing.T) { + res := &recordingTokenResolver{} + s := &PRService{tokenResolver: res} + sess := &Session{RepoURL: "https://github.com/acme/repo.git", TenantID: "acme", ProviderKey: "operator-github"} + if err := s.resolveAccessToken(ctx, sess); !errors.Is(err, errTenantNeedsAccessToken) { + t.Fatalf("err = %v, want errTenantNeedsAccessToken", err) + } + if sess.AccessToken != "" || res.calls != 0 { + t.Errorf("token = %q, resolver calls = %d; want neither", sess.AccessToken, res.calls) + } + }) + + t.Run("tenant session keeps its own token", func(t *testing.T) { + res := &recordingTokenResolver{} + s := &PRService{tokenResolver: res} + sess := &Session{TenantID: "acme", AccessToken: "tenant-token"} + if err := s.resolveAccessToken(ctx, sess); err != nil { + t.Fatal(err) + } + if sess.AccessToken != "tenant-token" || res.calls != 0 { + t.Errorf("token = %q, resolver calls = %d", sess.AccessToken, res.calls) + } + }) + + t.Run("operator session falls back to the resolver", func(t *testing.T) { + res := &recordingTokenResolver{} + s := &PRService{tokenResolver: res} + sess := &Session{RepoURL: "https://github.com/acme/repo.git"} + if err := s.resolveAccessToken(ctx, sess); err != nil { + t.Fatal(err) + } + if sess.AccessToken != "operator-token" { + t.Errorf("token = %q, want operator-token", sess.AccessToken) + } + }) +} diff --git a/internal/tool/git/branch.go b/internal/tool/git/branch.go index 85ae8f7..e0b5209 100644 --- a/internal/tool/git/branch.go +++ b/internal/tool/git/branch.go @@ -5,13 +5,13 @@ import ( "fmt" "log/slog" "os" - "os/exec" "strings" ) // BranchOptions configures branch creation and push. type BranchOptions struct { WorkDir string + RepoURL string // the URL the workspace was cloned from; pushed to, whatever origin says BranchName string BaseBranch string // If set, the feature branch is created from origin/ instead of the current HEAD. CommitMsg string @@ -24,6 +24,9 @@ type BranchOptions struct { // Token is passed via GIT_ASKPASS (never in URL or .git/config). func CreateBranchAndPush(ctx context.Context, opts BranchOptions) error { workDir := opts.WorkDir + if err := SanitizeRepoConfig(ctx, workDir, opts.RepoURL); err != nil { + return err + } // Create and checkout branch from current HEAD. // The branch is based on whatever was cloned — the MR/PR target branch @@ -59,19 +62,19 @@ func CreateBranchAndPush(ctx context.Context, opts BranchOptions) error { "GIT_COMMITTER_NAME=" + opts.AuthorName, "GIT_COMMITTER_EMAIL=" + opts.AuthorEmail, } - if err := gitCmd(ctx, workDir, commitEnv, "commit", "-m", opts.CommitMsg); err != nil { + if err := gitCmd(ctx, workDir, commitEnv, "commit", "--no-verify", "--no-gpg-sign", "-m", opts.CommitMsg); err != nil { return fmt.Errorf("committing changes: %w", err) } slog.Info("changes committed", "branch", opts.BranchName) // Push via GIT_ASKPASS - pushEnv, cleanup, err := AskPassEnv(opts.Token) + pushEnv, cleanup, err := AskPassEnv(opts.Token, opts.RepoURL) if err != nil { return fmt.Errorf("preparing push credentials: %w", err) } defer cleanup() - if err := gitCmd(ctx, workDir, pushEnv, "push", "-u", "origin", opts.BranchName); err != nil { + if err := gitCmd(ctx, workDir, pushEnv, "push", "--no-verify", "-u", "origin", opts.BranchName); err != nil { return fmt.Errorf("pushing branch: %w", err) } slog.Info("branch pushed", "branch", opts.BranchName) @@ -79,21 +82,21 @@ func CreateBranchAndPush(ctx context.Context, opts BranchOptions) error { return nil } -// AskPassEnv prepares GIT_ASKPASS environment for authenticated git operations. +// AskPassEnv prepares GIT_ASKPASS environment for authenticated git operations +// against repoURL; the script answers only for repoURL's host. // Returns extra env vars and a cleanup function. -func AskPassEnv(token string) ([]string, func(), error) { +func AskPassEnv(token, repoURL string) ([]string, func(), error) { if token == "" { return nil, func() {}, nil } - askPassFile, err := createAskPassScript(token, "") + askPassFile, err := createAskPassScript(token, "", repoURL) if err != nil { return nil, nil, err } env := []string{ "GIT_ASKPASS=" + askPassFile, - "GIT_TERMINAL_PROMPT=0", } cleanup := func() { os.Remove(askPassFile) } return env, cleanup, nil @@ -101,11 +104,7 @@ func AskPassEnv(token string) ([]string, func(), error) { // gitCmd runs a git command in the given directory with optional extra env vars. func gitCmd(ctx context.Context, workDir string, extraEnv []string, args ...string) error { - cmd := exec.CommandContext(ctx, "git", args...) - cmd.Dir = workDir - if len(extraEnv) > 0 { - cmd.Env = append(os.Environ(), extraEnv...) - } + cmd := Command(ctx, workDir, extraEnv, args...) var stderr strings.Builder cmd.Stderr = &stderr @@ -118,8 +117,7 @@ func gitCmd(ctx context.Context, workDir string, extraEnv []string, args ...stri // gitOutput runs a git command and returns stdout. func gitOutput(ctx context.Context, workDir string, args ...string) (string, error) { - cmd := exec.CommandContext(ctx, "git", args...) - cmd.Dir = workDir + cmd := Command(ctx, workDir, nil, args...) out, err := cmd.Output() if err != nil { @@ -129,6 +127,7 @@ func gitOutput(ctx context.Context, workDir string, args ...string) (string, err } // GenerateBranchName creates a branch name with prefix and slug, adding numeric suffix if needed. +// The caller must have run SanitizeRepoConfig on workDir. func GenerateBranchName(ctx context.Context, workDir, prefix, slug string) string { base := prefix + slug name := base @@ -147,6 +146,9 @@ func GenerateBranchName(ctx context.Context, workDir, prefix, slug string) strin // DefaultBranch detects the default branch of the cloned repository // by reading the symbolic-ref of origin/HEAD. func DefaultBranch(ctx context.Context, workDir string) (string, error) { + if err := SanitizeRepoConfig(ctx, workDir, ""); err != nil { + return "", err + } // Try symbolic-ref first (set by clone) out, err := gitOutput(ctx, workDir, "symbolic-ref", "refs/remotes/origin/HEAD") if err == nil { @@ -180,12 +182,16 @@ func branchExists(ctx context.Context, workDir, name string) bool { // GetUnstagedDiff returns the diff of all uncommitted changes in the workspace. func GetUnstagedDiff(ctx context.Context, workDir string) (string, error) { - return gitOutput(ctx, workDir, "diff", "HEAD") + if err := SanitizeRepoConfig(ctx, workDir, ""); err != nil { + return "", err + } + return gitOutput(ctx, workDir, "diff", "--no-ext-diff", "--no-textconv", "HEAD") } // PushExistingOptions configures pushing follow-up changes to an existing branch. type PushExistingOptions struct { WorkDir string + RepoURL string // the URL the workspace was cloned from; pushed to, whatever origin says BranchName string CommitMsg string AuthorName string @@ -197,6 +203,9 @@ type PushExistingOptions struct { // Returns an error if there are no new changes to push. func CommitAndPushToExisting(ctx context.Context, opts PushExistingOptions) error { workDir := opts.WorkDir + if err := SanitizeRepoConfig(ctx, workDir, opts.RepoURL); err != nil { + return err + } // Stage all changes if err := gitCmd(ctx, workDir, nil, "add", "-A"); err != nil { @@ -219,19 +228,19 @@ func CommitAndPushToExisting(ctx context.Context, opts PushExistingOptions) erro "GIT_COMMITTER_NAME=" + opts.AuthorName, "GIT_COMMITTER_EMAIL=" + opts.AuthorEmail, } - if err := gitCmd(ctx, workDir, commitEnv, "commit", "-m", opts.CommitMsg); err != nil { + if err := gitCmd(ctx, workDir, commitEnv, "commit", "--no-verify", "--no-gpg-sign", "-m", opts.CommitMsg); err != nil { return fmt.Errorf("committing changes: %w", err) } slog.Info("follow-up changes committed", "branch", opts.BranchName) // Push via GIT_ASKPASS - pushEnv, cleanup, err := AskPassEnv(opts.Token) + pushEnv, cleanup, err := AskPassEnv(opts.Token, opts.RepoURL) if err != nil { return fmt.Errorf("preparing push credentials: %w", err) } defer cleanup() - if err := gitCmd(ctx, workDir, pushEnv, "push", "origin", opts.BranchName); err != nil { + if err := gitCmd(ctx, workDir, pushEnv, "push", "--no-verify", "origin", opts.BranchName); err != nil { return fmt.Errorf("pushing to branch: %w", err) } slog.Info("follow-up changes pushed", "branch", opts.BranchName) diff --git a/internal/tool/git/clone.go b/internal/tool/git/clone.go index 3c41fc0..1ba7ab5 100644 --- a/internal/tool/git/clone.go +++ b/internal/tool/git/clone.go @@ -5,7 +5,6 @@ import ( "fmt" "log/slog" "os" - "os/exec" "strings" ) @@ -31,34 +30,26 @@ type CloneOptions struct { // Clone clones a git repository using GIT_ASKPASS for token authentication. // The token is never embedded in the URL or stored in .git/config. func Clone(ctx context.Context, opts CloneOptions) error { - args := []string{"clone"} + args := []string{"clone", "--no-recurse-submodules"} if opts.Shallow { args = append(args, "--depth", "1") } if opts.Branch != "" { args = append(args, "--branch", opts.Branch) } - args = append(args, opts.RepoURL, opts.DestDir) - - cmd := exec.CommandContext(ctx, "git", args...) + args = append(args, "--", opts.RepoURL, opts.DestDir) // Token via GIT_ASKPASS — never stored in .git/config - var askPassFile string + var env []string if opts.Token != "" { - var err error - askPassFile, err = createAskPassScript(opts.Token, opts.Username) + askPassFile, err := createAskPassScript(opts.Token, opts.Username, opts.RepoURL) if err != nil { return fmt.Errorf("creating askpass script: %w", err) } defer os.Remove(askPassFile) - - cmd.Env = append(os.Environ(), - "GIT_ASKPASS="+askPassFile, - "GIT_TERMINAL_PROMPT=0", - ) - } else { - cmd.Env = append(os.Environ(), "GIT_TERMINAL_PROMPT=0") + env = []string{"GIT_ASKPASS=" + askPassFile} } + cmd := Command(ctx, "", env, args...) var stderr strings.Builder cmd.Stderr = &stderr @@ -73,25 +64,31 @@ func Clone(ctx context.Context, opts CloneOptions) error { } // createAskPassScript creates a temporary script that answers git credential -// prompts. With an empty username, the token is echoed for both the username -// and password prompts (PAT behavior). With a username, git's "Username for -// ..." prompt gets the username and every other prompt gets the token. -func createAskPassScript(token, username string) (string, error) { +// prompts for repoURL's scheme and host only. With an empty username, the +// token answers both the username and password prompts (PAT behavior); with a +// username, the username prompt gets the username. A prompt for any other +// host — a remote URL or redirect the server did not choose — gets nothing. +// For a repoURL without a network host the script never answers. +func createAskPassScript(token, username, repoURL string) (string, error) { f, err := os.CreateTemp("", "codeforge-askpass-*.sh") if err != nil { return "", err } - // Shell-escape the credentials to prevent injection - escaped := shellEscape(token) - var script string if username == "" { - script = fmt.Sprintf("#!/bin/sh\necho '%s'\n", escaped) - } else { - script = fmt.Sprintf( - "#!/bin/sh\ncase \"$1\" in\n[Uu]sername*) echo '%s' ;;\n*) echo '%s' ;;\nesac\n", - shellEscape(username), escaped, - ) + username = token + } + script := "#!/bin/sh\nexit 1\n" + if scheme, host := askPassTarget(repoURL); host != "" { + // Git prompts with "Username for '://': " and then + // "Password for '://@': ". The scheme and host are + // validated to contain no shell metacharacters; credentials are + // shell-escaped to prevent injection. + script = fmt.Sprintf("#!/bin/sh\ncase \"$1\" in\n"+ + "\"Username for '%[1]s://%[2]s': \") echo '%[3]s' ;;\n"+ + "\"Password for '%[1]s://\"*\"@%[2]s': \") echo '%[4]s' ;;\n"+ + "*) exit 1 ;;\nesac\n", + scheme, host, shellEscape(username), shellEscape(token)) } if _, err := f.WriteString(script); err != nil { diff --git a/internal/tool/git/clone_test.go b/internal/tool/git/clone_test.go index b1440df..2c3599d 100644 --- a/internal/tool/git/clone_test.go +++ b/internal/tool/git/clone_test.go @@ -8,16 +8,15 @@ import ( ) // runAskPass executes a generated askpass script with the given git prompt -// and returns its trimmed stdout. -func runAskPass(t *testing.T, scriptPath, prompt string) string { +// and returns its trimmed stdout and whether it answered (exit status 0). +func runAskPass(t *testing.T, scriptPath, prompt string) (string, bool) { t.Helper() out, err := exec.Command("sh", scriptPath, prompt).Output() - if err != nil { - t.Fatalf("running askpass script: %v", err) - } - return strings.TrimSuffix(string(out), "\n") + return strings.TrimSuffix(string(out), "\n"), err == nil } +const askPassRepoURL = "https://gitlab.example.com/group/repo.git" + func TestCreateAskPassScript(t *testing.T) { const promptUser = "Username for 'https://gitlab.example.com': " const promptPass = "Password for 'https://gitlab-ci-token@gitlab.example.com': " @@ -64,18 +63,11 @@ func TestCreateAskPassScript(t *testing.T) { prompt: promptPass, want: "to'ken", }, - { - name: "empty prompt with username falls through to token", - token: "job-token-secret", - username: GitLabCIJobTokenUsername, - prompt: "", - want: "job-token-secret", - }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - path, err := createAskPassScript(tt.token, tt.username) + path, err := createAskPassScript(tt.token, tt.username, askPassRepoURL) if err != nil { t.Fatalf("createAskPassScript: %v", err) } @@ -89,9 +81,60 @@ func TestCreateAskPassScript(t *testing.T) { t.Errorf("script mode = %o, want 0700", info.Mode().Perm()) } - if got := runAskPass(t, path, tt.prompt); got != tt.want { - t.Errorf("askpass(%q) = %q, want %q", tt.prompt, got, tt.want) + if got, ok := runAskPass(t, path, tt.prompt); !ok || got != tt.want { + t.Errorf("askpass(%q) = %q (answered %v), want %q", tt.prompt, got, ok, tt.want) } }) } } + +// TestAskPassScriptOnlyAnswersRepoHost checks that the token is never handed +// to a host other than the one the server cloned from — e.g. when the +// workspace's origin URL or an HTTP redirect points somewhere else. +func TestAskPassScriptOnlyAnswersRepoHost(t *testing.T) { + tests := []struct { + name string + repoURL string + prompt string + }{ + {"other host username", askPassRepoURL, "Username for 'https://evil.example': "}, + {"other host password", askPassRepoURL, "Password for 'https://tok@evil.example': "}, + {"host as prefix of another host", askPassRepoURL, "Password for 'https://tok@gitlab.example.com.evil.example': "}, + {"host smuggled in username", askPassRepoURL, "Password for 'https://x@gitlab.example.com@evil.example': "}, + {"scheme downgrade", askPassRepoURL, "Password for 'http://tok@gitlab.example.com': "}, + {"other port", askPassRepoURL, "Password for 'https://tok@gitlab.example.com:8443': "}, + {"empty prompt", askPassRepoURL, ""}, + {"local repository never answers", "/srv/repos/project.git", "Password for 'https://tok@gitlab.example.com': "}, + {"file URL never answers", "file:///srv/repos/project.git", "Username for 'https://gitlab.example.com': "}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + path, err := createAskPassScript("tok", "", tt.repoURL) + if err != nil { + t.Fatalf("createAskPassScript: %v", err) + } + defer os.Remove(path) + + if got, ok := runAskPass(t, path, tt.prompt); ok || got != "" { + t.Errorf("askpass(%q) = %q (answered %v), want no answer", tt.prompt, got, ok) + } + }) + } +} + +func TestAskPassScriptHostWithPort(t *testing.T) { + path, err := createAskPassScript("tok", "", "http://127.0.0.1:8080/group/repo.git") + if err != nil { + t.Fatalf("createAskPassScript: %v", err) + } + defer os.Remove(path) + + for _, prompt := range []string{ + "Username for 'http://127.0.0.1:8080': ", + "Password for 'http://tok@127.0.0.1:8080': ", + } { + if got, ok := runAskPass(t, path, prompt); !ok || got != "tok" { + t.Errorf("askpass(%q) = %q (answered %v), want token", prompt, got, ok) + } + } +} diff --git a/internal/tool/git/diff.go b/internal/tool/git/diff.go index b1681de..a7d3466 100644 --- a/internal/tool/git/diff.go +++ b/internal/tool/git/diff.go @@ -3,7 +3,6 @@ package git import ( "context" "fmt" - "os/exec" "regexp" "strconv" "strings" @@ -20,9 +19,12 @@ type ChangesSummary struct { // CalculateChanges computes a summary of workspace changes after CLI execution. // It runs git status and git diff --shortstat (both staged and unstaged). func CalculateChanges(ctx context.Context, workDir string) (*ChangesSummary, error) { + if err := SanitizeRepoConfig(ctx, workDir, ""); err != nil { + return nil, err + } + // git status --porcelain for file counts - statusCmd := exec.CommandContext(ctx, "git", "status", "--porcelain") - statusCmd.Dir = workDir + statusCmd := Command(ctx, workDir, nil, "status", "--porcelain") statusOut, err := statusCmd.Output() if err != nil { return nil, fmt.Errorf("git status: %w", err) @@ -67,13 +69,12 @@ func CalculateChanges(ctx context.Context, workDir string) (*ChangesSummary, err var shortStatRegex = regexp.MustCompile(`(\d+) insertions?\(\+\).*?(\d+) deletions?\(-\)|(\d+) insertions?\(\+\)|(\d+) deletions?\(-\)`) func shortStat(ctx context.Context, workDir string, cached bool) (insertions, deletions int) { - args := []string{"diff", "--shortstat"} + args := []string{"diff", "--no-ext-diff", "--no-textconv", "--shortstat"} if cached { - args = []string{"diff", "--cached", "--shortstat"} + args = []string{"diff", "--no-ext-diff", "--no-textconv", "--cached", "--shortstat"} } - cmd := exec.CommandContext(ctx, "git", args...) - cmd.Dir = workDir + cmd := Command(ctx, workDir, nil, args...) out, err := cmd.Output() if err != nil { return 0, 0 diff --git a/internal/tool/git/gitlab.go b/internal/tool/git/gitlab.go index 9b28b1c..2daff9b 100644 --- a/internal/tool/git/gitlab.go +++ b/internal/tool/git/gitlab.go @@ -142,3 +142,57 @@ func (c *GitLabMRCreator) GetMRStatus(ctx context.Context, repo *RepoInfo, token } return status, nil } + +// GitLabDeveloperAccess is GitLab's Developer role, the lowest access level +// that can push to a project. +const GitLabDeveloperAccess = 30 + +// GitLabAccessLevel returns userID's effective access level on a GitLab +// project, including membership inherited from groups, or 0 when the user is +// not an active member. baseURL is the instance's "scheme://host[:port]". +// +// Redirects are refused so the token is never forwarded to another host. +func GitLabAccessLevel(ctx context.Context, baseURL, token string, projectID, userID int) (int, error) { + endpoint := fmt.Sprintf("%s/api/v4/projects/%d/members/all/%d", baseURL, projectID, userID) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) + if err != nil { + return 0, fmt.Errorf("creating member request: %w", err) + } + req.Header.Set("PRIVATE-TOKEN", token) + + client := &http.Client{ + Timeout: 15 * time.Second, + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + } + resp, err := client.Do(req) + if err != nil { + return 0, fmt.Errorf("gitlab API request: %w", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if err != nil { + return 0, fmt.Errorf("reading gitlab response: %w", err) + } + switch resp.StatusCode { + case http.StatusOK: + case http.StatusNotFound: + return 0, nil + default: + return 0, fmt.Errorf("gitlab API returned %d: %s", resp.StatusCode, truncateBytes(body, 500)) + } + + var member struct { + AccessLevel int `json:"access_level"` + State string `json:"state"` + } + if err := json.Unmarshal(body, &member); err != nil { + return 0, fmt.Errorf("parsing gitlab member response: %w", err) + } + if member.State != "" && member.State != "active" { + return 0, nil + } + return member.AccessLevel, nil +} diff --git a/internal/tool/git/gitlab_test.go b/internal/tool/git/gitlab_test.go index e53804c..5c5a9b3 100644 --- a/internal/tool/git/gitlab_test.go +++ b/internal/tool/git/gitlab_test.go @@ -72,3 +72,46 @@ func TestGetMRStatus(t *testing.T) { }) } } + +func TestGitLabAccessLevel(t *testing.T) { + tests := []struct { + name string + status int + body string + want int + wantErr bool + }{ + {name: "developer", status: http.StatusOK, body: `{"id":7,"access_level":30,"state":"active"}`, want: 30}, + {name: "reporter", status: http.StatusOK, body: `{"id":7,"access_level":20,"state":"active"}`, want: 20}, + {name: "inactive member counts as none", status: http.StatusOK, body: `{"id":7,"access_level":40,"state":"awaiting"}`, want: 0}, + {name: "not a member", status: http.StatusNotFound, body: `{"message":"404 Not found"}`, want: 0}, + {name: "api error", status: http.StatusUnauthorized, body: `{"message":"401 Unauthorized"}`, wantErr: true}, + {name: "redirect is not followed", status: http.StatusFound, body: ``, wantErr: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/v4/projects/42/members/all/7" { + t.Errorf("path = %s", r.URL.Path) + } + if r.Header.Get("PRIVATE-TOKEN") != "tok" { + t.Errorf("PRIVATE-TOKEN = %q", r.Header.Get("PRIVATE-TOKEN")) + } + if tt.status == http.StatusFound { + w.Header().Set("Location", "http://elsewhere.invalid/") + } + w.WriteHeader(tt.status) + _, _ = w.Write([]byte(tt.body)) + })) + defer srv.Close() + + got, err := GitLabAccessLevel(context.Background(), srv.URL, "tok", 42, 7) + if (err != nil) != tt.wantErr { + t.Fatalf("err = %v, wantErr %v", err, tt.wantErr) + } + if got != tt.want { + t.Errorf("level = %d, want %d", got, tt.want) + } + }) + } +} diff --git a/internal/tool/git/safe.go b/internal/tool/git/safe.go new file mode 100644 index 0000000..37e31fc --- /dev/null +++ b/internal/tool/git/safe.go @@ -0,0 +1,250 @@ +package git + +import ( + "context" + "errors" + "fmt" + "net/url" + "os" + "os/exec" + "path/filepath" + "regexp" + "strings" +) + +// Server-side git runs in workspaces the AI CLI has written to. The CLI runs +// code from the repository it works on, so everything it can reach is +// untrusted: the work tree, .git/config, .git/hooks, and the global git config +// in the HOME it shares with the server. Git runs programs named in its +// configuration (core.fsmonitor, hooks, filter and diff drivers, credential +// helpers, ...), and the server's git processes carry provider tokens and used +// to inherit the server's whole environment. +// +// Every git process the server starts therefore goes through Command, which +// +// - pins the settings that could run a program, recurse into submodules or +// hand credentials elsewhere on the command line, where repository config +// cannot override them; +// - ignores the global config (GIT_CONFIG_GLOBAL=/dev/null). The system +// config stays: it is owned by root, not by the user the CLI runs as; +// - builds the environment from an allowlist, so server secrets never reach +// git or anything git spawns. +// +// Settings that cannot be pinned that way (filter..*, include.path, +// url..insteadOf, core.worktree, ...) are removed by SanitizeRepoConfig, +// which rebuilds .git/config from an allowlist before the server runs git in a +// workspace. + +// hardenedArgs are global options prepended to every server-side git command. +var hardenedArgs = []string{ + "-c", "core.fsmonitor=false", + "-c", "core.hooksPath=/dev/null", + "-c", "credential.helper=", + "-c", "credential.useHttpPath=false", + "-c", "core.askPass=", + "-c", "protocol.ext.allow=never", + "-c", "protocol.file.allow=user", + "-c", "submodule.recurse=false", + "-c", "fetch.recurseSubmodules=false", + "-c", "push.recurseSubmodules=no", + "-c", "diff.ignoreSubmodules=all", + "-c", "commit.gpgSign=false", + // The global config that used to mark workspaces as safe is no longer + // read, and repository config is rebuilt before use, so ownership checks + // add nothing here. + "-c", "safe.directory=*", +} + +// gitEnvAllow are the server environment variables passed to git. +var gitEnvAllow = map[string]struct{}{ + "PATH": {}, + "HOME": {}, + "TMPDIR": {}, + "TZ": {}, + + // TLS trust for self-hosted instances. + "SSL_CERT_FILE": {}, + "SSL_CERT_DIR": {}, + "CURL_CA_BUNDLE": {}, + "GIT_SSL_CAINFO": {}, + "GIT_SSL_CAPATH": {}, + "GIT_SSL_NO_VERIFY": {}, + + // Outbound proxy configuration. + "HTTP_PROXY": {}, + "HTTPS_PROXY": {}, + "ALL_PROXY": {}, + "NO_PROXY": {}, + "http_proxy": {}, + "https_proxy": {}, + "all_proxy": {}, + "no_proxy": {}, +} + +// gitEnv returns the environment for a server-side git process. +func gitEnv(extra []string) []string { + env := []string{ + "GIT_CONFIG_GLOBAL=/dev/null", + "GIT_TERMINAL_PROMPT=0", + } + for _, kv := range os.Environ() { + name, _, ok := strings.Cut(kv, "=") + if !ok { + continue + } + if _, allowed := gitEnvAllow[name]; allowed { + env = append(env, kv) + } + } + return append(env, extra...) +} + +// Command returns a git command that runs in dir (the current directory when +// dir is empty) with the hardened configuration and an allowlisted +// environment plus extraEnv. +func Command(ctx context.Context, dir string, extraEnv []string, args ...string) *exec.Cmd { + full := make([]string, 0, len(hardenedArgs)+len(args)) + full = append(full, hardenedArgs...) + full = append(full, args...) + + cmd := exec.CommandContext(ctx, "git", full...) + cmd.Dir = dir + cmd.Env = gitEnv(extraEnv) + return cmd +} + +var ( + objectFormats = map[string]bool{"sha1": true, "sha256": true} + refStorages = map[string]bool{"files": true, "reftable": true} + boolValue = regexp.MustCompile(`^(true|false)$`) +) + +// SanitizeRepoConfig rebuilds workDir/.git/config from an allowlist so that +// none of the settings git would execute or follow survive from the workspace: +// only the repository format, a few filesystem probes from clone time, and the +// origin remote are kept. +// +// originURL, when non-empty, becomes remote.origin.url — callers that talk to +// the remote pass the URL the server cloned from, so a changed origin cannot +// redirect a fetch or push (and the token with it). With an empty originURL the +// existing origin URL is carried over as plain data, for local-only commands. +// +// A .git that is not a real directory (a symlink, or a gitfile pointing +// elsewhere) or that redirects to a common directory is refused. +func SanitizeRepoConfig(ctx context.Context, workDir, originURL string) error { + gitDir := filepath.Join(workDir, ".git") + info, err := os.Lstat(gitDir) + if err != nil { + return fmt.Errorf("inspecting .git: %w", err) + } + if !info.IsDir() { + return fmt.Errorf("refusing to run git: %s is not a directory", gitDir) + } + if _, err := os.Lstat(filepath.Join(gitDir, "commondir")); err == nil { + return errors.New("refusing to run git: .git/commondir redirects the repository") + } + if strings.ContainsAny(originURL, "\x00\r\n") { + return errors.New("refusing to run git: invalid origin URL") + } + + cfgPath := filepath.Join(gitDir, "config") + get := func(key string, typ string) string { + args := []string{"config", "--file", cfgPath, "--no-includes"} + if typ != "" { + args = append(args, "--type="+typ) + } + // Run outside the workspace so git does not discover the repository. + out, err := Command(ctx, os.TempDir(), nil, append(args, "--get", key)...).Output() + if err != nil { + return "" + } + return strings.TrimSpace(string(out)) + } + + var b strings.Builder + version := get("core.repositoryformatversion", "int") + if version != "1" { + version = "0" + } + b.WriteString("[core]\n") + fmt.Fprintf(&b, "\trepositoryformatversion = %s\n", version) + b.WriteString("\tbare = false\n") + b.WriteString("\tlogallrefupdates = true\n") + for _, key := range []string{"filemode", "ignorecase", "precomposeunicode", "symlinks"} { + if v := get("core."+key, "bool"); boolValue.MatchString(v) { + fmt.Fprintf(&b, "\t%s = %s\n", key, v) + } + } + + if version == "1" { + var ext strings.Builder + if v := get("extensions.objectformat", ""); objectFormats[v] { + fmt.Fprintf(&ext, "\tobjectformat = %s\n", v) + } + if v := get("extensions.refstorage", ""); refStorages[v] { + fmt.Fprintf(&ext, "\trefstorage = %s\n", v) + } + if ext.Len() > 0 { + b.WriteString("[extensions]\n") + b.WriteString(ext.String()) + } + } + + origin := originURL + if origin == "" { + origin = get("remote.origin.url", "") + if strings.ContainsAny(origin, "\x00\r\n") { + origin = "" + } + } + if origin != "" { + b.WriteString("[remote \"origin\"]\n") + fmt.Fprintf(&b, "\turl = %s\n", quoteConfigValue(origin)) + b.WriteString("\tfetch = +refs/heads/*:refs/remotes/origin/*\n") + } + + // Replace the file rather than editing it: rename swaps out a symlinked + // config instead of writing through it. + tmp, err := os.CreateTemp(gitDir, "config.codeforge-*") + if err != nil { + return fmt.Errorf("writing git config: %w", err) + } + defer os.Remove(tmp.Name()) + if _, err := tmp.WriteString(b.String()); err != nil { + _ = tmp.Close() + return fmt.Errorf("writing git config: %w", err) + } + if err := tmp.Close(); err != nil { + return fmt.Errorf("writing git config: %w", err) + } + if err := os.Chmod(tmp.Name(), 0o644); err != nil { + return fmt.Errorf("writing git config: %w", err) + } + if err := os.Rename(tmp.Name(), cfgPath); err != nil { + return fmt.Errorf("writing git config: %w", err) + } + return nil +} + +// quoteConfigValue quotes a value for a git config file. +func quoteConfigValue(v string) string { + v = strings.ReplaceAll(v, `\`, `\\`) + v = strings.ReplaceAll(v, `"`, `\"`) + return `"` + v + `"` +} + +// askPassTarget returns the scheme and host[:port] the askpass script may +// answer for, or empty strings when repoURL names no network host (a local +// path or file:// URL never prompts for credentials). +func askPassTarget(repoURL string) (scheme, host string) { + u, err := url.Parse(repoURL) + if err != nil || (u.Scheme != "https" && u.Scheme != "http") { + return "", "" + } + if !validAskPassHost.MatchString(u.Host) { + return "", "" + } + return u.Scheme, u.Host +} + +var validAskPassHost = regexp.MustCompile(`^[A-Za-z0-9.\-]+(:[0-9]+)?$|^\[[0-9A-Fa-f:.]+\](:[0-9]+)?$`) diff --git a/internal/tool/git/safe_test.go b/internal/tool/git/safe_test.go new file mode 100644 index 0000000..b94239b --- /dev/null +++ b/internal/tool/git/safe_test.go @@ -0,0 +1,299 @@ +package git + +import ( + "context" + "net/http" + "net/http/httptest" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "testing" +) + +// rawGit runs git without any hardening, the way the AI CLI (or a script it +// runs) would in the workspace. +func rawGit(t *testing.T, dir string, args ...string) string { + t.Helper() + cmd := exec.Command("git", append([]string{"-C", dir}, args...)...) + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("git %v: %v: %s", args, err, out) + } + return strings.TrimSpace(string(out)) +} + +// cloneWorkspace clones a fresh workspace from a local bare remote and returns +// (workspace, remote). +func cloneWorkspace(t *testing.T) (string, string) { + t.Helper() + if _, err := exec.LookPath("git"); err != nil { + t.Skip("git not installed") + } + src := initSourceRepo(t, "main") + remote := filepath.Join(t.TempDir(), "remote.git") + rawGit(t, src, "clone", "--bare", src, remote) + + ws := filepath.Join(t.TempDir(), "ws") + if err := Clone(context.Background(), CloneOptions{RepoURL: remote, DestDir: ws}); err != nil { + t.Fatalf("Clone: %v", err) + } + return ws, remote +} + +// writeExecutable writes a shell script and returns its path. +func writeExecutable(t *testing.T, dir, name, body string) string { + t.Helper() + path := filepath.Join(dir, name) + if err := os.WriteFile(path, []byte("#!/bin/sh\n"+body), 0o755); err != nil { + t.Fatal(err) + } + return path +} + +// TestServerGitIgnoresWorkspaceExecutors plants every kind of command the +// workspace's git configuration can make git run — fsmonitor, hooks, filter +// and diff drivers, in .git/config, in an included file and in the global +// config — then drives the server-side git operations over that workspace. +// None of the planted programs may run, and the server's secrets must not +// reach any git process. +func TestServerGitIgnoresWorkspaceExecutors(t *testing.T) { + ws, remote := cloneWorkspace(t) + ctx := context.Background() + + t.Setenv("CODEFORGE_ENCRYPTION__KEY", "server-secret-key") + t.Setenv("GITHUB_TOKEN", "server-github-token") + + tools := t.TempDir() + marker := filepath.Join(tools, "ran") + record := "echo \"$0 $*\" >> " + marker + "\nenv >> " + marker + "\n" + run := writeExecutable(t, tools, "run.sh", record+"exit 0\n") + filter := writeExecutable(t, tools, "filter.sh", record+"cat\n") + textconv := writeExecutable(t, tools, "textconv.sh", record+"cat \"$1\"\n") + + hooks := filepath.Join(tools, "hooks") + if err := os.MkdirAll(hooks, 0o755); err != nil { + t.Fatal(err) + } + for _, hook := range []string{"pre-commit", "commit-msg", "post-commit", "pre-push", "post-checkout", "reference-transaction"} { + writeExecutable(t, hooks, hook, record+"exit 0\n") + writeExecutable(t, filepath.Join(ws, ".git", "hooks"), hook, record+"exit 0\n") + } + + included := filepath.Join(tools, "included.gitconfig") + if err := os.WriteFile(included, []byte("[diff \"evil2\"]\n\ttextconv = "+textconv+"\n"), 0o644); err != nil { + t.Fatal(err) + } + + rawGit(t, ws, "config", "core.fsmonitor", run) + rawGit(t, ws, "config", "core.hooksPath", hooks) + rawGit(t, ws, "config", "filter.evil.clean", filter) + rawGit(t, ws, "config", "filter.evil.smudge", filter) + rawGit(t, ws, "config", "diff.evil.textconv", textconv) + rawGit(t, ws, "config", "diff.external", run) + rawGit(t, ws, "config", "include.path", included) + + // The CLI shares HOME with the server, so it can write the global config too. + home := t.TempDir() + t.Setenv("HOME", home) + t.Setenv("XDG_CONFIG_HOME", filepath.Join(home, ".config")) + globalCfg := "[filter \"evil3\"]\n\tclean = " + filter + "\n[core]\n\tfsmonitor = " + run + "\n" + if err := os.WriteFile(filepath.Join(home, ".gitconfig"), []byte(globalCfg), 0o644); err != nil { + t.Fatal(err) + } + + attrs := "*.md filter=evil diff=evil\n*.txt filter=evil3 diff=evil2\n" + writeTestFile(t, ws, ".gitattributes", attrs) + writeTestFile(t, ws, "README.md", "changed by the session\n") + writeTestFile(t, ws, "notes.txt", "new file\n") + + if _, err := CalculateChanges(ctx, ws); err != nil { + t.Fatalf("CalculateChanges: %v", err) + } + if _, err := Diff(ctx, ws); err != nil { + t.Fatalf("Diff: %v", err) + } + if _, err := GetUnstagedDiff(ctx, ws); err != nil { + t.Fatalf("GetUnstagedDiff: %v", err) + } + if _, err := DefaultBranch(ctx, ws); err != nil { + t.Fatalf("DefaultBranch: %v", err) + } + err := CreateBranchAndPush(ctx, BranchOptions{ + WorkDir: ws, + RepoURL: remote, + BranchName: "codeforge/test", + CommitMsg: "test", + AuthorName: "CodeForge", + AuthorEmail: "codeforge@example.com", + }) + if err != nil { + t.Fatalf("CreateBranchAndPush: %v", err) + } + + if out, err := os.ReadFile(marker); err == nil { + t.Fatalf("workspace-controlled program ran during server git operations:\n%s", out) + } + if got := rawGit(t, remote, "rev-parse", "--verify", "refs/heads/codeforge/test"); got == "" { + t.Error("branch was not pushed") + } +} + +// TestPushIgnoresWorkspaceOrigin points the workspace's origin (and an +// insteadOf rewrite of the real remote) at a server that asks for +// credentials. The push must go to the URL the server cloned from and the +// token must never reach the other server. +func TestPushIgnoresWorkspaceOrigin(t *testing.T) { + ws, remote := cloneWorkspace(t) + ctx := context.Background() + + var ( + mu sync.Mutex + requests []string + ) + evil := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + requests = append(requests, r.URL.String()+" auth="+r.Header.Get("Authorization")) + mu.Unlock() + w.Header().Set("WWW-Authenticate", `Basic realm="git"`) + w.WriteHeader(http.StatusUnauthorized) + })) + defer evil.Close() + + rawGit(t, ws, "remote", "set-url", "origin", evil.URL+"/group/repo.git") + rawGit(t, ws, "config", "remote.origin.pushurl", evil.URL+"/group/repo.git") + rawGit(t, ws, "config", "url."+evil.URL+"/.insteadOf", remote) + writeTestFile(t, ws, "README.md", "changed by the session\n") + + err := CreateBranchAndPush(ctx, BranchOptions{ + WorkDir: ws, + RepoURL: remote, + BranchName: "codeforge/test", + CommitMsg: "test", + AuthorName: "CodeForge", + AuthorEmail: "codeforge@example.com", + Token: "secret-token", + }) + + mu.Lock() + defer mu.Unlock() + if len(requests) > 0 { + t.Fatalf("push contacted the workspace-chosen remote: %v", requests) + } + if err != nil { + t.Fatalf("CreateBranchAndPush: %v", err) + } + if got := rawGit(t, ws, "config", "--get", "remote.origin.url"); got != remote { + t.Errorf("origin = %q, want %q", got, remote) + } +} + +func TestCommandEnvironmentExcludesServerSecrets(t *testing.T) { + secrets := map[string]string{ + "CODEFORGE_ENCRYPTION__KEY": "k", + "CODEFORGE_SERVER__AUTH_TOKEN": "t", + "CODEFORGE_REDIS__URL": "redis://:pw@redis:6379", + "GITHUB_TOKEN": "ghp", + "GITLAB_TOKEN": "glpat", + "AWS_SECRET_ACCESS_KEY": "aws", + "GIT_DIR": "/elsewhere", + "GIT_CONFIG_PARAMETERS": "'core.fsmonitor'='/tmp/x'", + "GIT_SSH_COMMAND": "/tmp/x", + } + for k, v := range secrets { + t.Setenv(k, v) + } + t.Setenv("HTTPS_PROXY", "http://proxy:3128") + + cmd := Command(context.Background(), t.TempDir(), []string{"GIT_ASKPASS=/tmp/askpass"}, "status") + env := map[string]string{} + for _, kv := range cmd.Env { + k, v, _ := strings.Cut(kv, "=") + env[k] = v + } + for k := range secrets { + if _, ok := env[k]; ok { + t.Errorf("%s reached the git environment", k) + } + } + for k, want := range map[string]string{ + "GIT_CONFIG_GLOBAL": "/dev/null", + "GIT_TERMINAL_PROMPT": "0", + "GIT_ASKPASS": "/tmp/askpass", + "HTTPS_PROXY": "http://proxy:3128", + } { + if env[k] != want { + t.Errorf("%s = %q, want %q", k, env[k], want) + } + } + if env["PATH"] == "" { + t.Error("PATH missing from the git environment") + } +} + +func TestSanitizeRepoConfig(t *testing.T) { + ctx := context.Background() + + t.Run("keeps origin as data for local commands", func(t *testing.T) { + ws, remote := cloneWorkspace(t) + rawGit(t, ws, "config", "user.name", "someone") + if err := SanitizeRepoConfig(ctx, ws, ""); err != nil { + t.Fatal(err) + } + if got := rawGit(t, ws, "config", "--get", "remote.origin.url"); got != remote { + t.Errorf("origin = %q, want %q", got, remote) + } + if out, err := exec.Command("git", "-C", ws, "config", "--local", "--get", "user.name").Output(); err == nil { + t.Errorf("user.name survived: %q", out) + } + // The repository still works. + rawGit(t, ws, "status", "--porcelain") + rawGit(t, ws, "rev-parse", "--verify", "refs/remotes/origin/main") + }) + + t.Run("sets the given origin", func(t *testing.T) { + ws, _ := cloneWorkspace(t) + want := `https://gitlab.example.com/group/re"po.git` + if err := SanitizeRepoConfig(ctx, ws, want); err != nil { + t.Fatal(err) + } + if got := rawGit(t, ws, "config", "--get", "remote.origin.url"); got != want { + t.Errorf("origin = %q, want %q", got, want) + } + }) + + t.Run("refuses a gitfile", func(t *testing.T) { + ws := t.TempDir() + writeTestFile(t, ws, ".git", "gitdir: /elsewhere\n") + if err := SanitizeRepoConfig(ctx, ws, ""); err == nil { + t.Error("expected an error for a .git file") + } + }) + + t.Run("refuses a symlinked .git", func(t *testing.T) { + other, _ := cloneWorkspace(t) + ws := t.TempDir() + if err := os.Symlink(filepath.Join(other, ".git"), filepath.Join(ws, ".git")); err != nil { + t.Fatal(err) + } + if err := SanitizeRepoConfig(ctx, ws, ""); err == nil { + t.Error("expected an error for a symlinked .git") + } + }) + + t.Run("refuses commondir", func(t *testing.T) { + ws, _ := cloneWorkspace(t) + writeTestFile(t, filepath.Join(ws, ".git"), "commondir", "/elsewhere\n") + if err := SanitizeRepoConfig(ctx, ws, ""); err == nil { + t.Error("expected an error for .git/commondir") + } + }) + + t.Run("refuses a multi-line origin", func(t *testing.T) { + ws, _ := cloneWorkspace(t) + if err := SanitizeRepoConfig(ctx, ws, "https://example.com/a\n[core]\nfsmonitor=x"); err == nil { + t.Error("expected an error for a multi-line origin") + } + }) +} diff --git a/internal/tool/git/workspace_diff.go b/internal/tool/git/workspace_diff.go index 2c7b774..dbdd0d0 100644 --- a/internal/tool/git/workspace_diff.go +++ b/internal/tool/git/workspace_diff.go @@ -44,6 +44,10 @@ type WorkspaceDiff struct { // (e.g. during create-pr) is unaffected. The unified diff is capped at 1 MB, // cut at a line boundary with Truncated set. func Diff(ctx context.Context, workDir string) (*WorkspaceDiff, error) { + if err := SanitizeRepoConfig(ctx, workDir, ""); err != nil { + return nil, err + } + // Intent-to-add so untracked files show up in `git diff HEAD`. if _, err := runGit(ctx, workDir, "add", "-A", "-N", "."); err != nil { return nil, err @@ -52,12 +56,12 @@ func Diff(ctx context.Context, workDir string) (*WorkspaceDiff, error) { // --no-renames keeps the unified diff, numstat, and porcelain status // consistent: a moved file is reported as one deletion + one addition // everywhere instead of a rename entry in some outputs only. - unified, err := runGit(ctx, workDir, "diff", "--no-renames", "HEAD") + unified, err := runGit(ctx, workDir, "diff", "--no-ext-diff", "--no-textconv", "--no-renames", "HEAD") if err != nil { return nil, err } - numstat, err := runGit(ctx, workDir, "diff", "--no-renames", "--numstat", "HEAD") + numstat, err := runGit(ctx, workDir, "diff", "--no-ext-diff", "--no-textconv", "--no-renames", "--numstat", "HEAD") if err != nil { return nil, err } @@ -107,12 +111,11 @@ func Diff(ctx context.Context, workDir string) (*WorkspaceDiff, error) { return result, nil } -// runGit executes git with explicit args against the given directory -// (git -C ...) and returns stdout. On failure, stderr from the git -// process is folded into the returned error. +// runGit executes git with explicit args in the given directory and returns +// stdout. On failure, stderr from the git process is folded into the returned +// error. func runGit(ctx context.Context, workDir string, args ...string) ([]byte, error) { - full := append([]string{"-C", workDir}, args...) - cmd := exec.CommandContext(ctx, "git", full...) + cmd := Command(ctx, workDir, nil, args...) out, err := cmd.Output() if err != nil { var exitErr *exec.ExitError diff --git a/internal/tool/runner/codex.go b/internal/tool/runner/codex.go index 82c8725..5e37276 100644 --- a/internal/tool/runner/codex.go +++ b/internal/tool/runner/codex.go @@ -34,8 +34,8 @@ func NewCodexRunner(binaryPath string) *CodexRunner { return &CodexRunner{binaryPath: binaryPath} } -// Run executes Codex CLI with JSON output, calling OnEvent for each line. -func (c *CodexRunner) Run(ctx context.Context, opts RunOptions) (*RunResult, error) { +// codexArgs builds the `codex exec` arguments for a run. +func codexArgs(opts RunOptions) []string { // --full-auto = --sandbox workspace-write + auto-approve on-request. // We use danger-full-access instead because Codex's Landlock sandbox // does not work inside Docker (missing kernel support / capabilities). @@ -61,7 +61,15 @@ func (c *CodexRunner) Run(ctx context.Context, opts RunOptions) (*RunResult, err prompt = opts.AppendSystemPrompt + "\n\n---\n\n" + prompt } - args = append(args, prompt) + // "--" ends option parsing: a prompt starting with "-" (e.g. "-c key=value") + // would otherwise be read as a Codex config override. + args = append(args, "--", prompt) + return args +} + +// Run executes Codex CLI with JSON output, calling OnEvent for each line. +func (c *CodexRunner) Run(ctx context.Context, opts RunOptions) (*RunResult, error) { + args := codexArgs(opts) cmd := exec.CommandContext(ctx, c.binaryPath, args...) cmd.Dir = opts.WorkDir diff --git a/internal/tool/runner/codex_test.go b/internal/tool/runner/codex_test.go index b7361b5..bfa1708 100644 --- a/internal/tool/runner/codex_test.go +++ b/internal/tool/runner/codex_test.go @@ -122,3 +122,21 @@ func TestNewCodexRunner_PathResolution(t *testing.T) { }) } } + +// TestCodexArgs_PromptCannotInjectOptions checks that the prompt is passed +// after "--", so a prompt like "-c model_provider=..." stays a prompt instead +// of becoming a Codex config override. +func TestCodexArgs_PromptCannotInjectOptions(t *testing.T) { + for _, prompt := range []string{"-cmodel_provider=evil", "--sandbox=read-only", "resume"} { + args := codexArgs(RunOptions{Prompt: prompt, WorkDir: "/ws", Model: "gpt-5"}) + n := len(args) + if n < 2 || args[n-2] != "--" || args[n-1] != prompt { + t.Errorf("prompt %q: args end with %q, want [\"--\" %q]", prompt, args[max(0, n-2):], prompt) + } + } + + args := codexArgs(RunOptions{Prompt: "-x", AppendSystemPrompt: "system"}) + if got := args[len(args)-1]; got != "system\n\n---\n\n-x" { + t.Errorf("prompt with system prefix = %q", got) + } +} diff --git a/internal/tools/resolver.go b/internal/tools/resolver.go index 50032b8..8ba0305 100644 --- a/internal/tools/resolver.go +++ b/internal/tools/resolver.go @@ -26,8 +26,20 @@ func NewResolver(registry Registry, keyReg ...keys.Registry) *Resolver { return r } -// Resolve converts a list of per-session tool requests into fully resolved ToolInstances. +// Resolve converts a list of per-session tool requests into fully resolved +// ToolInstances, filling missing config from the operator's provider keys. func (r *Resolver) Resolve(ctx context.Context, projectID string, sessionTools []SessionTool) ([]ToolInstance, error) { + return r.resolve(ctx, projectID, sessionTools, true) +} + +// ResolveOwnConfig is Resolve without the provider-key auto-fill: each tool +// gets only the config the session supplied. Used for subscription tenant +// sessions, which must not receive the operator's keys. +func (r *Resolver) ResolveOwnConfig(ctx context.Context, projectID string, sessionTools []SessionTool) ([]ToolInstance, error) { + return r.resolve(ctx, projectID, sessionTools, false) +} + +func (r *Resolver) resolve(ctx context.Context, projectID string, sessionTools []SessionTool, autoFill bool) ([]ToolInstance, error) { if len(sessionTools) == 0 { return nil, nil } @@ -41,7 +53,10 @@ func (r *Resolver) Resolve(ctx context.Context, projectID string, sessionTools [ } // Auto-fill missing config from Provider Keys - config := r.autoFillConfig(ctx, def, tt.Config) + config := tt.Config + if autoFill { + config = r.autoFillConfig(ctx, def, tt.Config) + } if err := ValidateConfig(def, config); err != nil { return nil, fmt.Errorf("tool %q: %w", tt.Name, err) diff --git a/internal/tools/resolver_test.go b/internal/tools/resolver_test.go index 35d42e1..ee6b340 100644 --- a/internal/tools/resolver_test.go +++ b/internal/tools/resolver_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/freema/codeforge/internal/apperror" + "github.com/freema/codeforge/internal/keys" ) // mockRegistry is a test double for Registry. @@ -168,3 +169,50 @@ func TestResolver_EmptyTools(t *testing.T) { t.Errorf("expected nil for empty tools, got %v", instances) } } + +// operatorKeyRegistry holds an operator GitHub key that auto-fill would use. +type operatorKeyRegistry struct{} + +func (operatorKeyRegistry) Create(context.Context, keys.Key) error { return nil } +func (operatorKeyRegistry) List(context.Context) ([]keys.Key, error) { return nil, nil } +func (operatorKeyRegistry) Delete(context.Context, string) error { return nil } +func (operatorKeyRegistry) Resolve(context.Context, string, string) (string, error) { + return "", apperror.NotFound("no key") +} +func (operatorKeyRegistry) Verify(context.Context, string) (*keys.VerifyResult, string, error) { + return nil, "", nil +} +func (operatorKeyRegistry) ResolveByName(_ context.Context, name string) (string, string, error) { + if name == "github-env" { + return "operator-token", "github", nil + } + return "", "", apperror.NotFound("no key") +} +func (operatorKeyRegistry) ResolveFullByName(context.Context, string) (string, string, string, error) { + return "", "", "", apperror.NotFound("no key") +} + +func TestResolver_ResolveOwnConfigSkipsAutoFill(t *testing.T) { + ctx := context.Background() + resolver := NewResolver(newMockRegistry(), operatorKeyRegistry{}) + + filled, err := resolver.Resolve(ctx, "", []SessionTool{{Name: "github"}}) + if err != nil { + t.Fatalf("Resolve: %v", err) + } + if filled[0].Config["token"] != "operator-token" { + t.Fatalf("Resolve did not auto-fill: %v", filled[0].Config) + } + + if _, err := resolver.ResolveOwnConfig(ctx, "", []SessionTool{{Name: "github"}}); err == nil { + t.Error("ResolveOwnConfig accepted a tool with no token, so it must have auto-filled one") + } + + own, err := resolver.ResolveOwnConfig(ctx, "", []SessionTool{{Name: "github", Config: map[string]string{"token": "tenant-token"}}}) + if err != nil { + t.Fatalf("ResolveOwnConfig: %v", err) + } + if own[0].Config["token"] != "tenant-token" { + t.Errorf("token = %q, want tenant-token", own[0].Config["token"]) + } +} diff --git a/internal/worker/executor.go b/internal/worker/executor.go index 804bc6e..13322a2 100644 --- a/internal/worker/executor.go +++ b/internal/worker/executor.go @@ -8,7 +8,6 @@ import ( "io/fs" "log/slog" "os" - "os/exec" "os/user" "path/filepath" "runtime/debug" @@ -274,8 +273,10 @@ func (e *Executor) resolveTimeout(t *session.Session) int { } // resolveToken resolves the access token from the key registry if not already set. +// Tenant sessions keep whatever token they brought (possibly none, for a +// public repository) and never fall back to the operator's credentials. func (e *Executor) resolveToken(ctx context.Context, t *session.Session, log *slog.Logger) { - if e.keyResolver == nil || t.AccessToken != "" { + if e.keyResolver == nil || t.AccessToken != "" || !t.UsesOperatorCredentials() { return } token, err := e.keyResolver.ResolveToken(ctx, t.RepoURL, t.AccessToken, t.ProviderKey) @@ -365,7 +366,11 @@ func (e *Executor) setupMCP(ctx context.Context, t *session.Session, workDir str // Resolve tool definitions → MCP servers var toolMCPServers []mcp.Server if e.toolResolver != nil && t.Config != nil && len(t.Config.Tools) > 0 { - instances, err := e.toolResolver.Resolve(ctx, t.RepoURL, t.Config.Tools) + resolve := e.toolResolver.Resolve + if !t.UsesOperatorCredentials() { + resolve = e.toolResolver.ResolveOwnConfig + } + instances, err := resolve(ctx, t.RepoURL, t.Config.Tools) if err != nil { // Fail-closed: session explicitly requested tools but resolve failed return "", fmt.Errorf("tool resolution failed: %w", err) @@ -402,7 +407,15 @@ func (e *Executor) setupMCP(ctx context.Context, t *session.Session, workDir str cli = t.Config.CLI } - if err := e.mcpInstaller.Setup(ctx, workDir, t.RepoURL, cli, taskMCPServers); err != nil { + var err error + if t.UsesOperatorCredentials() { + err = e.mcpInstaller.Setup(ctx, workDir, t.RepoURL, cli, taskMCPServers) + } else if len(taskMCPServers) > 0 { + // The operator's registered servers carry the operator's credentials; + // a tenant session gets only the servers it configured itself. + err = mcp.WriteMCPConfigForCLI(workDir, cli, taskMCPServers) + } + if err != nil { if len(taskMCPServers) > 0 { // Fail-closed: MCP servers were configured but install failed return "", fmt.Errorf("MCP setup failed: %w", err) @@ -935,18 +948,20 @@ func (e *Executor) cloneStep(ctx context.Context, t *session.Session, workDir st func (e *Executor) pullBranch(ctx context.Context, t *session.Session, workDir string, log *slog.Logger) { log.Info("pulling latest changes", "branch", t.Branch) - askPassEnv, cleanup, err := gitpkg.AskPassEnv(t.AccessToken) + // The workspace was written by the previous iteration's CLI run. + if err := gitpkg.SanitizeRepoConfig(ctx, workDir, t.RepoURL); err != nil { + log.Warn("refusing to pull into workspace (continuing with existing workspace)", "error", err) + return + } + + askPassEnv, cleanup, err := gitpkg.AskPassEnv(t.AccessToken, t.RepoURL) if err != nil { log.Warn("failed to create askpass for pull", "error", err) return } defer cleanup() - cmd := exec.CommandContext(ctx, "git", "pull", "origin", t.Branch) - cmd.Dir = workDir - if len(askPassEnv) > 0 { - cmd.Env = append(os.Environ(), askPassEnv...) - } + cmd := gitpkg.Command(ctx, workDir, askPassEnv, "pull", "--no-verify", "--no-recurse-submodules", "origin", t.Branch) if err := cmd.Run(); err != nil { log.Warn("git pull failed (continuing with existing workspace)", "error", err) @@ -1280,22 +1295,19 @@ func (e *Executor) sendWebhookAsync(ctx context.Context, sessionID, callbackURL // fetchAndCheckoutPR fetches a PR ref from origin and checks out a local branch. // This handles both same-repo and fork PRs via the pull/{number}/head ref. func (e *Executor) fetchAndCheckoutPR(ctx context.Context, t *session.Session, workDir, prRef, localBranch string, log *slog.Logger) error { - askPassEnv, cleanup, err := gitpkg.AskPassEnv(t.AccessToken) + if err := gitpkg.SanitizeRepoConfig(ctx, workDir, t.RepoURL); err != nil { + return fmt.Errorf("preparing workspace for PR fetch: %w", err) + } + + askPassEnv, cleanup, err := gitpkg.AskPassEnv(t.AccessToken, t.RepoURL) if err != nil { return fmt.Errorf("creating askpass for PR fetch: %w", err) } defer cleanup() - env := os.Environ() - if len(askPassEnv) > 0 { - env = append(env, askPassEnv...) - } - // git fetch origin pull/N/head:pr-N fetchRefSpec := fmt.Sprintf("%s:%s", prRef, localBranch) - fetchCmd := exec.CommandContext(ctx, "git", "fetch", "origin", fetchRefSpec) - fetchCmd.Dir = workDir - fetchCmd.Env = env + fetchCmd := gitpkg.Command(ctx, workDir, askPassEnv, "fetch", "--no-recurse-submodules", "origin", fetchRefSpec) var stderr strings.Builder fetchCmd.Stderr = &stderr @@ -1304,8 +1316,7 @@ func (e *Executor) fetchAndCheckoutPR(ctx context.Context, t *session.Session, w } // git checkout pr-N - checkoutCmd := exec.CommandContext(ctx, "git", "checkout", localBranch) - checkoutCmd.Dir = workDir + checkoutCmd := gitpkg.Command(ctx, workDir, nil, "checkout", localBranch) stderr.Reset() checkoutCmd.Stderr = &stderr @@ -1365,7 +1376,7 @@ func (e *Executor) handlePRReviewCompletion(ctx context.Context, t *session.Sess // Resolve token from provider key token := t.AccessToken - if token == "" && e.keyResolver != nil { + if token == "" && e.keyResolver != nil && t.UsesOperatorCredentials() { resolved, resolveErr := e.keyResolver.ResolveToken(ctx, t.RepoURL, "", t.ProviderKey) if resolveErr != nil { log.Error("pr_review: failed to resolve token for comment posting", "error", resolveErr) @@ -1652,7 +1663,7 @@ func (e *Executor) autoPostReview(ctx context.Context, t *session.Session, revie // autoPostReviewToPR posts review results to a specific PR number. func (e *Executor) autoPostReviewToPR(ctx context.Context, t *session.Session, prNumber int, reviewResult *review.ReviewResult, log *slog.Logger) { token := t.AccessToken - if token == "" && e.keyResolver != nil { + if token == "" && e.keyResolver != nil && t.UsesOperatorCredentials() { resolved, err := e.keyResolver.ResolveToken(ctx, t.RepoURL, "", t.ProviderKey) if err != nil { log.Error("auto-post: failed to resolve token", "error", err) diff --git a/internal/worker/tenant_credentials_test.go b/internal/worker/tenant_credentials_test.go new file mode 100644 index 0000000..0f00335 --- /dev/null +++ b/internal/worker/tenant_credentials_test.go @@ -0,0 +1,169 @@ +package worker + +import ( + "context" + "errors" + "log/slog" + "os" + "strings" + "testing" + + "github.com/freema/codeforge/internal/apperror" + "github.com/freema/codeforge/internal/keys" + "github.com/freema/codeforge/internal/session" + "github.com/freema/codeforge/internal/tool/mcp" + "github.com/freema/codeforge/internal/tools" +) + +// operatorKeys is a key registry holding the operator's credentials. +type operatorKeys struct{} + +func (operatorKeys) Create(context.Context, keys.Key) error { return nil } +func (operatorKeys) List(context.Context) ([]keys.Key, error) { + return []keys.Key{{Name: "operator-github", Provider: "github"}}, nil +} +func (operatorKeys) Delete(context.Context, string) error { return nil } +func (operatorKeys) Resolve(_ context.Context, _, name string) (string, error) { + if name == "operator-github" { + return "operator-registry-token", nil + } + return "", errors.New("not found") +} +func (operatorKeys) Verify(context.Context, string) (*keys.VerifyResult, string, error) { + return nil, "", nil +} +func (operatorKeys) ResolveByName(_ context.Context, name string) (string, string, error) { + if name == "operator-github" || name == "github-env" { + return "operator-registry-token", "github", nil + } + return "", "", errors.New("not found") +} +func (operatorKeys) ResolveFullByName(ctx context.Context, name string) (string, string, string, error) { + token, provider, err := operatorKeys{}.ResolveByName(ctx, name) + return token, provider, "", err +} + +// noTools is an empty tool registry; sessions fall back to the built-in catalog. +type noTools struct{} + +func (noTools) Create(context.Context, string, tools.ToolDefinition) error { return nil } +func (noTools) Get(_ context.Context, _, name string) (*tools.ToolDefinition, error) { + return nil, apperror.NotFound("tool %s not found", name) +} +func (noTools) List(context.Context, string) ([]tools.ToolDefinition, error) { return nil, nil } +func (noTools) Delete(context.Context, string, string) error { return nil } + +// operatorMCP holds one operator-registered MCP server carrying a secret. +type operatorMCP struct{} + +var operatorServer = mcp.Server{Name: "operator-sentry", Command: "npx", Package: "sentry-mcp", Env: map[string]string{"SENTRY_TOKEN": "operator-mcp-secret"}} + +func (operatorMCP) CreateGlobal(context.Context, mcp.Server) error { return nil } +func (operatorMCP) ListGlobal(context.Context) ([]mcp.Server, error) { + return []mcp.Server{operatorServer}, nil +} +func (operatorMCP) DeleteGlobal(context.Context, string) error { return nil } +func (operatorMCP) ResolveGlobal(context.Context, string) (*mcp.Server, error) { + return &operatorServer, nil +} +func (operatorMCP) CreateProject(context.Context, string, mcp.Server) error { return nil } +func (operatorMCP) ListProject(context.Context, string) ([]mcp.Server, error) { return nil, nil } +func (operatorMCP) DeleteProject(context.Context, string, string) error { return nil } +func (operatorMCP) ResolveMCPServers(_ context.Context, _ string, task []mcp.Server) ([]mcp.Server, error) { + return append([]mcp.Server{operatorServer}, task...), nil +} + +const tenantTestRepo = "https://github.com/acme/repo.git" + +func TestResolveToken_TenantSessionGetsNoOperatorToken(t *testing.T) { + t.Setenv("GITHUB_TOKEN", "operator-env-token") + e := &Executor{keyResolver: keys.NewResolver(operatorKeys{}, nil)} + log := slog.Default() + + tenantSession := &session.Session{RepoURL: tenantTestRepo, TenantID: "acme", ProviderKey: "operator-github"} + e.resolveToken(context.Background(), tenantSession, log) + if tenantSession.AccessToken != "" { + t.Errorf("tenant session got token %q", tenantSession.AccessToken) + } + + tenantNoKey := &session.Session{RepoURL: tenantTestRepo, TenantID: "acme"} + e.resolveToken(context.Background(), tenantNoKey, log) + if tenantNoKey.AccessToken != "" { + t.Errorf("tenant session got env token %q", tenantNoKey.AccessToken) + } + + own := &session.Session{RepoURL: tenantTestRepo, TenantID: "acme", AccessToken: "tenant-token"} + e.resolveToken(context.Background(), own, log) + if own.AccessToken != "tenant-token" { + t.Errorf("tenant's own token replaced with %q", own.AccessToken) + } + + operator := &session.Session{RepoURL: tenantTestRepo, ProviderKey: "operator-github"} + e.resolveToken(context.Background(), operator, log) + if operator.AccessToken != "operator-registry-token" { + t.Errorf("operator session token = %q, want the registry token", operator.AccessToken) + } +} + +func TestSetupMCP_TenantSessionGetsNoOperatorCredentials(t *testing.T) { + t.Setenv("GITHUB_TOKEN", "operator-env-token") + e := &Executor{ + mcpInstaller: mcp.NewInstaller(operatorMCP{}), + toolResolver: tools.NewResolver(noTools{}, operatorKeys{}), + } + log := slog.Default() + + readConfig := func(t *testing.T, dir string) string { + t.Helper() + out, err := os.ReadFile(mcp.ConfigPath(dir, "")) + if err != nil { + t.Fatalf("reading MCP config: %v", err) + } + return string(out) + } + + t.Run("operator registry servers and tool auto-fill are not used", func(t *testing.T) { + dir := t.TempDir() + tenantSession := &session.Session{RepoURL: tenantTestRepo, TenantID: "acme", Config: &session.Config{ + MCPServers: []session.MCPServer{{Name: "own", Command: "npx", Package: "own-mcp"}}, + }} + if _, err := e.setupMCP(context.Background(), tenantSession, dir, log); err != nil { + t.Fatalf("setupMCP: %v", err) + } + cfg := readConfig(t, dir) + if !strings.Contains(cfg, "own-mcp") { + t.Errorf("tenant's own server missing from %s", cfg) + } + if strings.Contains(cfg, "operator-sentry") || strings.Contains(cfg, "operator-mcp-secret") { + t.Errorf("operator server written into tenant workspace: %s", cfg) + } + }) + + t.Run("tool config is not auto-filled from operator keys", func(t *testing.T) { + dir := t.TempDir() + tenantSession := &session.Session{RepoURL: tenantTestRepo, TenantID: "acme", Config: &session.Config{ + Tools: []tools.SessionTool{{Name: "github"}}, + }} + _, err := e.setupMCP(context.Background(), tenantSession, dir, log) + if err == nil || !strings.Contains(err.Error(), "token") { + t.Fatalf("err = %v, want tool resolution to fail on the missing token", err) + } + if out, readErr := os.ReadFile(mcp.ConfigPath(dir, "")); readErr == nil && strings.Contains(string(out), "operator-registry-token") { + t.Errorf("operator key written into tenant workspace: %s", out) + } + }) + + t.Run("operator sessions keep registry servers and auto-fill", func(t *testing.T) { + dir := t.TempDir() + operator := &session.Session{RepoURL: tenantTestRepo, Config: &session.Config{ + Tools: []tools.SessionTool{{Name: "github"}}, + }} + if _, err := e.setupMCP(context.Background(), operator, dir, log); err != nil { + t.Fatalf("setupMCP: %v", err) + } + cfg := readConfig(t, dir) + if !strings.Contains(cfg, "operator-sentry") || !strings.Contains(cfg, "operator-registry-token") { + t.Errorf("operator session lost registry servers or auto-fill: %s", cfg) + } + }) +}