fix: preserve instrumented client contracts

Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
bestony
2026-07-14 01:29:30 +08:00
co-authored by multica-agent
parent 54d228dda5
commit 966afd1b4b
5 changed files with 28 additions and 28 deletions
+8 -10
View File
@@ -280,8 +280,6 @@ func NewClient(proxyURL string) (*Client, error) {
}
client.Transport = transport
}
client.Transport = servertiming.WrapRoundTripper(client.Transport)
return &Client{
httpClient: client,
}, nil
@@ -343,7 +341,7 @@ func (c *Client) ExchangeCode(ctx context.Context, code, codeVerifier string) (*
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("token 交换请求失败: %w", err)
}
@@ -385,7 +383,7 @@ func (c *Client) RefreshToken(ctx context.Context, refreshToken string) (*TokenR
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("token 刷新请求失败: %w", err)
}
@@ -416,7 +414,7 @@ func (c *Client) GetUserInfo(ctx context.Context, accessToken string) (*UserInfo
}
req.Header.Set("Authorization", "Bearer "+accessToken)
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("用户信息请求失败: %w", err)
}
@@ -467,7 +465,7 @@ func (c *Client) LoadCodeAssist(ctx context.Context, accessToken string) (*LoadC
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
lastErr = fmt.Errorf("loadCodeAssist 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -546,7 +544,7 @@ func (c *Client) OnboardUser(ctx context.Context, accessToken, tierID string) (s
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
lastErr = fmt.Errorf("onboardUser 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -685,7 +683,7 @@ func (c *Client) FetchAvailableModels(ctx context.Context, accessToken, projectI
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
resp, err := fetchClient.Do(req)
resp, err := servertiming.Do(fetchClient, req)
if err != nil {
lastErr = fmt.Errorf("fetchAvailableModels 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -844,7 +842,7 @@ func (c *Client) SetUserSettings(ctx context.Context, accessToken string) (*SetU
req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1")
req.Host = "daily-cloudcode-pa.googleapis.com"
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("setUserSettings 请求失败: %w", err)
}
@@ -887,7 +885,7 @@ func (c *Client) FetchUserInfo(ctx context.Context, accessToken, projectID strin
req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1")
req.Host = "daily-cloudcode-pa.googleapis.com"
resp, err := c.httpClient.Do(req)
resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("fetchUserInfo 请求失败: %w", err)
}
@@ -276,9 +276,9 @@ func normalizeMetricName(name string) string {
}
switch {
case r >= 'a' && r <= 'z', r >= '0' && r <= '9':
b.WriteRune(r)
_, _ = b.WriteRune(r)
case r == '_' || r == '-':
b.WriteByte('_')
_ = b.WriteByte('_')
}
}
return strings.Trim(b.String(), "_")
@@ -120,13 +120,10 @@ func TestCollectorConcurrentRecording(t *testing.T) {
}
func TestContextHelpersHandleMissingCollector(t *testing.T) {
if Active(nil) || Active(context.Background()) {
if Active(context.Background()) {
t.Fatal("context without collector reported active")
}
if got := HeaderValue(context.Background(), time.Now(), "hit"); got != "" {
t.Fatalf("HeaderValue() = %q without collector, want empty", got)
}
if got := WithCollector(nil, nil); got == nil {
t.Fatal("WithCollector(nil, nil) returned nil context")
}
}
@@ -103,7 +103,9 @@ func (c *serverTimingConn) BeginTx(ctx context.Context, opts driver.TxOptions) (
if opts.ReadOnly {
return nil, errors.New("driver does not support read-only transactions")
}
tx, err = c.Conn.Begin()
// The wrapper exposes ConnBeginTx, so it must retain database/sql's
// legacy fallback for drivers that only implement Conn.Begin.
tx, err = c.Conn.Begin() //nolint:staticcheck // Required driver compatibility fallback.
}
servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now())
if err != nil || tx == nil {
@@ -162,7 +164,9 @@ func (s *serverTimingStmt) ExecContext(ctx context.Context, args []driver.NamedV
var values []driver.Value
values, err = namedValues(args)
if err == nil {
result, err = s.Stmt.Exec(values)
// The wrapper exposes StmtExecContext and must preserve the fallback
// database/sql would use for a legacy driver statement.
result, err = s.Stmt.Exec(values) //nolint:staticcheck // Required driver compatibility fallback.
}
}
servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
@@ -181,7 +185,9 @@ func (s *serverTimingStmt) QueryContext(ctx context.Context, args []driver.Named
var values []driver.Value
values, err = namedValues(args)
if err == nil {
rows, err = s.Stmt.Query(values)
// The wrapper exposes StmtQueryContext and must preserve the fallback
// database/sql would use for a legacy driver statement.
rows, err = s.Stmt.Query(values) //nolint:staticcheck // Required driver compatibility fallback.
}
}
servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
@@ -198,13 +204,6 @@ func (s *serverTimingStmt) CheckNamedValue(value *driver.NamedValue) error {
return driver.ErrSkip
}
func (s *serverTimingStmt) ColumnConverter(index int) driver.ValueConverter {
if converter, ok := s.Stmt.(driver.ColumnConverter); ok {
return converter.ColumnConverter(index)
}
return driver.DefaultParameterConverter
}
func namedValues(args []driver.NamedValue) ([]driver.Value, error) {
values := make([]driver.Value, len(args))
for i, arg := range args {
@@ -159,7 +159,10 @@ func TestServerTimingConnectorRecordsDriverCallsWithoutRowLifetime(t *testing.T)
if err != nil {
t.Fatal(err)
}
conn := rawConn.(*serverTimingConn)
conn, ok := rawConn.(*serverTimingConn)
if !ok {
t.Fatalf("Connect() returned %T, want *serverTimingConn", rawConn)
}
if _, err := conn.ExecContext(ctx, "sensitive update", nil); err != nil {
t.Fatal(err)
@@ -203,7 +206,10 @@ func TestServerTimingPreparedStatementsAndTransactions(t *testing.T) {
if err != nil {
t.Fatal(err)
}
timedStmt := stmt.(*serverTimingStmt)
timedStmt, ok := stmt.(*serverTimingStmt)
if !ok {
t.Fatalf("PrepareContext() returned %T, want *serverTimingStmt", stmt)
}
if _, err := timedStmt.ExecContext(ctx, nil); err != nil {
t.Fatal(err)
}