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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 0 additions & 17 deletions internal/auth/header.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,16 +36,13 @@ package auth

import (
"errors"
"fmt"
"strings"

"github.com/github/gh-aw-mcpg/internal/logger"
"github.com/github/gh-aw-mcpg/internal/sanitize"
"github.com/github/gh-aw-mcpg/internal/util"
)

var logAuth = logger.ForFile()
var logAPIKey = logger.New("auth:apikey")

var (
// ErrMissingAuthHeader is returned when the Authorization header is missing
Expand Down Expand Up @@ -200,17 +197,3 @@ func IsMalformedHeader(header string) bool {
}
return false
}

// GenerateRandomAgentID generates a cryptographically random agent ID.
// Per spec §7.3, the gateway SHOULD generate a random agent ID on startup
// if none is provided. Returns a 32-byte hex-encoded string (64 chars).
func GenerateRandomAgentID() (string, error) {
logAPIKey.Print("Generating random agent ID")
key, err := util.RandomHex(32)
if err != nil {
logAPIKey.Printf("Random agent ID generation failed: %v", err)
return "", fmt.Errorf("failed to generate random agent ID: %w", err)
}
logAPIKey.Print("Random agent ID generated successfully")
return key, nil
}
49 changes: 0 additions & 49 deletions internal/auth/header_test.go
Original file line number Diff line number Diff line change
@@ -1,8 +1,6 @@
package auth

import (
"crypto/rand"
"errors"
"testing"

"github.com/stretchr/testify/assert"
Expand Down Expand Up @@ -596,50 +594,3 @@ func TestStripAuthScheme(t *testing.T) {
})
}
}

// errorReader is a test helper io.Reader that always returns the configured error.
type errorReader struct {
err error
}

func (r *errorReader) Read(_ []byte) (int, error) {
return 0, r.err
}

// TestGenerateRandomAgentID_RandomFailure verifies that GenerateRandomAgentID
// correctly wraps and propagates errors from the underlying random source.
// This test must NOT run in parallel because it temporarily replaces the
// global crypto/rand.Reader.
func TestGenerateRandomAgentID_RandomFailure(t *testing.T) {
syntheticErr := errors.New("synthetic entropy failure")

origReader := rand.Reader
rand.Reader = &errorReader{err: syntheticErr}
defer func() { rand.Reader = origReader }()

key, err := GenerateRandomAgentID()

assert.Empty(t, key, "key should be empty when random generation fails")
require.Error(t, err, "should return an error when the random source fails")
assert.ErrorIs(t, err, syntheticErr, "error should wrap the underlying source error")
assert.Contains(t, err.Error(), "failed to generate random agent ID",
"error message should describe the failure context")
}

// TestGenerateRandomAgentID_RecoveryAfterFailure verifies that
// GenerateRandomAgentID works correctly after the random source is restored,
// confirming that no state is leaked between calls.
// This test must NOT run in parallel because it temporarily replaces the
// global crypto/rand.Reader.
func TestGenerateRandomAgentID_RecoveryAfterFailure(t *testing.T) {
origReader := rand.Reader
rand.Reader = &errorReader{err: errors.New("transient failure")}
_, err := GenerateRandomAgentID()
require.Error(t, err, "should fail with broken reader")

// Restore and verify subsequent call succeeds.
rand.Reader = origReader
key, err := GenerateRandomAgentID()
require.NoError(t, err, "should succeed after reader is restored")
assert.Len(t, key, 64, "restored call should return 64-char hex key")
}
27 changes: 27 additions & 0 deletions internal/auth/id.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
package auth

import (
"fmt"

"github.com/github/gh-aw-mcpg/internal/logger"
"github.com/github/gh-aw-mcpg/internal/util"
)

// logAPIKey is the debug logger for API-key / agent-ID generation.
// It uses the custom namespace "auth:apikey" so callers can filter these
// debug logs independently with DEBUG=auth:apikey.
var logAPIKey = logger.New("auth:apikey")

// GenerateRandomAgentID generates a cryptographically random agent ID.
// Per spec §7.3, the gateway SHOULD generate a random agent ID on startup
// if none is provided. Returns a 32-byte hex-encoded string (64 chars).
func GenerateRandomAgentID() (string, error) {

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done in 3ffc166 — moved the ID-generation failure/recovery tests and the errorReader helper into a new internal/auth/id_test.go, and dropped the now-unused crypto/rand and errors imports from header_test.go.

logAPIKey.Print("Generating random agent ID")
key, err := util.RandomHex(32)
if err != nil {
logAPIKey.Printf("Random agent ID generation failed: %v", err)
return "", fmt.Errorf("failed to generate random agent ID: %w", err)
}
logAPIKey.Print("Random agent ID generated successfully")
return key, nil
}
57 changes: 57 additions & 0 deletions internal/auth/id_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
package auth

import (
"crypto/rand"
"errors"
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

// errorReader is a test helper io.Reader that always returns the configured error.
type errorReader struct {
err error
}

func (r *errorReader) Read(_ []byte) (int, error) {
return 0, r.err
}

// TestGenerateRandomAgentID_RandomFailure verifies that GenerateRandomAgentID
// correctly wraps and propagates errors from the underlying random source.
// This test must NOT run in parallel because it temporarily replaces the
// global crypto/rand.Reader.
func TestGenerateRandomAgentID_RandomFailure(t *testing.T) {
syntheticErr := errors.New("synthetic entropy failure")

origReader := rand.Reader
rand.Reader = &errorReader{err: syntheticErr}
defer func() { rand.Reader = origReader }()

key, err := GenerateRandomAgentID()

assert.Empty(t, key, "key should be empty when random generation fails")
require.Error(t, err, "should return an error when the random source fails")
assert.ErrorIs(t, err, syntheticErr, "error should wrap the underlying source error")
assert.Contains(t, err.Error(), "failed to generate random agent ID",
"error message should describe the failure context")
}

// TestGenerateRandomAgentID_RecoveryAfterFailure verifies that
// GenerateRandomAgentID works correctly after the random source is restored,
// confirming that no state is leaked between calls.
// This test must NOT run in parallel because it temporarily replaces the
// global crypto/rand.Reader.
func TestGenerateRandomAgentID_RecoveryAfterFailure(t *testing.T) {
origReader := rand.Reader
rand.Reader = &errorReader{err: errors.New("transient failure")}
_, err := GenerateRandomAgentID()
require.Error(t, err, "should fail with broken reader")

// Restore and verify subsequent call succeeds.
rand.Reader = origReader
key, err := GenerateRandomAgentID()
require.NoError(t, err, "should succeed after reader is restored")
assert.Len(t, key, 64, "restored call should return 64-char hex key")
}
31 changes: 31 additions & 0 deletions internal/logger/fileutil.go
Original file line number Diff line number Diff line change
Expand Up @@ -123,3 +123,34 @@ func writeJSONToFile(logDir, fileName string, data any, perm os.FileMode) error
}
return atomicWriteFile(filepath.Join(logDir, fileName), jsonData, perm)
}

// jsonFileSink holds the shared state common to stateful JSON-file loggers
// (logDir, fileName, useFallback). Embed this struct in logger types that
// persist an in-memory data structure to a JSON file so the three repeated
// fields and the writeJSON helper do not have to be duplicated.
//
// Usage:
//
// type MyLogger struct {
// lockable
// jsonFileSink
// data MyData
// }
//
// The embedded jsonFileSink.writeJSON method can then be called from
// writeToFile to write data to the configured JSON file:
//
// func (l *MyLogger) writeToFile() error {
// return l.writeJSON(l.data, 0644)
// }
type jsonFileSink struct {
logDir string
fileName string
useFallback bool
}

// writeJSON marshals data as indented JSON and atomically writes it to the
// file at s.logDir/s.fileName with the given permissions.
func (s *jsonFileSink) writeJSON(data any, perm os.FileMode) error {
return writeJSONToFile(s.logDir, s.fileName, data, perm)
}
5 changes: 2 additions & 3 deletions internal/logger/logger_namespace_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,9 @@ func TestLoggerNamespacesMatchFileConventions(t *testing.T) {
internalRoot := filepath.Join(repoRoot, "internal")

exceptionNamespaces := map[string][]string{
// header.go defines two loggers: one for general auth (auto-derived via ForFile as
// "auth:header") and one for API-key auth which uses the custom namespace "auth:apikey"
// id.go defines the logAPIKey logger with the custom namespace "auth:apikey"
// so callers can filter API-key debug logs independently with DEBUG=auth:apikey.
"internal/auth/header.go": {"auth:apikey"},
"internal/auth/id.go": {"auth:apikey"},

// The following files use intentionally shorter or semantically clearer namespaces
// instead of the full file-name-derived form. These are preserved for backward
Expand Down
19 changes: 7 additions & 12 deletions internal/logger/observed_url_domains_logger.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,10 +27,8 @@ func URLDomainAuditEnabled() bool {
// ObservedURLDomainsLogger manages unique observed URL domains grouped by server ID.
type ObservedURLDomainsLogger struct {
lockable
logDir string
fileName string
data map[string]map[string]struct{}
useFallback bool
jsonFileSink
data map[string]map[string]struct{}
}

var (
Expand All @@ -46,9 +44,8 @@ var observedURLDomainsLoggerFactory = newLoggerFactory(
}

l := &ObservedURLDomainsLogger{
logDir: logDir,
fileName: fileName,
data: make(map[string]map[string]struct{}),
jsonFileSink: jsonFileSink{logDir: logDir, fileName: fileName},
data: make(map[string]map[string]struct{}),
}
if err := l.writeToFile(); err != nil {
return nil, err
Expand All @@ -58,10 +55,8 @@ var observedURLDomainsLoggerFactory = newLoggerFactory(
},
func(err error, logDir, fileName string) (*ObservedURLDomainsLogger, error) {
return fallbackLoggerOnInitError(err, "Failed to initialize observed URL domains log file", "Observed URL domains logging disabled", &ObservedURLDomainsLogger{
logDir: logDir,
fileName: fileName,
data: make(map[string]map[string]struct{}),
useFallback: true,
jsonFileSink: jsonFileSink{logDir: logDir, fileName: fileName, useFallback: true},
data: make(map[string]map[string]struct{}),
})
},
)
Expand Down Expand Up @@ -112,7 +107,7 @@ func (l *ObservedURLDomainsLogger) writeToFile() error {
for serverID, domains := range l.data {
serialized[serverID] = util.SortedSetKeys(domains)
}
return writeJSONToFile(l.logDir, l.fileName, serialized, 0600)
return l.writeJSON(serialized, 0600)
}

func (l *ObservedURLDomainsLogger) Close() error { return nil }
Expand Down
8 changes: 4 additions & 4 deletions internal/logger/observed_url_domains_logger_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -155,8 +155,8 @@ func TestLogDomains_NilDomains(t *testing.T) {

func TestLogDomains_FallbackMode_ReturnsNil(t *testing.T) {
l := &ObservedURLDomainsLogger{
data: make(map[string]map[string]struct{}),
useFallback: true,
data: make(map[string]map[string]struct{}),
jsonFileSink: jsonFileSink{useFallback: true},
}

// In fallback mode LogDomains should silently succeed without writing.
Expand Down Expand Up @@ -320,8 +320,8 @@ func TestLogObservedURLDomains_FallbackMode_NoPanic(t *testing.T) {
globalObservedURLDomainsMu.Lock()
prev := globalObservedURLDomainsLogger
globalObservedURLDomainsLogger = &ObservedURLDomainsLogger{
data: make(map[string]map[string]struct{}),
useFallback: true,
data: make(map[string]map[string]struct{}),
jsonFileSink: jsonFileSink{useFallback: true},
}
globalObservedURLDomainsMu.Unlock()
t.Cleanup(func() {
Expand Down
15 changes: 5 additions & 10 deletions internal/logger/tools_logger.go
Original file line number Diff line number Diff line change
Expand Up @@ -26,10 +26,8 @@ type ToolsData struct {
// ToolsLogger manages logging of MCP server tools to a JSON file
type ToolsLogger struct {
lockable
logDir string
fileName string
data *ToolsData
useFallback bool
jsonFileSink
data *ToolsData
}

var (
Expand All @@ -47,8 +45,7 @@ var toolsLoggerFactory = newLoggerFactory(
}

tl := &ToolsLogger{
logDir: logDir,
fileName: fileName,
jsonFileSink: jsonFileSink{logDir: logDir, fileName: fileName},
data: &ToolsData{
Servers: make(map[string][]ToolInfo),
},
Expand All @@ -58,9 +55,7 @@ var toolsLoggerFactory = newLoggerFactory(
},
func(err error, logDir, fileName string) (*ToolsLogger, error) {
return fallbackLoggerOnInitError(err, "Failed to initialize tools log file", "Tools logging disabled", &ToolsLogger{
logDir: logDir,
fileName: fileName,
useFallback: true,
jsonFileSink: jsonFileSink{logDir: logDir, fileName: fileName, useFallback: true},
data: &ToolsData{
Servers: make(map[string][]ToolInfo),
},
Expand Down Expand Up @@ -92,7 +87,7 @@ func (tl *ToolsLogger) LogTools(serverID string, tools []ToolInfo) error {
// writeToFile writes the current tools data to the JSON file.
// Caller must hold tl.mu lock.
func (tl *ToolsLogger) writeToFile() error {
return writeJSONToFile(tl.logDir, tl.fileName, tl.data, 0644)
return tl.writeJSON(tl.data, 0644)
}

// Close is a no-op for ToolsLogger (implements closableLogger interface)
Expand Down
24 changes: 9 additions & 15 deletions internal/logger/tools_logger_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -221,8 +221,7 @@ func TestWriteToFile_Success(t *testing.T) {

tmpDir := t.TempDir()
tl := &ToolsLogger{
logDir: tmpDir,
fileName: "tools.json",
jsonFileSink: jsonFileSink{logDir: tmpDir, fileName: "tools.json"},
data: &ToolsData{
Servers: map[string][]ToolInfo{
"server1": {
Expand Down Expand Up @@ -252,9 +251,8 @@ func TestWriteToFile_WriteFileFails(t *testing.T) {
assert := assert.New(t)

tl := &ToolsLogger{
logDir: "/nonexistent/dir/that/does/not/exist",
fileName: "tools.json",
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
jsonFileSink: jsonFileSink{logDir: "/nonexistent/dir/that/does/not/exist", fileName: "tools.json"},
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
}

err := tl.writeToFile()
Expand All @@ -275,9 +273,8 @@ func TestWriteToFile_RenameFails(t *testing.T) {
require.NoError(os.MkdirAll(targetPath, 0755))

tl := &ToolsLogger{
logDir: tmpDir,
fileName: "tools.json",
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
jsonFileSink: jsonFileSink{logDir: tmpDir, fileName: "tools.json"},
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
}

err := tl.writeToFile()
Expand All @@ -299,9 +296,8 @@ func TestLogToolsForServer_ErrorIsLogged(t *testing.T) {
oldLogger := globalToolsLogger
// Point the global logger at a nonexistent directory so writeToFile fails.
globalToolsLogger = &ToolsLogger{
logDir: "/nonexistent/path/for/test",
fileName: "tools.json",
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
jsonFileSink: jsonFileSink{logDir: "/nonexistent/path/for/test", fileName: "tools.json"},
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
}
globalToolsMu.Unlock()
t.Cleanup(func() {
Expand All @@ -325,10 +321,8 @@ func TestLogToolsForServer_FallbackSkipsErrors(t *testing.T) {
globalToolsMu.Lock()
oldLogger := globalToolsLogger
globalToolsLogger = &ToolsLogger{
logDir: "/nonexistent/path",
fileName: "tools.json",
useFallback: true,
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
jsonFileSink: jsonFileSink{logDir: "/nonexistent/path", fileName: "tools.json", useFallback: true},
data: &ToolsData{Servers: make(map[string][]ToolInfo)},
}
globalToolsMu.Unlock()
t.Cleanup(func() {
Expand Down
Loading
Loading