diff --git a/client.go b/client.go index bd8f514..7da1e6c 100644 --- a/client.go +++ b/client.go @@ -145,21 +145,11 @@ func WithRetryPolicy(maxRetries int, backoff Backoff) ClientOption { } } -// request makes an HTTP request to the Replicate API. -func (r *Client) request(ctx context.Context, method, path string, body interface{}, out interface{}) error { - bodyBuffer := &bytes.Buffer{} - if body != nil { - bodyBytes, err := json.Marshal(body) - if err != nil { - return fmt.Errorf("failed to marshal request body: %w", err) - } - bodyBuffer = bytes.NewBuffer(bodyBytes) - } - +func (r *Client) newRequest(ctx context.Context, method, path string, body io.Reader) (*http.Request, error) { url := constructURL(r.options.baseURL, path) - request, err := http.NewRequestWithContext(ctx, method, url, bodyBuffer) + request, err := http.NewRequestWithContext(ctx, method, url, body) if err != nil { - return fmt.Errorf("failed to create request: %w", err) + return nil, fmt.Errorf("failed to create request: %w", err) } request.Header.Set("Content-Type", "application/json") @@ -168,6 +158,10 @@ func (r *Client) request(ctx context.Context, method, path string, body interfac request.Header.Set("User-Agent", *r.options.userAgent) } + return request, nil +} + +func (r *Client) do(ctx context.Context, request *http.Request, out interface{}) error { maxRetries := r.options.retryPolicy.maxRetries backoff := r.options.retryPolicy.backoff @@ -187,7 +181,7 @@ func (r *Client) request(ctx context.Context, method, path string, body interfac if response.StatusCode < 200 || response.StatusCode >= 400 { apiError = unmarshalAPIError(response, responseBytes) - if !r.shouldRetry(response, method) { + if !r.shouldRetry(response, request.Method) { return apiError } @@ -229,6 +223,25 @@ func (r *Client) request(ctx context.Context, method, path string, body interfac return fmt.Errorf("request failed") } +// fetch makes an HTTP request to Replicate's API. +func (r *Client) fetch(ctx context.Context, method, path string, body interface{}, out interface{}) error { + bodyBuffer := &bytes.Buffer{} + if body != nil { + bodyBytes, err := json.Marshal(body) + if err != nil { + return fmt.Errorf("failed to marshal request body: %w", err) + } + bodyBuffer = bytes.NewBuffer(bodyBytes) + } + + request, err := r.newRequest(ctx, method, path, bodyBuffer) + if err != nil { + return err + } + + return r.do(ctx, request, out) +} + // shouldRetry returns true if the request should be retried. // // - GET requests should be retried if the response status code is 429 or 5xx. diff --git a/client_test.go b/client_test.go index e212750..eb5ffb9 100644 --- a/client_test.go +++ b/client_test.go @@ -1,12 +1,20 @@ package replicate_test import ( + "bytes" "context" + "crypto/md5" // nolint:gosec + "crypto/sha256" + "encoding/hex" "encoding/json" "fmt" "io" + "mime" + "mime/multipart" "net/http" "net/http/httptest" + "os" + "path/filepath" "testing" "time" @@ -1139,7 +1147,6 @@ func TestAutomaticallyRetryGetRequests(t *testing.T) { i := 0 mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - status := statuses[i] i++ @@ -1318,3 +1325,241 @@ func TestStream(t *testing.T) { } } } + +func TestCreateFile(t *testing.T) { + fileID := "file-id" + options := &replicate.CreateFileOptions{ + Filename: "hello.txt", + ContentType: "text/plain", + Metadata: map[string]string{"foo": "bar"}, + } + + mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/files", r.URL.Path) + assert.Equal(t, http.MethodPost, r.Method) + + _, params, err := mime.ParseMediaType(r.Header.Get("Content-Type")) + if err != nil { + t.Fatal(err) + } + + mr := multipart.NewReader(r.Body, params["boundary"]) + defer r.Body.Close() + + part, err := mr.NextPart() + if err != nil { + t.Fatal(err) + } + + assert.Equal(t, "form-data; name=\"content\"; filename=\"hello.txt\"", part.Header.Get("Content-Disposition")) + assert.Equal(t, "text/plain", part.Header.Get("Content-Type")) + + content, err := io.ReadAll(part) + if err != nil { + t.Fatal(err) + } + + etag := fmt.Sprintf("%x", md5.Sum(content)) // nolint:gosec + checksum := sha256.Sum256(content) + file := &replicate.File{ + ID: fileID, + Name: "hello.txt", + ContentType: "text/plain", + Size: len(content), + Etag: etag, + Checksums: map[string]string{"sha256": hex.EncodeToString(checksum[:])}, + Metadata: map[string]string{"foo": "bar"}, + CreatedAt: "2022-04-26T22:13:06.224088Z", + URLs: map[string]string{"get": "https://api.replicate.com/v1/files/" + fileID}, + } + + responseBytes, err := json.Marshal(file) + if err != nil { + t.Fatal(err) + } + + w.WriteHeader(http.StatusCreated) + w.Write(responseBytes) + })) + defer mockServer.Close() + + client, err := replicate.NewClient( + replicate.WithToken("test-token"), + replicate.WithBaseURL(mockServer.URL), + ) + require.NotNil(t, client) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + t.Run("CreateFileFromBytes", func(t *testing.T) { + content := []byte("Hello, world!") + file, err := client.CreateFileFromBytes(ctx, content, options) + if err != nil { + t.Fatal(err) + } + assertCreatedFile(t, fileID, file) + }) + + t.Run("CreateFileFromBuffer", func(t *testing.T) { + buf := bytes.NewBufferString("Hello, world!") + file, err := client.CreateFileFromBuffer(ctx, buf, options) + if err != nil { + t.Fatal(err) + } + assertCreatedFile(t, fileID, file) + }) + + t.Run("CreateFileFromPath", func(t *testing.T) { + content := []byte("Hello, world!") + tmpFilePath := filepath.Join(t.TempDir(), "hello.txt") + if err := os.WriteFile(tmpFilePath, content, 0o644); err != nil { + t.Fatal(err) + } + file, err := client.CreateFileFromPath(ctx, tmpFilePath, options) + if err != nil { + t.Fatal(err) + } + assertCreatedFile(t, fileID, file) + }) +} + +func assertCreatedFile(t *testing.T, fileID string, file *replicate.File) { + assert.Equal(t, fileID, file.ID) + assert.Equal(t, "hello.txt", file.Name) + assert.Equal(t, "text/plain", file.ContentType) + assert.Equal(t, 13, file.Size) + assert.Equal(t, "6cd3556deb0da54bca060b4c39479839", file.Etag) + assert.Equal(t, "315f5bdb76d078c43b8ac0064e4a0164612b1fce77c869345bfc94c75894edd3", file.Checksums["sha256"]) + assert.Equal(t, map[string]string{"foo": "bar"}, file.Metadata) + assert.Equal(t, "2022-04-26T22:13:06.224088Z", file.CreatedAt) + assert.Equal(t, "https://api.replicate.com/v1/files/"+fileID, file.URLs["get"]) +} + +func TestListFiles(t *testing.T) { + fileID := "file-id" + mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/files", r.URL.Path) + assert.Equal(t, http.MethodGet, r.Method) + + response := replicate.Page[replicate.File]{ + Results: []replicate.File{ + { + ID: fileID, + Name: "hello.txt", + ContentType: "text/plain", + Size: 13, + CreatedAt: "2022-04-26T22:13:06.224088Z", + URLs: map[string]string{"get": "https://api.replicate.com/v1/files/" + fileID}, + }, + }, + } + + responseBytes, err := json.Marshal(response) + if err != nil { + t.Fatal(err) + } + + w.WriteHeader(http.StatusOK) + w.Write(responseBytes) + })) + defer mockServer.Close() + + client, err := replicate.NewClient( + replicate.WithToken("test-token"), + replicate.WithBaseURL(mockServer.URL), + ) + require.NotNil(t, client) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + files, err := client.ListFiles(ctx) + if err != nil { + t.Fatal(err) + } + + assert.Equal(t, 1, len(files.Results)) + assert.Nil(t, files.Previous) + assert.Nil(t, files.Next) + + file := files.Results[0] + assert.Equal(t, fileID, file.ID) + assert.Equal(t, "hello.txt", file.Name) + assert.Equal(t, "text/plain", file.ContentType) + assert.Equal(t, 13, file.Size) + assert.Equal(t, "2022-04-26T22:13:06.224088Z", file.CreatedAt) + assert.Equal(t, "https://api.replicate.com/v1/files/"+fileID, file.URLs["get"]) +} +func TestGetFile(t *testing.T) { + fileID := "file-id" + mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/files/"+fileID, r.URL.Path) + assert.Equal(t, http.MethodGet, r.Method) + + file := &replicate.File{ + ID: fileID, + Name: "hello.txt", + ContentType: "text/plain", + Size: 13, + CreatedAt: "2022-04-26T22:13:06.224088Z", + URLs: map[string]string{"get": "https://api.replicate.com/v1/files/" + fileID}, + } + + responseBytes, err := json.Marshal(file) + if err != nil { + t.Fatal(err) + } + + w.WriteHeader(http.StatusOK) + w.Write(responseBytes) + })) + defer mockServer.Close() + + client, err := replicate.NewClient( + replicate.WithToken("test-token"), + replicate.WithBaseURL(mockServer.URL), + ) + require.NotNil(t, client) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + file, err := client.GetFile(ctx, fileID) + if err != nil { + t.Fatal(err) + } + + assert.Equal(t, fileID, file.ID) + assert.Equal(t, "hello.txt", file.Name) + assert.Equal(t, "text/plain", file.ContentType) + assert.Equal(t, 13, file.Size) + assert.Equal(t, "2022-04-26T22:13:06.224088Z", file.CreatedAt) + assert.Equal(t, "https://api.replicate.com/v1/files/"+fileID, file.URLs["get"]) +} + +func TestDeleteFile(t *testing.T) { + fileID := "file-id" + mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/files/"+fileID, r.URL.Path) + assert.Equal(t, http.MethodDelete, r.Method) + w.WriteHeader(http.StatusOK) + })) + defer mockServer.Close() + + client, err := replicate.NewClient( + replicate.WithToken("test-token"), + replicate.WithBaseURL(mockServer.URL), + ) + require.NotNil(t, client) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + err = client.DeleteFile(ctx, fileID) + assert.NoError(t, err) +} diff --git a/collection.go b/collection.go index e570eec..399bce6 100644 --- a/collection.go +++ b/collection.go @@ -34,7 +34,7 @@ func (c *Collection) UnmarshalJSON(data []byte) error { // ListCollections returns a list of all collections. func (r *Client) ListCollections(ctx context.Context) (*Page[Collection], error) { response := &Page[Collection]{} - err := r.request(ctx, "GET", "/collections", nil, response) + err := r.fetch(ctx, "GET", "/collections", nil, response) if err != nil { return nil, fmt.Errorf("failed to list collections: %w", err) } @@ -44,7 +44,7 @@ func (r *Client) ListCollections(ctx context.Context) (*Page[Collection], error) // GetCollection returns a collection by slug. func (r *Client) GetCollection(ctx context.Context, slug string) (*Collection, error) { collection := &Collection{} - err := r.request(ctx, "GET", fmt.Sprintf("/collections/%s", slug), nil, collection) + err := r.fetch(ctx, "GET", fmt.Sprintf("/collections/%s", slug), nil, collection) if err != nil { return nil, fmt.Errorf("failed to get collection: %w", err) } diff --git a/deployment.go b/deployment.go index 07e9614..5f7f522 100644 --- a/deployment.go +++ b/deployment.go @@ -22,7 +22,7 @@ func (r *Client) CreatePredictionWithDeployment(ctx context.Context, deployment_ prediction := &Prediction{} path := fmt.Sprintf("/deployments/%s/%s/predictions", deployment_owner, deployment_name) - err := r.request(ctx, "POST", path, data, prediction) + err := r.fetch(ctx, "POST", path, data, prediction) if err != nil { return nil, fmt.Errorf("failed to create prediction: %w", err) } diff --git a/files.go b/files.go new file mode 100644 index 0000000..8df8bc5 --- /dev/null +++ b/files.go @@ -0,0 +1,163 @@ +package replicate + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "mime" + "mime/multipart" + "net/http" + "net/textproto" + "os" + "path/filepath" +) + +type File struct { + ID string `json:"id"` + Name string `json:"name"` + ContentType string `json:"content_type"` + Size int `json:"size"` + Etag string `json:"etag"` + Checksums map[string]string `json:"checksums"` + Metadata map[string]string `json:"metadata"` + CreatedAt string `json:"created_at"` + ExpiresAt string `json:"expires_at"` + URLs map[string]string `json:"urls"` +} + +type CreateFileOptions struct { + Filename string `json:"filename"` + ContentType string `json:"content_type"` + Metadata map[string]string `json:"metadata"` +} + +// CreateFileFromPath creates a new file from a file path. +func (r *Client) CreateFileFromPath(ctx context.Context, filePath string, options *CreateFileOptions) (*File, error) { + f, err := os.Open(filePath) + if err != nil { + return nil, fmt.Errorf("failed to open file: %w", err) + } + defer f.Close() + + if options.Filename == "" { + _, options.Filename = filepath.Split(filePath) + } + + if options.ContentType == "" { + if options.Filename != "" { + ext := filepath.Ext(options.Filename) + options.ContentType = mime.TypeByExtension(ext) + } + } + + return r.createFile(ctx, f, options) +} + +// CreateFileFromBytes creates a new file from bytes. +func (r *Client) CreateFileFromBytes(ctx context.Context, data []byte, options *CreateFileOptions) (*File, error) { + buf := bytes.NewBuffer(data) + + if options.ContentType == "" { + options.ContentType = http.DetectContentType(data) + } + + return r.createFile(ctx, buf, options) +} + +// CreateFileFromBuffer creates a new file from a buffer. +func (r *Client) CreateFileFromBuffer(ctx context.Context, buf *bytes.Buffer, options *CreateFileOptions) (*File, error) { + return r.createFile(ctx, buf, options) +} + +// CreateFile creates a new file. +func (r *Client) createFile(ctx context.Context, reader io.Reader, options *CreateFileOptions) (*File, error) { + body := &bytes.Buffer{} + writer := multipart.NewWriter(body) + + filename := options.Filename + if filename == "" { + filename = "file" + } + + contentType := options.ContentType + if contentType == "" { + contentType = "application/octet-stream" + } + + h := make(textproto.MIMEHeader) + h.Set("Content-Disposition", fmt.Sprintf(`form-data; name="content"; filename="%s"`, filename)) + h.Set("Content-Type", contentType) + + content, err := writer.CreatePart(h) + if err != nil { + return nil, fmt.Errorf("failed to create form file: %w", err) + } + _, err = io.Copy(content, reader) + if err != nil { + return nil, fmt.Errorf("failed to write file to form: %w", err) + } + + if options.Metadata != nil { + metadata, err := json.Marshal(options.Metadata) + if err != nil { + return nil, fmt.Errorf("failed to marshal metadata: %w", err) + } + err = writer.WriteField("metadata", string(metadata)) + if err != nil { + return nil, fmt.Errorf("failed to write metadata to form: %w", err) + } + } + + err = writer.Close() + if err != nil { + return nil, fmt.Errorf("failed to close writer: %w", err) + } + + req, err := r.newRequest(ctx, "POST", "/files", body) + if err != nil { + return nil, fmt.Errorf("failed to create request: %w", err) + } + req.Header.Set("Content-Type", writer.FormDataContentType()) + + file := &File{} + err = r.do(ctx, req, file) + if err != nil { + return nil, fmt.Errorf("failed to create file: %w", err) + } + + return file, nil +} + +// ListFiles lists your files. +func (r *Client) ListFiles(ctx context.Context) (*Page[File], error) { + response := &Page[File]{} + err := r.fetch(ctx, "GET", "/files", nil, response) + if err != nil { + return nil, fmt.Errorf("failed to list files: %w", err) + } + + return response, nil +} + +// GetFile retrieves information about a file. +func (r *Client) GetFile(ctx context.Context, fileID string) (*File, error) { + file := &File{} + err := r.fetch(ctx, "GET", fmt.Sprintf("/files/%s", fileID), nil, file) + if err != nil { + return nil, fmt.Errorf("failed to get file: %w", err) + } + + return file, nil +} + +// DeleteFile deletes a file. +func (r *Client) DeleteFile(ctx context.Context, fileID string) error { + err := r.fetch(ctx, "DELETE", fmt.Sprintf("/files/%s", fileID), nil, nil) + if err != nil { + return fmt.Errorf("failed to delete file: %w", err) + } + + return nil +} diff --git a/hardware.go b/hardware.go index 54856f7..3cc5cda 100644 --- a/hardware.go +++ b/hardware.go @@ -32,7 +32,7 @@ func (h *Hardware) UnmarshalJSON(data []byte) error { // ListHardware returns a list of available hardware. func (r *Client) ListHardware(ctx context.Context) (*[]Hardware, error) { response := &[]Hardware{} - err := r.request(ctx, "GET", "/hardware", nil, response) + err := r.fetch(ctx, "GET", "/hardware", nil, response) if err != nil { return nil, fmt.Errorf("failed to list collections: %w", err) } diff --git a/model.go b/model.go index bdff5bd..42b6bb6 100644 --- a/model.go +++ b/model.go @@ -77,7 +77,7 @@ func (m *ModelVersion) UnmarshalJSON(data []byte) error { // ListModels lists public models. func (r *Client) ListModels(ctx context.Context) (*Page[Model], error) { response := &Page[Model]{} - err := r.request(ctx, "GET", "/models", nil, response) + err := r.fetch(ctx, "GET", "/models", nil, response) if err != nil { return nil, fmt.Errorf("failed to list models: %w", err) } @@ -87,7 +87,7 @@ func (r *Client) ListModels(ctx context.Context) (*Page[Model], error) { // GetModel retrieves information about a model. func (r *Client) GetModel(ctx context.Context, modelOwner string, modelName string) (*Model, error) { model := &Model{} - err := r.request(ctx, "GET", fmt.Sprintf("/models/%s/%s", modelOwner, modelName), nil, model) + err := r.fetch(ctx, "GET", fmt.Sprintf("/models/%s/%s", modelOwner, modelName), nil, model) if err != nil { return nil, fmt.Errorf("failed to get model: %w", err) } @@ -108,7 +108,7 @@ func (r *Client) CreateModel(ctx context.Context, modelOwner string, modelName s CreateModelOptions: options, } - err := r.request(ctx, "POST", "/models", body, model) + err := r.fetch(ctx, "POST", "/models", body, model) if err != nil { return nil, fmt.Errorf("failed to create model: %w", err) } @@ -118,7 +118,7 @@ func (r *Client) CreateModel(ctx context.Context, modelOwner string, modelName s // ListModelVersions lists the versions of a model. func (r *Client) ListModelVersions(ctx context.Context, modelOwner string, modelName string) (*Page[ModelVersion], error) { response := &Page[ModelVersion]{} - err := r.request(ctx, "GET", fmt.Sprintf("/models/%s/%s/versions", modelOwner, modelName), nil, response) + err := r.fetch(ctx, "GET", fmt.Sprintf("/models/%s/%s/versions", modelOwner, modelName), nil, response) if err != nil { return nil, fmt.Errorf("failed to list model versions: %w", err) } @@ -128,7 +128,7 @@ func (r *Client) ListModelVersions(ctx context.Context, modelOwner string, model // GetModelVersion retrieves a specific version of a model. func (r *Client) GetModelVersion(ctx context.Context, modelOwner string, modelName string, versionID string) (*ModelVersion, error) { version := &ModelVersion{} - err := r.request(ctx, "GET", fmt.Sprintf("/models/%s/%s/versions/%s", modelOwner, modelName, versionID), nil, version) + err := r.fetch(ctx, "GET", fmt.Sprintf("/models/%s/%s/versions/%s", modelOwner, modelName, versionID), nil, version) if err != nil { return nil, fmt.Errorf("failed to get model version: %w", err) } @@ -151,7 +151,7 @@ func (r *Client) CreatePredictionWithModel(ctx context.Context, modelOwner strin } prediction := &Prediction{} - err := r.request(ctx, "POST", fmt.Sprintf("/models/%s/%s/predictions", modelOwner, modelName), data, prediction) + err := r.fetch(ctx, "POST", fmt.Sprintf("/models/%s/%s/predictions", modelOwner, modelName), data, prediction) if err != nil { return nil, err } diff --git a/paginate.go b/paginate.go index a2daefd..91a9085 100644 --- a/paginate.go +++ b/paginate.go @@ -44,7 +44,7 @@ func Paginate[T any](ctx context.Context, client *Client, initialPage *Page[T]) for nextURL != nil { page := &Page[T]{} - err := client.request(ctx, "GET", *nextURL, nil, page) + err := client.fetch(ctx, "GET", *nextURL, nil, page) if err != nil { errChan <- err return diff --git a/prediction.go b/prediction.go index dcc63b3..1310e9d 100644 --- a/prediction.go +++ b/prediction.go @@ -114,7 +114,7 @@ func (r *Client) CreatePrediction(ctx context.Context, version string, input Pre } prediction := &Prediction{} - err := r.request(ctx, "POST", "/predictions", data, prediction) + err := r.fetch(ctx, "POST", "/predictions", data, prediction) if err != nil { return nil, fmt.Errorf("failed to create prediction: %w", err) } @@ -125,7 +125,7 @@ func (r *Client) CreatePrediction(ctx context.Context, version string, input Pre // ListPredictions returns a paginated list of predictions. func (r *Client) ListPredictions(ctx context.Context) (*Page[Prediction], error) { response := &Page[Prediction]{} - err := r.request(ctx, "GET", "/predictions", nil, response) + err := r.fetch(ctx, "GET", "/predictions", nil, response) if err != nil { return nil, fmt.Errorf("failed to list predictions: %w", err) } @@ -135,7 +135,7 @@ func (r *Client) ListPredictions(ctx context.Context) (*Page[Prediction], error) // GetPrediction retrieves a prediction from the Replicate API by its ID. func (r *Client) GetPrediction(ctx context.Context, id string) (*Prediction, error) { prediction := &Prediction{} - err := r.request(ctx, "GET", fmt.Sprintf("/predictions/%s", id), nil, prediction) + err := r.fetch(ctx, "GET", fmt.Sprintf("/predictions/%s", id), nil, prediction) if err != nil { return nil, fmt.Errorf("failed to get prediction: %w", err) } diff --git a/training.go b/training.go index 08b04a5..5a4c011 100644 --- a/training.go +++ b/training.go @@ -23,7 +23,7 @@ func (r *Client) CreateTraining(ctx context.Context, model_owner string, model_n training := &Training{} path := fmt.Sprintf("/models/%s/%s/versions/%s/trainings", model_owner, model_name, version) - err := r.request(ctx, "POST", path, data, training) + err := r.fetch(ctx, "POST", path, data, training) if err != nil { return nil, fmt.Errorf("failed to create training: %w", err) } @@ -34,7 +34,7 @@ func (r *Client) CreateTraining(ctx context.Context, model_owner string, model_n // GetTraining sends a request to the Replicate API to get a training. func (r *Client) GetTraining(ctx context.Context, trainingID string) (*Training, error) { training := &Training{} - err := r.request(ctx, "GET", fmt.Sprintf("/trainings/%s", trainingID), nil, training) + err := r.fetch(ctx, "GET", fmt.Sprintf("/trainings/%s", trainingID), nil, training) if err != nil { return nil, fmt.Errorf("failed to get training: %w", err) } @@ -45,7 +45,7 @@ func (r *Client) GetTraining(ctx context.Context, trainingID string) (*Training, // CancelTraining sends a request to the Replicate API to cancel a training. func (r *Client) CancelTraining(ctx context.Context, trainingID string) (*Training, error) { training := &Training{} - err := r.request(ctx, "POST", fmt.Sprintf("/trainings/%s/cancel", trainingID), nil, training) + err := r.fetch(ctx, "POST", fmt.Sprintf("/trainings/%s/cancel", trainingID), nil, training) if err != nil { return nil, fmt.Errorf("failed to cancel training: %w", err) } @@ -56,7 +56,7 @@ func (r *Client) CancelTraining(ctx context.Context, trainingID string) (*Traini // ListTrainings returns a list of trainings. func (r *Client) ListTrainings(ctx context.Context) (*Page[Training], error) { response := &Page[Training]{} - err := r.request(ctx, "GET", "/trainings", nil, response) + err := r.fetch(ctx, "GET", "/trainings", nil, response) if err != nil { return nil, fmt.Errorf("failed to list trainings: %w", err) }