diff --git a/internal/storage/webdav.go b/internal/storage/webdav.go index 15d0cd1d..5e080e66 100644 --- a/internal/storage/webdav.go +++ b/internal/storage/webdav.go @@ -7,6 +7,7 @@ import ( "context" "fmt" "io" + "net/http" "path" "strings" @@ -14,40 +15,62 @@ import ( "github.com/studio-b12/gowebdav" ) +type contextTransport struct { + ctx context.Context + parent http.RoundTripper +} + +func (t *contextTransport) RoundTrip(req *http.Request) (*http.Response, error) { + return t.parent.RoundTrip(req.WithContext(t.ctx)) +} + type webDAVBackend struct { - client *gowebdav.Client + endpoint string + username string + password string basePath string } func newWebDAVBackend(cfg WebDAVConfig) (*webDAVBackend, error) { - client := gowebdav.NewClient(strings.TrimRight(cfg.Endpoint, "/"), cfg.Username, cfg.Password) - client.SetTransport(httppool.DefaultTransport()) return &webDAVBackend{ - client: client, + endpoint: strings.TrimRight(cfg.Endpoint, "/"), + username: cfg.Username, + password: cfg.Password, basePath: strings.Trim(cfg.BasePath, "/"), }, nil } -func (b *webDAVBackend) Put(_ context.Context, key string, body io.Reader, size int64, _ string) (PutResult, error) { +func (b *webDAVBackend) newClient(ctx context.Context) *gowebdav.Client { + client := gowebdav.NewClient(b.endpoint, b.username, b.password) + client.SetTransport(&contextTransport{ + ctx: ctx, + parent: httppool.DefaultTransport(), + }) + return client +} + +func (b *webDAVBackend) Put(ctx context.Context, key string, body io.Reader, size int64, _ string) (PutResult, error) { key = b.key(key) + client := b.newClient(ctx) if dir := path.Dir(key); dir != "." && dir != "/" { - if err := b.client.MkdirAll(dir, storageDirPerm); err != nil { + if err := client.MkdirAll(dir, storageDirPerm); err != nil { return PutResult{}, fmt.Errorf("create WebDAV directory: %w", err) } } - if err := b.client.WriteStreamWithLength(key, body, size, storageFilePerm); err != nil { + if err := client.WriteStreamWithLength(key, body, size, storageFilePerm); err != nil { return PutResult{}, fmt.Errorf("put WebDAV object: %w", err) } return PutResult{Key: key}, nil } -func (b *webDAVBackend) Get(_ context.Context, key string) (*Object, error) { +func (b *webDAVBackend) Get(ctx context.Context, key string) (*Object, error) { key = b.key(key) - info, err := b.client.Stat(key) + client := b.newClient(ctx) + info, err := client.Stat(key) if err != nil { return nil, fmt.Errorf("stat WebDAV object: %w", err) } - body, err := b.client.ReadStream(key) + body, err := client.ReadStream(key) if err != nil { return nil, fmt.Errorf("get WebDAV object: %w", err) } @@ -58,15 +81,17 @@ func (b *webDAVBackend) Get(_ context.Context, key string) (*Object, error) { return &Object{Body: body, ContentLength: info.Size(), ContentType: contentType}, nil } -func (b *webDAVBackend) Delete(_ context.Context, key string) error { - if err := b.client.Remove(b.key(key)); err != nil { +func (b *webDAVBackend) Delete(ctx context.Context, key string) error { + client := b.newClient(ctx) + if err := client.Remove(b.key(key)); err != nil { return fmt.Errorf("delete WebDAV object: %w", err) } return nil } -func (b *webDAVBackend) Test(_ context.Context) error { - if err := b.client.Connect(); err != nil { +func (b *webDAVBackend) Test(ctx context.Context) error { + client := b.newClient(ctx) + if err := client.Connect(); err != nil { return fmt.Errorf("connect WebDAV: %w", err) } return nil diff --git a/pkg/logger/utils.go b/pkg/logger/utils.go index 83aa2592..d1a48bdf 100644 --- a/pkg/logger/utils.go +++ b/pkg/logger/utils.go @@ -89,6 +89,9 @@ func getLogLevelForConfig(cfg Config) zapcore.Level { func getTraceIDFields(ctx context.Context) []zap.Field { span := trace.SpanFromContext(ctx) spanContext := span.SpanContext() + if !spanContext.IsValid() { + return nil + } return []zap.Field{ zap.String("traceID", spanContext.TraceID().String()), zap.String("spanID", spanContext.SpanID().String()), diff --git a/pkg/trace/sampler.go b/pkg/trace/sampler.go index d3efbbc0..388b400b 100644 --- a/pkg/trace/sampler.go +++ b/pkg/trace/sampler.go @@ -8,11 +8,11 @@ import ( sdktrace "go.opentelemetry.io/otel/sdk/trace" ) -// ParentBasedErrorAwareSampler 创建父级感知的概率采样器 +// ParentBasedRatioSampler 创建父级感知的概率采样器 // - 如果父 Span 已采样,则子 Span 也采样 // - 如果父 Span 未采样,则子 Span 也不采样 // - 如果是根 Span,按 samplingRate 概率采样 -func ParentBasedErrorAwareSampler(samplingRate float64) sdktrace.Sampler { +func ParentBasedRatioSampler(samplingRate float64) sdktrace.Sampler { return sdktrace.ParentBased( sdktrace.TraceIDRatioBased(samplingRate), ) diff --git a/pkg/trace/trace_provider.go b/pkg/trace/trace_provider.go index 8aeb2e34..4f34df04 100644 --- a/pkg/trace/trace_provider.go +++ b/pkg/trace/trace_provider.go @@ -47,7 +47,7 @@ func newTracerProvider(cfg Config) (*sdktrace.TracerProvider, error) { tracerProvider := sdktrace.NewTracerProvider( sdktrace.WithBatcher(traceExporter), sdktrace.WithResource(r), - sdktrace.WithSampler(ParentBasedErrorAwareSampler(cfg.SamplingRate)), + sdktrace.WithSampler(ParentBasedRatioSampler(cfg.SamplingRate)), ) return tracerProvider, nil }