diff --git a/go.mod b/go.mod index 18362a8..309474f 100644 --- a/go.mod +++ b/go.mod @@ -10,14 +10,14 @@ require ( github.com/redis/go-redis/v9 v9.18.0 github.com/testcontainers/testcontainers-go/modules/postgres v0.42.0 github.com/testcontainers/testcontainers-go/modules/redis v0.42.0 - golang.org/x/oauth2 v0.30.0 - google.golang.org/api v0.247.0 + golang.org/x/oauth2 v0.36.0 + google.golang.org/api v0.276.0 ) require ( - cloud.google.com/go/auth v0.16.4 // indirect + cloud.google.com/go/auth v0.20.0 // indirect cloud.google.com/go/auth/oauth2adapt v0.2.8 // indirect - cloud.google.com/go/compute/metadata v0.8.0 // indirect + cloud.google.com/go/compute/metadata v0.9.0 // indirect dario.cat/mergo v1.0.2 // indirect github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect github.com/Microsoft/go-winio v0.6.2 // indirect @@ -39,8 +39,8 @@ require ( github.com/go-logr/stdr v1.2.2 // indirect github.com/go-ole/go-ole v1.2.6 // indirect github.com/google/s2a-go v0.1.9 // indirect - github.com/googleapis/enterprise-certificate-proxy v0.3.6 // indirect - github.com/googleapis/gax-go/v2 v2.15.0 // indirect + github.com/googleapis/enterprise-certificate-proxy v0.3.14 // indirect + github.com/googleapis/gax-go/v2 v2.21.0 // indirect github.com/jackc/pgerrcode v0.0.0-20220416144525-469b46aa5efa // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect @@ -73,17 +73,16 @@ require ( go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.67.0 // indirect go.opentelemetry.io/otel v1.43.0 // indirect go.opentelemetry.io/otel/metric v1.43.0 // indirect - go.opentelemetry.io/otel/sdk v1.43.0 // indirect go.opentelemetry.io/otel/sdk/metric v1.43.0 // indirect go.opentelemetry.io/otel/trace v1.43.0 // indirect go.uber.org/atomic v1.11.0 // indirect golang.org/x/crypto v0.49.0 // indirect - golang.org/x/net v0.51.0 // indirect + golang.org/x/net v0.52.0 // indirect golang.org/x/sync v0.20.0 // indirect golang.org/x/sys v0.42.0 // indirect golang.org/x/text v0.35.0 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20250818200422-3122310a409c // indirect - google.golang.org/grpc v1.74.2 // indirect - google.golang.org/protobuf v1.36.7 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260401024825-9d38bb4040a9 // indirect + google.golang.org/grpc v1.80.0 // indirect + google.golang.org/protobuf v1.36.11 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/go.sum b/go.sum index ba87ef8..9e08f9a 100644 --- a/go.sum +++ b/go.sum @@ -1,10 +1,9 @@ -cloud.google.com/go v0.121.6 h1:waZiuajrI28iAf40cWgycWNgaXPO06dupuS+sgibK6c= -cloud.google.com/go/auth v0.16.4 h1:fXOAIQmkApVvcIn7Pc2+5J8QTMVbUGLscnSVNl11su8= -cloud.google.com/go/auth v0.16.4/go.mod h1:j10ncYwjX/g3cdX7GpEzsdM+d+ZNsXAbb6qXA7p1Y5M= +cloud.google.com/go/auth v0.20.0 h1:kXTssoVb4azsVDoUiF8KvxAqrsQcQtB53DcSgta74CA= +cloud.google.com/go/auth v0.20.0/go.mod h1:942/yi/itH1SsmpyrbnTMDgGfdy2BUqIKyd0cyYLc5Q= cloud.google.com/go/auth/oauth2adapt v0.2.8 h1:keo8NaayQZ6wimpNSmW5OPc283g65QNIiLpZnkHRbnc= cloud.google.com/go/auth/oauth2adapt v0.2.8/go.mod h1:XQ9y31RkqZCcwJWNSx2Xvric3RrU88hAYYbjDWYDL+c= -cloud.google.com/go/compute/metadata v0.8.0 h1:HxMRIbao8w17ZX6wBnjhcDkW6lTFpgcaobyVfZWqRLA= -cloud.google.com/go/compute/metadata v0.8.0/go.mod h1:sYOGTp851OV9bOFJ9CH7elVvyzopvWQFNNghtDQ/Biw= +cloud.google.com/go/compute/metadata v0.9.0 h1:pDUj4QMoPejqq20dK0Pg2N4yG9zIkYGdBtwLoEkH9Zs= +cloud.google.com/go/compute/metadata v0.9.0/go.mod h1:E0bWwX5wTnLPedCKqk3pJmVgCBSM6qQI1yTBdEb3C10= dario.cat/mergo v1.0.2 h1:85+piFYR1tMbRrLcDwR18y4UKJ3aH1Tbzi24VRW1TK8= dario.cat/mergo v1.0.2/go.mod h1:E/hbnu0NxMFBjpMIE34DRGLWqDy0g5FuKDhCb31ngxA= github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6 h1:He8afgbRMd7mFxO99hRNu+6tazq8nFF9lIwo9JFroBk= @@ -65,6 +64,8 @@ github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q= github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= github.com/golang-migrate/migrate/v4 v4.19.1 h1:OCyb44lFuQfYXYLx1SCxPZQGU7mcaZ7gH9yH4jSFbBA= github.com/golang-migrate/migrate/v4 v4.19.1/go.mod h1:CTcgfjxhaUtsLipnLoQRWCrjYXycRz/g5+RWDuYgPrE= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= @@ -72,10 +73,10 @@ github.com/google/s2a-go v0.1.9 h1:LGD7gtMgezd8a/Xak7mEWL0PjoTQFvpRudN895yqKW0= github.com/google/s2a-go v0.1.9/go.mod h1:YA0Ei2ZQL3acow2O62kdp9UlnvMmU7kA6Eutn0dXayM= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= -github.com/googleapis/enterprise-certificate-proxy v0.3.6 h1:GW/XbdyBFQ8Qe+YAmFU9uHLo7OnF5tL52HFAgMmyrf4= -github.com/googleapis/enterprise-certificate-proxy v0.3.6/go.mod h1:MkHOF77EYAE7qfSuSS9PU6g4Nt4e11cnsDUowfwewLA= -github.com/googleapis/gax-go/v2 v2.15.0 h1:SyjDc1mGgZU5LncH8gimWo9lW1DtIfPibOG81vgd/bo= -github.com/googleapis/gax-go/v2 v2.15.0/go.mod h1:zVVkkxAQHa1RQpg9z2AUCMnKhi0Qld9rcmyfL1OZhoc= +github.com/googleapis/enterprise-certificate-proxy v0.3.14 h1:yh8ncqsbUY4shRD5dA6RlzjJaT4hi3kII+zYw8wmLb8= +github.com/googleapis/enterprise-certificate-proxy v0.3.14/go.mod h1:vqVt9yG9480NtzREnTlmGSBmFrA+bzb0yl0TxoBQXOg= +github.com/googleapis/gax-go/v2 v2.21.0 h1:h45NjjzEO3faG9Lg/cFrBh2PgegVVgzqKzuZl/wMbiI= +github.com/googleapis/gax-go/v2 v2.21.0/go.mod h1:But/NJU6TnZsrLai/xBAQLLz+Hc7fHZJt/hsCz3Fih4= github.com/jackc/pgerrcode v0.0.0-20220416144525-469b46aa5efa h1:s+4MhCQ6YrzisK6hFJUX53drDT4UsSW3DEhKn0ifuHw= github.com/jackc/pgerrcode v0.0.0-20220416144525-469b46aa5efa/go.mod h1:a/s9Lp5W7n/DD0VrVoyJ00FbP2ytTPDVOivvn2bMlds= github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= @@ -180,10 +181,10 @@ go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4= golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA= -golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo= -golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y= -golang.org/x/oauth2 v0.30.0 h1:dnDm7JmhM45NNpd8FDDeLhK6FwqbOf4MLCM9zb1BOHI= -golang.org/x/oauth2 v0.30.0/go.mod h1:B++QgG3ZKulg6sRPGD/mqlHQs5rB3Ml9erfeDY7xKlU= +golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0= +golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw= +golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= +golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -196,15 +197,20 @@ golang.org/x/term v0.41.0/go.mod h1:3pfBgksrReYfZ5lvYM0kSO0LIkAl4Yl2bXOkKP7Ec2A= golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -google.golang.org/api v0.247.0 h1:tSd/e0QrUlLsrwMKmkbQhYVa109qIintOls2Wh6bngc= -google.golang.org/api v0.247.0/go.mod h1:r1qZOPmxXffXg6xS5uhx16Fa/UFY8QU/K4bfKrnvovM= -google.golang.org/genproto v0.0.0-20250603155806-513f23925822 h1:rHWScKit0gvAPuOnu87KpaYtjK5zBMLcULh7gxkCXu4= -google.golang.org/genproto/googleapis/rpc v0.0.0-20250818200422-3122310a409c h1:qXWI/sQtv5UKboZ/zUk7h+mrf/lXORyI+n9DKDAusdg= -google.golang.org/genproto/googleapis/rpc v0.0.0-20250818200422-3122310a409c/go.mod h1:gw1tLEfykwDz2ET4a12jcXt4couGAm7IwsVaTy0Sflo= -google.golang.org/grpc v1.74.2 h1:WoosgB65DlWVC9FqI82dGsZhWFNBSLjQ84bjROOpMu4= -google.golang.org/grpc v1.74.2/go.mod h1:CtQ+BGjaAIXHs/5YS3i473GqwBBa1zGQNevxdeBEXrM= -google.golang.org/protobuf v1.36.7 h1:IgrO7UwFQGJdRNXH/sQux4R1Dj1WAKcLElzeeRaXV2A= -google.golang.org/protobuf v1.36.7/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/api v0.276.0 h1:nVArUtfLEihtW+b0DdcqRGK1xoEm2+ltAihyztq7MKY= +google.golang.org/api v0.276.0/go.mod h1:Fnag/EWUPIcJXuIkP1pjoTgS5vdxlk3eeemL7Do6bvw= +google.golang.org/genproto v0.0.0-20260319201613-d00831a3d3e7 h1:XzmzkmB14QhVhgnawEVsOn6OFsnpyxNPRY9QV01dNB0= +google.golang.org/genproto v0.0.0-20260319201613-d00831a3d3e7/go.mod h1:L43LFes82YgSonw6iTXTxXUX1OlULt4AQtkik4ULL/I= +google.golang.org/genproto/googleapis/api v0.0.0-20260319201613-d00831a3d3e7 h1:41r6JMbpzBMen0R/4TZeeAmGXSJC7DftGINUodzTkPI= +google.golang.org/genproto/googleapis/api v0.0.0-20260319201613-d00831a3d3e7/go.mod h1:EIQZ5bFCfRQDV4MhRle7+OgjNtZ6P1PiZBgAKuxXu/Y= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260401024825-9d38bb4040a9 h1:m8qni9SQFH0tJc1X0vmnpw/0t+AImlSvp30sEupozUg= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260401024825-9d38bb4040a9/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.80.0 h1:Xr6m2WmWZLETvUNvIUmeD5OAagMw3FiKmMlTdViWsHM= +google.golang.org/grpc v1.80.0/go.mod h1:ho/dLnxwi3EDJA4Zghp7k2Ec1+c2jqup0bFkw07bwF4= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= diff --git a/internal/drive/cache.go b/internal/drive/cache.go new file mode 100644 index 0000000..30c6aac --- /dev/null +++ b/internal/drive/cache.go @@ -0,0 +1,39 @@ +package drive + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/google/uuid" + goredis "github.com/redis/go-redis/v9" +) + +type FolderCache struct { + cli *goredis.Client +} + +func NewFolderCache(cli *goredis.Client) *FolderCache { + return &FolderCache{cli: cli} +} + +func (c *FolderCache) Get(ctx context.Context, userID uuid.UUID, path string) (string, bool, error) { + key := fmt.Sprintf("drive_folder_cache:%s:%s", userID, path) + val, err := c.cli.Get(ctx, key).Result() + if errors.Is(err, goredis.Nil) { + return "", false, nil + } + if err != nil { + return "", false, fmt.Errorf("cache get: %w", err) + } + return val, true, nil +} + +func (c *FolderCache) Set(ctx context.Context, userID uuid.UUID, path, folderID string, ttl time.Duration) error { + key := fmt.Sprintf("drive_folder_cache:%s:%s", userID, path) + if err := c.cli.Set(ctx, key, folderID, ttl).Err(); err != nil { + return fmt.Errorf("cache set: %w", err) + } + return nil +} diff --git a/internal/drive/cache_integration_test.go b/internal/drive/cache_integration_test.go new file mode 100644 index 0000000..48069aa --- /dev/null +++ b/internal/drive/cache_integration_test.go @@ -0,0 +1,48 @@ +//go:build integration + +package drive_test + +import ( + "context" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/smallchungus/disttaskqueue/internal/drive" + "github.com/smallchungus/disttaskqueue/internal/testutil" +) + +func TestFolderCache_SetThenGet(t *testing.T) { + cli := testutil.StartRedis(t) + cache := drive.NewFolderCache(cli) + ctx := context.Background() + uid := uuid.New() + + if err := cache.Set(ctx, uid, "2026/04/17", "folder-id-abc", time.Minute); err != nil { + t.Fatal(err) + } + + got, ok, err := cache.Get(ctx, uid, "2026/04/17") + if err != nil { + t.Fatal(err) + } + if !ok || got != "folder-id-abc" { + t.Fatalf("got %q, ok=%v", got, ok) + } +} + +func TestFolderCache_GetReturnsFalseOnMiss(t *testing.T) { + cli := testutil.StartRedis(t) + cache := drive.NewFolderCache(cli) + ctx := context.Background() + uid := uuid.New() + + got, ok, err := cache.Get(ctx, uid, "nope") + if err != nil { + t.Fatal(err) + } + if ok || got != "" { + t.Fatalf("got %q, ok=%v, want miss", got, ok) + } +} diff --git a/internal/drive/client.go b/internal/drive/client.go new file mode 100644 index 0000000..049c93b --- /dev/null +++ b/internal/drive/client.go @@ -0,0 +1,97 @@ +package drive + +import ( + "bytes" + "context" + "fmt" + + "github.com/google/uuid" + "golang.org/x/oauth2" + driveapi "google.golang.org/api/drive/v3" + "google.golang.org/api/googleapi" + "google.golang.org/api/option" + + "github.com/smallchungus/disttaskqueue/internal/oauth" + "github.com/smallchungus/disttaskqueue/internal/store" +) + +const folderMime = "application/vnd.google-apps.folder" + +type Config struct { + Store *store.Store + UserID uuid.UUID + EncryptionKey []byte + OAuth2 *oauth2.Config + Endpoint string +} + +type Client struct { + svc *driveapi.Service +} + +func New(ctx context.Context, cfg Config) (*Client, error) { + tok, err := oauth.LoadToken(ctx, cfg.Store, cfg.UserID, cfg.EncryptionKey, "google") + if err != nil { + return nil, err + } + + base := cfg.OAuth2.TokenSource(ctx, tok) + saving := oauth.NewSavingSource(base, func(t *oauth2.Token) error { + return oauth.SaveToken(ctx, cfg.Store, cfg.UserID, cfg.EncryptionKey, "google", t) + }, tok) + httpClient := oauth2.NewClient(ctx, saving) + + opts := []option.ClientOption{option.WithHTTPClient(httpClient)} + if cfg.Endpoint != "" { + opts = append(opts, option.WithEndpoint(cfg.Endpoint)) + } + svc, err := driveapi.NewService(ctx, opts...) + if err != nil { + return nil, fmt.Errorf("drive svc: %w", err) + } + return &Client{svc: svc}, nil +} + +func (c *Client) EnsureFolder(ctx context.Context, parentID, name string) (string, error) { + q := fmt.Sprintf("name = '%s' and '%s' in parents and mimeType = '%s' and trashed = false", escape(name), parentID, folderMime) + resp, err := c.svc.Files.List().Q(q).Fields("files(id,name)").Context(ctx).Do() + if err != nil { + return "", fmt.Errorf("list folders: %w", err) + } + if len(resp.Files) > 0 { + return resp.Files[0].Id, nil + } + + created, err := c.svc.Files.Create(&driveapi.File{ + Name: name, + Parents: []string{parentID}, + MimeType: folderMime, + }).Fields("id").Context(ctx).Do() + if err != nil { + return "", fmt.Errorf("create folder: %w", err) + } + return created.Id, nil +} + +func (c *Client) Upload(ctx context.Context, parentID, name, contentType string, content []byte) (string, error) { + created, err := c.svc.Files.Create(&driveapi.File{ + Name: name, + Parents: []string{parentID}, + }).Media(bytes.NewReader(content), googleapi.ContentType(contentType)).Fields("id").Context(ctx).Do() + if err != nil { + return "", fmt.Errorf("upload: %w", err) + } + return created.Id, nil +} + +func escape(s string) string { + out := make([]byte, 0, len(s)) + for i := 0; i < len(s); i++ { + if s[i] == '\'' { + out = append(out, '\\', '\'') + continue + } + out = append(out, s[i]) + } + return string(out) +} diff --git a/internal/drive/client_integration_test.go b/internal/drive/client_integration_test.go new file mode 100644 index 0000000..ecda651 --- /dev/null +++ b/internal/drive/client_integration_test.go @@ -0,0 +1,143 @@ +//go:build integration + +package drive_test + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "golang.org/x/oauth2" + + "github.com/smallchungus/disttaskqueue/internal/drive" + "github.com/smallchungus/disttaskqueue/internal/oauth" + "github.com/smallchungus/disttaskqueue/internal/store" + "github.com/smallchungus/disttaskqueue/internal/testutil" +) + +type driveMock struct { + listResp string + createResp string + uploadResp string + listHits int + createHits int + uploadHits int + lastUpload []byte +} + +func (m *driveMock) server(t *testing.T) *httptest.Server { + t.Helper() + mux := http.NewServeMux() + mux.HandleFunc("/files", func(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + m.listHits++ + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(m.listResp)) + case http.MethodPost: + m.createHits++ + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(m.createResp)) + } + }) + mux.HandleFunc("/upload/drive/v3/files", func(w http.ResponseWriter, r *http.Request) { + m.uploadHits++ + body, _ := io.ReadAll(r.Body) + m.lastUpload = body + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(m.uploadResp)) + }) + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + return srv +} + +func newKey32() []byte { + k := make([]byte, 32) + for i := range k { + k[i] = byte(i) + } + return k +} + +func setupClient(t *testing.T, m *driveMock) *drive.Client { + t.Helper() + pool := testutil.StartPostgres(t) + if err := store.Migrate(context.Background(), pool.Config().ConnString()); err != nil { + t.Fatal(err) + } + s := store.New(pool) + u, _ := s.CreateUser(context.Background(), fmt.Sprintf("drive+%d@example.com", time.Now().UnixNano())) + + key := newKey32() + tok := &oauth2.Token{AccessToken: "x", RefreshToken: "y", Expiry: time.Now().Add(time.Hour)} + if err := oauth.SaveToken(context.Background(), s, u.ID, key, "google", tok); err != nil { + t.Fatal(err) + } + + srv := m.server(t) + cfg := &oauth2.Config{ClientID: "x", ClientSecret: "y"} + c, err := drive.New(context.Background(), drive.Config{ + Store: s, + UserID: u.ID, + EncryptionKey: key, + OAuth2: cfg, + Endpoint: srv.URL, + }) + if err != nil { + t.Fatalf("new client: %v", err) + } + return c +} + +func TestEnsureFolder_ReturnsExistingFolderID(t *testing.T) { + listJSON, _ := json.Marshal(map[string]any{ + "files": []map[string]any{{"id": "existing-folder-id", "name": "myfolder"}}, + }) + c := setupClient(t, &driveMock{listResp: string(listJSON)}) + + id, err := c.EnsureFolder(context.Background(), "parent-id", "myfolder") + if err != nil { + t.Fatalf("ensure: %v", err) + } + if id != "existing-folder-id" { + t.Fatalf("got %q, want existing-folder-id", id) + } +} + +func TestEnsureFolder_CreatesIfMissing(t *testing.T) { + listJSON, _ := json.Marshal(map[string]any{"files": []map[string]any{}}) + createJSON, _ := json.Marshal(map[string]any{"id": "new-folder-id", "name": "newfolder"}) + c := setupClient(t, &driveMock{listResp: string(listJSON), createResp: string(createJSON)}) + + id, err := c.EnsureFolder(context.Background(), "parent-id", "newfolder") + if err != nil { + t.Fatalf("ensure: %v", err) + } + if id != "new-folder-id" { + t.Fatalf("got %q, want new-folder-id", id) + } +} + +func TestUpload_PostsContent(t *testing.T) { + uploadJSON, _ := json.Marshal(map[string]any{"id": "uploaded-file-id"}) + m := &driveMock{uploadResp: string(uploadJSON)} + c := setupClient(t, m) + + id, err := c.Upload(context.Background(), "parent-id", "report.pdf", "application/pdf", []byte("PDF-DATA")) + if err != nil { + t.Fatalf("upload: %v", err) + } + if id != "uploaded-file-id" { + t.Fatalf("got %q, want uploaded-file-id", id) + } + if !strings.Contains(string(m.lastUpload), "PDF-DATA") { + t.Fatalf("upload body did not contain PDF-DATA") + } +} diff --git a/internal/gmail/client.go b/internal/gmail/client.go index 3a50b96..e575974 100644 --- a/internal/gmail/client.go +++ b/internal/gmail/client.go @@ -11,6 +11,7 @@ import ( gmailapi "google.golang.org/api/gmail/v1" "google.golang.org/api/option" + "github.com/smallchungus/disttaskqueue/internal/oauth" "github.com/smallchungus/disttaskqueue/internal/store" ) @@ -30,14 +31,14 @@ type Client struct { } func New(ctx context.Context, cfg Config) (*Client, error) { - tok, err := LoadToken(ctx, cfg.Store, cfg.UserID, cfg.EncryptionKey) + tok, err := oauth.LoadToken(ctx, cfg.Store, cfg.UserID, cfg.EncryptionKey, "google") if err != nil { return nil, err } base := cfg.OAuth2.TokenSource(ctx, tok) - saving := newSavingSource(base, func(t *oauth2.Token) error { - return SaveToken(ctx, cfg.Store, cfg.UserID, cfg.EncryptionKey, t) + saving := oauth.NewSavingSource(base, func(t *oauth2.Token) error { + return oauth.SaveToken(ctx, cfg.Store, cfg.UserID, cfg.EncryptionKey, "google", t) }, tok) httpClient := oauth2.NewClient(ctx, saving) diff --git a/internal/gmail/client_integration_test.go b/internal/gmail/client_integration_test.go index f129dec..6a9321d 100644 --- a/internal/gmail/client_integration_test.go +++ b/internal/gmail/client_integration_test.go @@ -15,6 +15,7 @@ import ( "golang.org/x/oauth2" "github.com/smallchungus/disttaskqueue/internal/gmail" + "github.com/smallchungus/disttaskqueue/internal/oauth" "github.com/smallchungus/disttaskqueue/internal/store" "github.com/smallchungus/disttaskqueue/internal/testutil" ) @@ -71,7 +72,7 @@ func setupClient(t *testing.T, m *gmailMock) (*gmail.Client, *store.Store) { RefreshToken: "fake-refresh", Expiry: time.Now().Add(1 * time.Hour), } - if err := gmail.SaveToken(context.Background(), s, u.ID, key, tok); err != nil { + if err := oauth.SaveToken(context.Background(), s, u.ID, key, "google", tok); err != nil { t.Fatal(err) } diff --git a/internal/gmail/token.go b/internal/oauth/token.go similarity index 67% rename from internal/gmail/token.go rename to internal/oauth/token.go index 73fd508..9cea90f 100644 --- a/internal/gmail/token.go +++ b/internal/oauth/token.go @@ -1,4 +1,4 @@ -package gmail +package oauth import ( "context" @@ -8,50 +8,49 @@ import ( "github.com/google/uuid" "golang.org/x/oauth2" - "github.com/smallchungus/disttaskqueue/internal/oauth" "github.com/smallchungus/disttaskqueue/internal/store" ) -func LoadToken(ctx context.Context, s *store.Store, userID uuid.UUID, key []byte) (*oauth2.Token, error) { - rec, err := s.GetOAuthToken(ctx, userID, "google") +func LoadToken(ctx context.Context, s *store.Store, userID uuid.UUID, key []byte, provider string) (*oauth2.Token, error) { + rec, err := s.GetOAuthToken(ctx, userID, provider) if err != nil { return nil, fmt.Errorf("load token: %w", err) } - return decryptToken(rec.AccessCT, rec.RefreshCT, rec.ExpiresAt, key) + return DecryptToken(rec.AccessCT, rec.RefreshCT, rec.ExpiresAt, key) } -func SaveToken(ctx context.Context, s *store.Store, userID uuid.UUID, key []byte, tok *oauth2.Token) error { - accessCT, refreshCT, err := encryptToken(tok, key) +func SaveToken(ctx context.Context, s *store.Store, userID uuid.UUID, key []byte, provider string, tok *oauth2.Token) error { + accessCT, refreshCT, err := EncryptToken(tok, key) if err != nil { return fmt.Errorf("save token: %w", err) } return s.SaveOAuthToken(ctx, store.OAuthToken{ UserID: userID, - Provider: "google", + Provider: provider, AccessCT: accessCT, RefreshCT: refreshCT, ExpiresAt: tok.Expiry, }) } -func encryptToken(tok *oauth2.Token, key []byte) (access, refresh []byte, err error) { - access, err = oauth.Encrypt([]byte(tok.AccessToken), key) +func EncryptToken(tok *oauth2.Token, key []byte) (access, refresh []byte, err error) { + access, err = Encrypt([]byte(tok.AccessToken), key) if err != nil { return nil, nil, err } - refresh, err = oauth.Encrypt([]byte(tok.RefreshToken), key) + refresh, err = Encrypt([]byte(tok.RefreshToken), key) if err != nil { return nil, nil, err } return access, refresh, nil } -func decryptToken(access, refresh []byte, expiry time.Time, key []byte) (*oauth2.Token, error) { - a, err := oauth.Decrypt(access, key) +func DecryptToken(access, refresh []byte, expiry time.Time, key []byte) (*oauth2.Token, error) { + a, err := Decrypt(access, key) if err != nil { return nil, fmt.Errorf("decrypt access: %w", err) } - r, err := oauth.Decrypt(refresh, key) + r, err := Decrypt(refresh, key) if err != nil { return nil, fmt.Errorf("decrypt refresh: %w", err) } @@ -69,7 +68,7 @@ type savingSource struct { last *oauth2.Token } -func newSavingSource(base oauth2.TokenSource, save func(*oauth2.Token) error, seed *oauth2.Token) oauth2.TokenSource { +func NewSavingSource(base oauth2.TokenSource, save func(*oauth2.Token) error, seed *oauth2.Token) oauth2.TokenSource { return &savingSource{base: base, save: save, last: seed} } diff --git a/internal/gmail/token_test.go b/internal/oauth/token_test.go similarity index 80% rename from internal/gmail/token_test.go rename to internal/oauth/token_test.go index f00a004..ee4abbc 100644 --- a/internal/gmail/token_test.go +++ b/internal/oauth/token_test.go @@ -1,19 +1,13 @@ -package gmail +package oauth import ( - "context" "errors" "testing" "time" - "github.com/google/uuid" "golang.org/x/oauth2" - - "github.com/smallchungus/disttaskqueue/internal/store" ) -type fakeStore struct{} - type sequenceSource struct { tokens []*oauth2.Token idx int @@ -38,7 +32,7 @@ func TestSavingSource_SavesWhenTokenChanges(t *testing.T) { base := &sequenceSource{tokens: []*oauth2.Token{tok1, tok2}} var saved []*oauth2.Token - src := newSavingSource(base, func(t *oauth2.Token) error { + src := NewSavingSource(base, func(t *oauth2.Token) error { saved = append(saved, t) return nil }, tok1) @@ -68,7 +62,7 @@ func TestSavingSource_SavesWhenTokenChanges(t *testing.T) { func TestSavingSource_PropagatesBaseError(t *testing.T) { base := &sequenceSource{err: errors.New("refresh failed")} - src := newSavingSource(base, func(*oauth2.Token) error { return nil }, nil) + src := NewSavingSource(base, func(*oauth2.Token) error { return nil }, nil) if _, err := src.Token(); err == nil { t.Fatal("expected error") } @@ -85,11 +79,11 @@ func TestEncryptToken_DecryptToken_RoundTrip(t *testing.T) { Expiry: time.Now().Add(time.Hour).UTC().Truncate(time.Second), } - accessCT, refreshCT, err := encryptToken(in, key) + accessCT, refreshCT, err := EncryptToken(in, key) if err != nil { t.Fatalf("encrypt: %v", err) } - out, err := decryptToken(accessCT, refreshCT, in.Expiry, key) + out, err := DecryptToken(accessCT, refreshCT, in.Expiry, key) if err != nil { t.Fatalf("decrypt: %v", err) } @@ -97,8 +91,3 @@ func TestEncryptToken_DecryptToken_RoundTrip(t *testing.T) { t.Fatalf("mismatch: %+v vs %+v", out, in) } } - -var _ = fakeStore{} -var _ = uuid.UUID{} -var _ = context.Background() -var _ = store.OAuthToken{}