@@ -0,0 +1,356 @@
|
||||
package customer
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type DockerClient struct {
|
||||
hc *http.Client
|
||||
base string
|
||||
}
|
||||
|
||||
func NewDockerClient(raw string) (*DockerClient, error) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
raw = "unix:///var/run/docker.sock"
|
||||
}
|
||||
if strings.HasPrefix(raw, "unix://") {
|
||||
sock := strings.TrimPrefix(raw, "unix://")
|
||||
tr := &http.Transport{DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
var d net.Dialer
|
||||
return d.DialContext(ctx, "unix", sock)
|
||||
}}
|
||||
return &DockerClient{hc: &http.Client{Transport: tr, Timeout: 30 * time.Second}, base: "http://docker"}, nil
|
||||
}
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if u.Scheme != "http" && u.Scheme != "https" {
|
||||
return nil, errors.New("DOCKER_HOST must be unix://, http:// or https://")
|
||||
}
|
||||
return &DockerClient{hc: &http.Client{Timeout: 30 * time.Second}, base: strings.TrimRight(raw, "/")}, nil
|
||||
}
|
||||
|
||||
func (d *DockerClient) req(ctx context.Context, method, path string, in, out any) error {
|
||||
var body io.Reader
|
||||
if in != nil {
|
||||
b, err := json.Marshal(in)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
body = bytes.NewReader(b)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, d.base+path, body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if in != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
resp, err := d.hc.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
b, readErr := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
|
||||
if readErr != nil {
|
||||
return readErr
|
||||
}
|
||||
if resp.StatusCode/100 != 2 {
|
||||
return fmt.Errorf("docker API %s %s HTTP %d: %s", method, path, resp.StatusCode, strings.TrimSpace(string(b)))
|
||||
}
|
||||
if out != nil && len(bytes.TrimSpace(b)) > 0 {
|
||||
return json.Unmarshal(b, out)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (d *DockerClient) Ping(ctx context.Context) error {
|
||||
return d.req(ctx, http.MethodGet, "/_ping", nil, nil)
|
||||
}
|
||||
|
||||
// ImageExists checks the local Docker image cache without pulling anything.
|
||||
func (d *DockerClient) ImageExists(ctx context.Context, image string) (bool, error) {
|
||||
image = strings.TrimSpace(image)
|
||||
if image == "" {
|
||||
return false, errors.New("worker image is empty")
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, d.base+"/images/"+url.PathEscape(image)+"/json", nil)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
resp, err := d.hc.Do(req)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20))
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
return false, nil
|
||||
}
|
||||
if resp.StatusCode/100 != 2 {
|
||||
return false, fmt.Errorf("docker image inspect HTTP %d", resp.StatusCode)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// PullImage asks Docker Engine to pull a public/configured registry image.
|
||||
// Private-registry credentials are passed as Docker's X-Registry-Auth header.
|
||||
// The stream is inspected for daemon-side pull errors.
|
||||
func (d *DockerClient) PullImage(ctx context.Context, image, registryAuth string) error {
|
||||
image = strings.TrimSpace(image)
|
||||
if image == "" {
|
||||
return errors.New("worker image is empty")
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, d.base+"/images/create?fromImage="+url.QueryEscape(image), nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(registryAuth) != "" {
|
||||
req.Header.Set("X-Registry-Auth", strings.TrimSpace(registryAuth))
|
||||
}
|
||||
pullClient := *d.hc
|
||||
pullClient.Timeout = 10 * time.Minute
|
||||
resp, err := pullClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode/100 != 2 {
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10))
|
||||
return fmt.Errorf("docker image pull HTTP %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
|
||||
}
|
||||
dec := json.NewDecoder(io.LimitReader(resp.Body, 32<<20))
|
||||
for {
|
||||
var msg struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
if err := dec.Decode(&msg); err != nil {
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
return fmt.Errorf("docker image pull stream: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(msg.Error) != "" {
|
||||
return fmt.Errorf("docker image pull: %s", strings.TrimSpace(msg.Error))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegistryAuthHeader builds Docker Engine's X-Registry-Auth value. Use a
|
||||
// registry-scoped read-only deploy token instead of a personal password.
|
||||
func RegistryAuthHeader(username, password, serverAddress string) (string, error) {
|
||||
username = strings.TrimSpace(username)
|
||||
password = strings.TrimSpace(password)
|
||||
serverAddress = strings.TrimSpace(serverAddress)
|
||||
if username == "" && password == "" && serverAddress == "" {
|
||||
return "", nil
|
||||
}
|
||||
if username == "" || password == "" {
|
||||
return "", errors.New("both worker registry username and password/token are required")
|
||||
}
|
||||
payload := map[string]string{"username": username, "password": password}
|
||||
if serverAddress != "" {
|
||||
payload["serveraddress"] = serverAddress
|
||||
}
|
||||
b, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.URLEncoding.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
func (d *DockerClient) EnsureImage(ctx context.Context, image string, autoPull bool, registryAuth string) error {
|
||||
ok, err := d.ImageExists(ctx, image)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if ok {
|
||||
return nil
|
||||
}
|
||||
if !autoPull {
|
||||
return fmt.Errorf("worker image %q is not present on the Docker host and CS_WORKER_AUTO_PULL is disabled", image)
|
||||
}
|
||||
if err := d.PullImage(ctx, image, registryAuth); err != nil {
|
||||
return fmt.Errorf("pull worker image %q: %w", image, err)
|
||||
}
|
||||
ok, err = d.ImageExists(ctx, image)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return fmt.Errorf("worker image %q is still unavailable after pull", image)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (d *DockerClient) CreateVolume(ctx context.Context, name string) error {
|
||||
var out map[string]any
|
||||
return d.req(ctx, http.MethodPost, "/volumes/create", map[string]any{"Name": name, "Labels": map[string]string{"neuralhunt.managed": "true"}}, &out)
|
||||
}
|
||||
|
||||
type WorkerContainerConfig struct {
|
||||
Image, Entrypoint, Network, GameURL, RegisterURL, WorkerID, RegisterToken, TaskID, BeaconPath, Volume, Name string
|
||||
}
|
||||
|
||||
func (d *DockerClient) CreateWorker(ctx context.Context, c WorkerContainerConfig) (string, error) {
|
||||
name := url.QueryEscape(c.Name)
|
||||
body := map[string]any{
|
||||
"Image": c.Image,
|
||||
// Named Docker volumes are root-owned when first mounted. Managed workers
|
||||
// therefore run uid 0 only inside their own locked-down container so they
|
||||
// can create the 0600 identity file. They receive no Docker socket, all
|
||||
// Linux capabilities are dropped and the image root filesystem is read-only.
|
||||
"User": "0:0",
|
||||
"Cmd": []string{"-url", c.GameURL, "-identity", "/identity/identity.json", "-non-interactive", "-quiet", "-task", c.TaskID, "-beacon-path", c.BeaconPath},
|
||||
"Env": []string{
|
||||
"NEURALHUNT_WORKER_REGISTER_URL=" + c.RegisterURL,
|
||||
"NEURALHUNT_WORKER_LEASE_URL=" + strings.TrimSuffix(c.RegisterURL, "/register") + "/lease",
|
||||
"NEURALHUNT_WORKER_REGISTER_TOKEN=" + c.RegisterToken,
|
||||
"NEURALHUNT_WORKER_ID=" + c.WorkerID,
|
||||
},
|
||||
"Labels": map[string]string{"neuralhunt.managed": "true", "neuralhunt.worker_id": c.WorkerID},
|
||||
"HostConfig": map[string]any{
|
||||
"Mounts": []map[string]any{{"Type": "volume", "Source": c.Volume, "Target": "/identity"}},
|
||||
"NetworkMode": c.Network,
|
||||
"ReadonlyRootfs": true,
|
||||
"CapDrop": []string{"ALL"},
|
||||
"SecurityOpt": []string{"no-new-privileges"},
|
||||
"PidsLimit": 128,
|
||||
"Memory": 256 * 1024 * 1024,
|
||||
"NanoCpus": int64(1_000_000_000),
|
||||
},
|
||||
}
|
||||
// A dedicated worker image already declares /app/neuralhunt-client as its
|
||||
// ENTRYPOINT. Leaving Entrypoint unset makes CS_WORKER_IMAGE genuinely
|
||||
// pluggable. CS_WORKER_ENTRYPOINT exists only as a compatibility override
|
||||
// for older monolithic images.
|
||||
if strings.TrimSpace(c.Entrypoint) != "" {
|
||||
body["Entrypoint"] = []string{strings.TrimSpace(c.Entrypoint)}
|
||||
}
|
||||
var out struct {
|
||||
ID string `json:"Id"`
|
||||
}
|
||||
if err := d.req(ctx, http.MethodPost, "/containers/create?name="+name, body, &out); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if out.ID == "" {
|
||||
return "", errors.New("docker returned empty container id")
|
||||
}
|
||||
return out.ID, nil
|
||||
}
|
||||
func (d *DockerClient) Start(ctx context.Context, id string) error {
|
||||
return d.req(ctx, http.MethodPost, "/containers/"+url.PathEscape(id)+"/start", nil, nil)
|
||||
}
|
||||
func (d *DockerClient) Stop(ctx context.Context, id string, seconds int) error {
|
||||
if id == "" {
|
||||
return nil
|
||||
}
|
||||
if seconds < 1 {
|
||||
seconds = 10
|
||||
}
|
||||
return d.req(ctx, http.MethodPost, "/containers/"+url.PathEscape(id)+"/stop?t="+fmt.Sprint(seconds), nil, nil)
|
||||
}
|
||||
func (d *DockerClient) Remove(ctx context.Context, id string) error {
|
||||
if id == "" {
|
||||
return nil
|
||||
}
|
||||
err := d.req(ctx, http.MethodDelete, "/containers/"+url.PathEscape(id)+"?force=true&v=false", nil, nil)
|
||||
if err != nil && strings.Contains(err.Error(), "404") {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
func (d *DockerClient) Running(ctx context.Context, id string) (bool, error) {
|
||||
var out struct {
|
||||
State struct {
|
||||
Running bool `json:"Running"`
|
||||
} `json:"State"`
|
||||
}
|
||||
if err := d.req(ctx, http.MethodGet, "/containers/"+url.PathEscape(id)+"/json", nil, &out); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return out.State.Running, nil
|
||||
}
|
||||
|
||||
func (d *DockerClient) GetFile(ctx context.Context, containerID, path string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, d.base+"/containers/"+url.PathEscape(containerID)+"/archive?path="+url.QueryEscape(path), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := d.hc.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode/100 != 2 {
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return nil, fmt.Errorf("docker archive HTTP %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
|
||||
}
|
||||
tr := tar.NewReader(io.LimitReader(resp.Body, 8<<20))
|
||||
for {
|
||||
h, err := tr.Next()
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if filepath.Base(h.Name) == filepath.Base(path) && h.Typeflag == tar.TypeReg {
|
||||
return io.ReadAll(io.LimitReader(tr, 2<<20))
|
||||
}
|
||||
}
|
||||
return nil, errors.New("identity file not found in container volume")
|
||||
}
|
||||
func (d *DockerClient) PutFile(ctx context.Context, containerID, dir, name string, data []byte) error {
|
||||
var buf bytes.Buffer
|
||||
tw := tar.NewWriter(&buf)
|
||||
if err := tw.WriteHeader(&tar.Header{Name: name, Mode: 0600, Size: int64(len(data)), ModTime: time.Now()}); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := tw.Write(data); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tw.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPut, d.base+"/containers/"+url.PathEscape(containerID)+"/archive?path="+url.QueryEscape(dir), bytes.NewReader(buf.Bytes()))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/x-tar")
|
||||
resp, err := d.hc.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode/100 != 2 {
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return fmt.Errorf("docker put archive HTTP %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *DockerClient) RemoveVolume(ctx context.Context, name string) error {
|
||||
if name == "" {
|
||||
return nil
|
||||
}
|
||||
err := d.req(ctx, http.MethodDelete, "/volumes/"+url.PathEscape(name)+"?force=true", nil, nil)
|
||||
if err != nil && strings.Contains(err.Error(), "404") {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package customer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEnsureImagePullsMissingImage(t *testing.T) {
|
||||
var present atomic.Bool
|
||||
var pulls atomic.Int32
|
||||
image := "registry.example.com/neuralhunt/worker:v4.1"
|
||||
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.Method == http.MethodGet && strings.HasPrefix(r.URL.Path, "/images/") && strings.HasSuffix(r.URL.Path, "/json"):
|
||||
if !present.Load() {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"Id":"sha256:test"}`))
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/images/create":
|
||||
if got := r.URL.Query().Get("fromImage"); got != image {
|
||||
t.Fatalf("fromImage=%q want %q", got, image)
|
||||
}
|
||||
pulls.Add(1)
|
||||
present.Store(true)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte("{\"status\":\"Pull complete\"}\n"))
|
||||
default:
|
||||
http.Error(w, "unexpected request", http.StatusBadRequest)
|
||||
}
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
d := &DockerClient{hc: ts.Client(), base: ts.URL}
|
||||
if err := d.EnsureImage(context.Background(), image, true, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if pulls.Load() != 1 {
|
||||
t.Fatalf("pulls=%d want 1", pulls.Load())
|
||||
}
|
||||
if err := d.EnsureImage(context.Background(), image, true, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if pulls.Load() != 1 {
|
||||
t.Fatalf("second ensure pulled again: pulls=%d", pulls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureImageCanRequirePrePulledImage(t *testing.T) {
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
defer ts.Close()
|
||||
d := &DockerClient{hc: ts.Client(), base: ts.URL}
|
||||
if err := d.EnsureImage(context.Background(), "neuralhunt-worker:local", false, ""); err == nil || !strings.Contains(err.Error(), "CS_WORKER_AUTO_PULL") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegistryAuthHeader(t *testing.T) {
|
||||
h, err := RegistryAuthHeader("robot", "token", "registry.example.com")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if h == "" {
|
||||
t.Fatal("expected registry auth header")
|
||||
}
|
||||
if _, err := RegistryAuthHeader("robot", "", "registry.example.com"); err == nil {
|
||||
t.Fatal("incomplete registry credentials should fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateWorkerUsesImageEntrypointByDefault(t *testing.T) {
|
||||
var got map[string]any
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/containers/create" {
|
||||
http.Error(w, "unexpected", 400)
|
||||
return
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"Id":"container-1"}`))
|
||||
}))
|
||||
defer ts.Close()
|
||||
d := &DockerClient{hc: ts.Client(), base: ts.URL}
|
||||
|
||||
cfg := WorkerContainerConfig{Image: "neuralhunt-worker:local", Network: "nh", GameURL: "http://app:8080", RegisterURL: "http://cs:8092/internal/workers/register", WorkerID: "wrk_1", RegisterToken: "secret", TaskID: "task_1", BeaconPath: "auto", Volume: "vol_1", Name: "worker-1"}
|
||||
if _, err := d.CreateWorker(context.Background(), cfg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, exists := got["Entrypoint"]; exists {
|
||||
t.Fatalf("dedicated worker image should keep its image ENTRYPOINT: %#v", got["Entrypoint"])
|
||||
}
|
||||
|
||||
cfg.Name = "worker-2"
|
||||
cfg.Entrypoint = "/app/neuralhunt-client"
|
||||
if _, err := d.CreateWorker(context.Background(), cfg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, exists := got["Entrypoint"]; !exists {
|
||||
t.Fatal("compatibility Entrypoint override was not sent")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package customer
|
||||
|
||||
import (
|
||||
"crypto/elliptic"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
|
||||
"neuralhunt/internal/auth"
|
||||
)
|
||||
|
||||
var identityRawURL = base64.RawURLEncoding
|
||||
|
||||
type rawPrivateJWK struct {
|
||||
Kty string `json:"kty"`
|
||||
Crv string `json:"crv"`
|
||||
X string `json:"x"`
|
||||
Y string `json:"y"`
|
||||
D string `json:"d"`
|
||||
}
|
||||
|
||||
type rawIdentityFile struct {
|
||||
Version int `json:"version"`
|
||||
PublicJWK auth.PublicJWK `json:"publicJwk"`
|
||||
PrivateJWK rawPrivateJWK `json:"privateJwk"`
|
||||
}
|
||||
|
||||
func padP256(b []byte) []byte {
|
||||
out := make([]byte, 32)
|
||||
if len(b) > 32 {
|
||||
b = b[len(b)-32:]
|
||||
}
|
||||
copy(out[32-len(b):], b)
|
||||
return out
|
||||
}
|
||||
|
||||
// ValidateRawIdentity validates the portable raw CLI identity before it is
|
||||
// written into a managed worker's private Docker volume. It intentionally does
|
||||
// not accept the encrypted browser-export envelope because the service never
|
||||
// needs or asks for the customer's export passphrase.
|
||||
func ValidateRawIdentity(b []byte) error {
|
||||
var id rawIdentityFile
|
||||
if err := json.Unmarshal(b, &id); err != nil {
|
||||
return fmt.Errorf("invalid identity JSON: %w", err)
|
||||
}
|
||||
if id.Version != 1 || id.PublicJWK.Kty != "EC" || id.PublicJWK.Crv != "P-256" || id.PrivateJWK.Kty != "EC" || id.PrivateJWK.Crv != "P-256" || id.PrivateJWK.D == "" {
|
||||
return errors.New("unsupported identity; expected Neural Hunt version 1 P-256 raw identity")
|
||||
}
|
||||
db, err := identityRawURL.DecodeString(id.PrivateJWK.D)
|
||||
if err != nil {
|
||||
return errors.New("invalid private JWK encoding")
|
||||
}
|
||||
d := new(big.Int).SetBytes(db)
|
||||
curve := elliptic.P256()
|
||||
if d.Sign() <= 0 || d.Cmp(curve.Params().N) >= 0 {
|
||||
return errors.New("invalid P-256 private scalar")
|
||||
}
|
||||
x, y := curve.ScalarBaseMult(padP256(db))
|
||||
xs := identityRawURL.EncodeToString(padP256(x.Bytes()))
|
||||
ys := identityRawURL.EncodeToString(padP256(y.Bytes()))
|
||||
if xs != id.PublicJWK.X || ys != id.PublicJWK.Y || xs != id.PrivateJWK.X || ys != id.PrivateJWK.Y {
|
||||
return errors.New("identity public/private key mismatch")
|
||||
}
|
||||
if _, err := auth.ClientID(id.PublicJWK); err != nil {
|
||||
return fmt.Errorf("invalid public identity: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package customer
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type PayPalClient struct {
|
||||
ClientID, Secret, BaseURL string
|
||||
hc *http.Client
|
||||
mu sync.Mutex
|
||||
token string
|
||||
tokenExp time.Time
|
||||
}
|
||||
|
||||
func NewPayPalClient(clientID, secret, environment string) *PayPalClient {
|
||||
base := "https://api-m.sandbox.paypal.com"
|
||||
if strings.EqualFold(strings.TrimSpace(environment), "live") {
|
||||
base = "https://api-m.paypal.com"
|
||||
}
|
||||
return &PayPalClient{ClientID: strings.TrimSpace(clientID), Secret: strings.TrimSpace(secret), BaseURL: base, hc: &http.Client{Timeout: 20 * time.Second}}
|
||||
}
|
||||
func (p *PayPalClient) Ready() bool { return p.ClientID != "" && p.Secret != "" }
|
||||
func (p *PayPalClient) accessToken(ctx context.Context) (string, error) {
|
||||
p.mu.Lock()
|
||||
if p.token != "" && time.Until(p.tokenExp) > time.Minute {
|
||||
v := p.token
|
||||
p.mu.Unlock()
|
||||
return v, nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
if !p.Ready() {
|
||||
return "", errors.New("PayPal client ID/secret not configured")
|
||||
}
|
||||
form := url.Values{"grant_type": {"client_credentials"}}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.BaseURL+"/v1/oauth2/token", strings.NewReader(form.Encode()))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.SetBasicAuth(p.ClientID, p.Secret)
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
resp, err := p.hc.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
if resp.StatusCode/100 != 2 {
|
||||
return "", fmt.Errorf("PayPal OAuth HTTP %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
|
||||
}
|
||||
var out struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
ExpiresIn int `json:"expires_in"`
|
||||
}
|
||||
if err := json.Unmarshal(b, &out); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if out.AccessToken == "" {
|
||||
return "", errors.New("PayPal OAuth returned empty token")
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.token = out.AccessToken
|
||||
p.tokenExp = time.Now().Add(time.Duration(out.ExpiresIn) * time.Second)
|
||||
p.mu.Unlock()
|
||||
return out.AccessToken, nil
|
||||
}
|
||||
func (p *PayPalClient) call(ctx context.Context, method, path string, in, out any) error {
|
||||
tok, err := p.accessToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var rd io.Reader
|
||||
if in != nil {
|
||||
b, err := json.Marshal(in)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rd = bytes.NewReader(b)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, p.BaseURL+path, rd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
resp, err := p.hc.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
|
||||
if resp.StatusCode/100 != 2 {
|
||||
return fmt.Errorf("PayPal API %s HTTP %d: %s", path, resp.StatusCode, strings.TrimSpace(string(b)))
|
||||
}
|
||||
if out != nil && len(bytes.TrimSpace(b)) > 0 {
|
||||
return json.Unmarshal(b, out)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type PayPalOrderView struct {
|
||||
ID string `json:"id"`
|
||||
Status string `json:"status"`
|
||||
Links []struct{ Href, Rel, Method string } `json:"links"`
|
||||
PurchaseUnits []struct {
|
||||
Amount struct {
|
||||
CurrencyCode string `json:"currency_code"`
|
||||
Value string `json:"value"`
|
||||
} `json:"amount"`
|
||||
Payments struct {
|
||||
Captures []struct {
|
||||
ID, Status string
|
||||
Amount struct {
|
||||
CurrencyCode string `json:"currency_code"`
|
||||
Value string `json:"value"`
|
||||
} `json:"amount"`
|
||||
} `json:"captures"`
|
||||
} `json:"payments"`
|
||||
} `json:"purchase_units"`
|
||||
}
|
||||
|
||||
func (p *PayPalClient) CreateOrder(ctx context.Context, reference, description, amount, currency, returnURL, cancelURL string) (PayPalOrderView, error) {
|
||||
body := map[string]any{"intent": "CAPTURE", "purchase_units": []map[string]any{{"reference_id": reference, "description": description, "amount": map[string]string{"currency_code": currency, "value": amount}}}, "payment_source": map[string]any{"paypal": map[string]any{"experience_context": map[string]any{"shipping_preference": "NO_SHIPPING", "user_action": "PAY_NOW", "return_url": returnURL, "cancel_url": cancelURL}}}}
|
||||
var out PayPalOrderView
|
||||
err := p.call(ctx, http.MethodPost, "/v2/checkout/orders", body, &out)
|
||||
return out, err
|
||||
}
|
||||
func (p *PayPalClient) CaptureOrder(ctx context.Context, id string) (PayPalOrderView, error) {
|
||||
var out PayPalOrderView
|
||||
err := p.call(ctx, http.MethodPost, "/v2/checkout/orders/"+url.PathEscape(id)+"/capture", map[string]any{}, &out)
|
||||
return out, err
|
||||
}
|
||||
func (p *PayPalClient) GetOrder(ctx context.Context, id string) (PayPalOrderView, error) {
|
||||
var out PayPalOrderView
|
||||
err := p.call(ctx, http.MethodGet, "/v2/checkout/orders/"+url.PathEscape(id), nil, &out)
|
||||
return out, err
|
||||
}
|
||||
func ApprovalURL(o PayPalOrderView) string {
|
||||
for _, l := range o.Links {
|
||||
if l.Rel == "payer-action" || l.Rel == "approve" {
|
||||
return l.Href
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
func CaptureID(o PayPalOrderView) string {
|
||||
for _, u := range o.PurchaseUnits {
|
||||
for _, c := range u.Payments.Captures {
|
||||
if strings.EqualFold(c.Status, "COMPLETED") {
|
||||
return c.ID
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (p *PayPalClient) VerifyWebhook(ctx context.Context, webhookID string, h http.Header, event json.RawMessage) (bool, error) {
|
||||
if strings.TrimSpace(webhookID) == "" {
|
||||
return false, errors.New("PAYPAL_WEBHOOK_ID not configured")
|
||||
}
|
||||
var ev any
|
||||
if err := json.Unmarshal(event, &ev); err != nil {
|
||||
return false, err
|
||||
}
|
||||
body := map[string]any{"auth_algo": h.Get("PAYPAL-AUTH-ALGO"), "cert_url": h.Get("PAYPAL-CERT-URL"), "transmission_id": h.Get("PAYPAL-TRANSMISSION-ID"), "transmission_sig": h.Get("PAYPAL-TRANSMISSION-SIG"), "transmission_time": h.Get("PAYPAL-TRANSMISSION-TIME"), "webhook_id": webhookID, "webhook_event": ev}
|
||||
var out struct {
|
||||
VerificationStatus string `json:"verification_status"`
|
||||
}
|
||||
if err := p.call(ctx, http.MethodPost, "/v1/notifications/verify-webhook-signature", body, &out); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return strings.EqualFold(out.VerificationStatus, "SUCCESS"), nil
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package customer
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var rawURL = base64.RawURLEncoding
|
||||
|
||||
func RandomToken(n int) string {
|
||||
if n < 16 {
|
||||
n = 16
|
||||
}
|
||||
b := make([]byte, n)
|
||||
_, _ = rand.Read(b)
|
||||
return rawURL.EncodeToString(b)
|
||||
}
|
||||
|
||||
func NewPasswordHash(password string) (salt string, hash string, err error) {
|
||||
if len(strings.TrimSpace(password)) < 12 {
|
||||
return "", "", errors.New("password must be at least 12 characters")
|
||||
}
|
||||
s := make([]byte, 16)
|
||||
if _, err := rand.Read(s); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
h := pbkdf2SHA256([]byte(password), s, 310000, 32)
|
||||
return rawURL.EncodeToString(s), rawURL.EncodeToString(h), nil
|
||||
}
|
||||
func VerifyPassword(password, salt, hash string) bool {
|
||||
s, err := rawURL.DecodeString(salt)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
want, err := rawURL.DecodeString(hash)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
got := pbkdf2SHA256([]byte(password), s, 310000, len(want))
|
||||
return len(got) == len(want) && subtle.ConstantTimeCompare(got, want) == 1
|
||||
}
|
||||
func pbkdf2SHA256(password, salt []byte, iterations, keyLen int) []byte {
|
||||
hLen := sha256.Size
|
||||
blocks := (keyLen + hLen - 1) / hLen
|
||||
out := make([]byte, 0, blocks*hLen)
|
||||
for block := 1; block <= blocks; block++ {
|
||||
mac := hmac.New(sha256.New, password)
|
||||
mac.Write(salt)
|
||||
var n [4]byte
|
||||
binary.BigEndian.PutUint32(n[:], uint32(block))
|
||||
mac.Write(n[:])
|
||||
u := mac.Sum(nil)
|
||||
t := append([]byte(nil), u...)
|
||||
for i := 1; i < iterations; i++ {
|
||||
mac = hmac.New(sha256.New, password)
|
||||
mac.Write(u)
|
||||
u = mac.Sum(nil)
|
||||
for j := range t {
|
||||
t[j] ^= u[j]
|
||||
}
|
||||
}
|
||||
out = append(out, t...)
|
||||
}
|
||||
return out[:keyLen]
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,24 @@
|
||||
package customer
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseMoneyCentsExact(t *testing.T) {
|
||||
cases := map[string]int64{
|
||||
"0": 0,
|
||||
"1": 100,
|
||||
"1.2": 120,
|
||||
"1.20": 120,
|
||||
"499.99": 49999,
|
||||
}
|
||||
for in, want := range cases {
|
||||
got, err := parseMoneyCents(in)
|
||||
if err != nil || got != want {
|
||||
t.Fatalf("parseMoneyCents(%q) = %d, %v; want %d", in, got, err, want)
|
||||
}
|
||||
}
|
||||
for _, in := range []string{"", "-1.00", "+1.00", "1.234", "1,00", "abc"} {
|
||||
if _, err := parseMoneyCents(in); err == nil {
|
||||
t.Fatalf("parseMoneyCents(%q) should fail", in)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,425 @@
|
||||
package customer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
const schema = `
|
||||
PRAGMA foreign_keys=ON;
|
||||
CREATE TABLE IF NOT EXISTS customers(
|
||||
id TEXT PRIMARY KEY,
|
||||
username TEXT NOT NULL UNIQUE COLLATE NOCASE,
|
||||
password_salt TEXT NOT NULL,
|
||||
password_hash TEXT NOT NULL,
|
||||
reward_client_id TEXT NOT NULL DEFAULT '',
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS customer_sessions(
|
||||
id TEXT PRIMARY KEY,
|
||||
customer_id TEXT NOT NULL REFERENCES customers(id) ON DELETE CASCADE,
|
||||
expires_at INTEGER NOT NULL,
|
||||
created_at INTEGER NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS customer_sessions_exp_idx ON customer_sessions(expires_at);
|
||||
CREATE TABLE IF NOT EXISTS credit_ledger(
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
customer_id TEXT NOT NULL REFERENCES customers(id) ON DELETE CASCADE,
|
||||
delta_micros INTEGER NOT NULL,
|
||||
reason TEXT NOT NULL,
|
||||
reference TEXT NOT NULL UNIQUE,
|
||||
created_at INTEGER NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS credit_ledger_customer_idx ON credit_ledger(customer_id,created_at DESC);
|
||||
CREATE TABLE IF NOT EXISTS workers(
|
||||
id TEXT PRIMARY KEY,
|
||||
customer_id TEXT NOT NULL REFERENCES customers(id) ON DELETE CASCADE,
|
||||
task_id TEXT NOT NULL,
|
||||
beacon_path TEXT NOT NULL DEFAULT 'auto',
|
||||
docker_container_id TEXT NOT NULL DEFAULT '',
|
||||
docker_volume TEXT NOT NULL,
|
||||
status TEXT NOT NULL DEFAULT 'stopped' CHECK(status IN ('stopped','starting','running','error')),
|
||||
worker_client_id TEXT NOT NULL DEFAULT '',
|
||||
register_token TEXT NOT NULL,
|
||||
rate_micros_per_minute INTEGER NOT NULL,
|
||||
last_charge_at INTEGER,
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL,
|
||||
last_error TEXT NOT NULL DEFAULT ''
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS workers_customer_idx ON workers(customer_id,created_at);
|
||||
CREATE TABLE IF NOT EXISTS paypal_orders(
|
||||
order_id TEXT PRIMARY KEY,
|
||||
customer_id TEXT NOT NULL REFERENCES customers(id) ON DELETE CASCADE,
|
||||
package_id TEXT NOT NULL,
|
||||
amount_cents INTEGER NOT NULL,
|
||||
currency TEXT NOT NULL,
|
||||
credits_micros INTEGER NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
capture_id TEXT NOT NULL DEFAULT '',
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS paypal_orders_customer_idx ON paypal_orders(customer_id,created_at DESC);
|
||||
`
|
||||
|
||||
type Store struct{ DB *sql.DB }
|
||||
|
||||
func Open(ctx context.Context, path string) (*Store, error) {
|
||||
if strings.TrimSpace(path) == "" {
|
||||
path = "/customer-data/customer-service.db"
|
||||
}
|
||||
abs, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(abs), 0o750); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u := &url.URL{Scheme: "file", Path: filepath.ToSlash(abs)}
|
||||
q := u.Query()
|
||||
q.Add("_pragma", "busy_timeout(10000)")
|
||||
q.Add("_pragma", "foreign_keys(ON)")
|
||||
q.Add("_pragma", "synchronous(NORMAL)")
|
||||
q.Set("_txlock", "immediate")
|
||||
u.RawQuery = q.Encode()
|
||||
db, err := sql.Open("sqlite", u.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db.SetMaxOpenConns(4)
|
||||
if err := db.PingContext(ctx); err != nil {
|
||||
db.Close()
|
||||
return nil, err
|
||||
}
|
||||
if _, err := db.ExecContext(ctx, "PRAGMA journal_mode=WAL"); err != nil {
|
||||
db.Close()
|
||||
return nil, err
|
||||
}
|
||||
for i, stmt := range strings.Split(schema, ";") {
|
||||
stmt = strings.TrimSpace(stmt)
|
||||
if stmt == "" {
|
||||
continue
|
||||
}
|
||||
if _, err := db.ExecContext(ctx, stmt); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("customer schema %d: %w", i+1, err)
|
||||
}
|
||||
}
|
||||
return &Store{DB: db}, nil
|
||||
}
|
||||
|
||||
type Customer struct {
|
||||
ID string `json:"id"`
|
||||
Username string `json:"username"`
|
||||
RewardClientID string `json:"reward_client_id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type Worker struct {
|
||||
ID string `json:"id"`
|
||||
CustomerID string `json:"customer_id"`
|
||||
TaskID string `json:"task_id"`
|
||||
BeaconPath string `json:"beacon_path"`
|
||||
ContainerID string `json:"container_id"`
|
||||
Volume string `json:"volume"`
|
||||
Status string `json:"status"`
|
||||
WorkerClientID string `json:"worker_client_id"`
|
||||
RegisterToken string `json:"-"`
|
||||
RateMicrosPerMinute int64 `json:"rate_micros_per_minute"`
|
||||
LastChargeAt *time.Time `json:"last_charge_at,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
LastError string `json:"last_error,omitempty"`
|
||||
}
|
||||
|
||||
func (s *Store) CreateCustomer(ctx context.Context, id, username, salt, hash string) error {
|
||||
now := time.Now().UTC().UnixMilli()
|
||||
_, err := s.DB.ExecContext(ctx, `INSERT INTO customers(id,username,password_salt,password_hash,created_at,updated_at) VALUES(?,?,?,?,?,?)`, id, strings.TrimSpace(username), salt, hash, now, now)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) CustomerByUsername(ctx context.Context, username string) (Customer, string, string, error) {
|
||||
var c Customer
|
||||
var salt, hash string
|
||||
var created int64
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT id,username,reward_client_id,created_at,password_salt,password_hash FROM customers WHERE username=?`, strings.TrimSpace(username)).Scan(&c.ID, &c.Username, &c.RewardClientID, &created, &salt, &hash)
|
||||
c.CreatedAt = time.UnixMilli(created).UTC()
|
||||
return c, salt, hash, err
|
||||
}
|
||||
func (s *Store) CustomerByID(ctx context.Context, id string) (Customer, error) {
|
||||
var c Customer
|
||||
var created int64
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT id,username,reward_client_id,created_at FROM customers WHERE id=?`, id).Scan(&c.ID, &c.Username, &c.RewardClientID, &created)
|
||||
c.CreatedAt = time.UnixMilli(created).UTC()
|
||||
return c, err
|
||||
}
|
||||
func (s *Store) SetRewardClientID(ctx context.Context, id, cid string) error {
|
||||
_, err := s.DB.ExecContext(ctx, `UPDATE customers SET reward_client_id=?,updated_at=? WHERE id=?`, strings.TrimSpace(cid), time.Now().UTC().UnixMilli(), id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) CreateSession(ctx context.Context, sid, cid string, ttl time.Duration) error {
|
||||
now := time.Now().UTC()
|
||||
_, err := s.DB.ExecContext(ctx, `INSERT INTO customer_sessions(id,customer_id,expires_at,created_at) VALUES(?,?,?,?)`, sid, cid, now.Add(ttl).UnixMilli(), now.UnixMilli())
|
||||
return err
|
||||
}
|
||||
func (s *Store) SessionCustomer(ctx context.Context, sid string) (string, error) {
|
||||
var cid string
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT customer_id FROM customer_sessions WHERE id=? AND expires_at>?`, sid, time.Now().UTC().UnixMilli()).Scan(&cid)
|
||||
return cid, err
|
||||
}
|
||||
func (s *Store) DeleteSession(ctx context.Context, sid string) {
|
||||
_, _ = s.DB.ExecContext(ctx, `DELETE FROM customer_sessions WHERE id=?`, sid)
|
||||
}
|
||||
|
||||
func (s *Store) BalanceMicros(ctx context.Context, cid string) (int64, error) {
|
||||
var v sql.NullInt64
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT sum(delta_micros) FROM credit_ledger WHERE customer_id=?`, cid).Scan(&v)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return v.Int64, nil
|
||||
}
|
||||
func (s *Store) AddLedger(ctx context.Context, cid string, delta int64, reason, ref string) error {
|
||||
if delta == 0 {
|
||||
return errors.New("zero credit change")
|
||||
}
|
||||
_, err := s.DB.ExecContext(ctx, `INSERT INTO credit_ledger(customer_id,delta_micros,reason,reference,created_at) VALUES(?,?,?,?,?)`, cid, delta, reason, ref, time.Now().UTC().UnixMilli())
|
||||
return err
|
||||
}
|
||||
|
||||
type LedgerItem struct {
|
||||
DeltaMicros int64 `json:"delta_micros"`
|
||||
Reason string `json:"reason"`
|
||||
Reference string `json:"reference"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
func (s *Store) Ledger(ctx context.Context, cid string, limit int) ([]LedgerItem, error) {
|
||||
if limit < 1 || limit > 200 {
|
||||
limit = 50
|
||||
}
|
||||
rows, err := s.DB.QueryContext(ctx, `SELECT delta_micros,reason,reference,created_at FROM credit_ledger WHERE customer_id=? ORDER BY id DESC LIMIT ?`, cid, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []LedgerItem
|
||||
for rows.Next() {
|
||||
var x LedgerItem
|
||||
var ms int64
|
||||
if err := rows.Scan(&x.DeltaMicros, &x.Reason, &x.Reference, &ms); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x.CreatedAt = time.UnixMilli(ms).UTC()
|
||||
out = append(out, x)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func scanWorker(row interface{ Scan(...any) error }) (Worker, error) {
|
||||
var w Worker
|
||||
var last sql.NullInt64
|
||||
var created, updated int64
|
||||
err := row.Scan(&w.ID, &w.CustomerID, &w.TaskID, &w.BeaconPath, &w.ContainerID, &w.Volume, &w.Status, &w.WorkerClientID, &w.RegisterToken, &w.RateMicrosPerMinute, &last, &created, &updated, &w.LastError)
|
||||
if err != nil {
|
||||
return w, err
|
||||
}
|
||||
if last.Valid {
|
||||
v := time.UnixMilli(last.Int64).UTC()
|
||||
w.LastChargeAt = &v
|
||||
}
|
||||
w.CreatedAt = time.UnixMilli(created).UTC()
|
||||
w.UpdatedAt = time.UnixMilli(updated).UTC()
|
||||
return w, nil
|
||||
}
|
||||
|
||||
const workerCols = `id,customer_id,task_id,beacon_path,docker_container_id,docker_volume,status,worker_client_id,register_token,rate_micros_per_minute,last_charge_at,created_at,updated_at,last_error`
|
||||
|
||||
func (s *Store) CreateWorker(ctx context.Context, w Worker) error {
|
||||
now := time.Now().UTC().UnixMilli()
|
||||
_, err := s.DB.ExecContext(ctx, `INSERT INTO workers(id,customer_id,task_id,beacon_path,docker_volume,status,register_token,rate_micros_per_minute,created_at,updated_at) VALUES(?,?,?,?,?,'stopped',?,?,?,?)`, w.ID, w.CustomerID, w.TaskID, w.BeaconPath, w.Volume, w.RegisterToken, w.RateMicrosPerMinute, now, now)
|
||||
return err
|
||||
}
|
||||
func (s *Store) Worker(ctx context.Context, cid, wid string) (Worker, error) {
|
||||
return scanWorker(s.DB.QueryRowContext(ctx, `SELECT `+workerCols+` FROM workers WHERE id=? AND customer_id=?`, wid, cid))
|
||||
}
|
||||
func (s *Store) WorkerByID(ctx context.Context, wid string) (Worker, error) {
|
||||
return scanWorker(s.DB.QueryRowContext(ctx, `SELECT `+workerCols+` FROM workers WHERE id=?`, wid))
|
||||
}
|
||||
func (s *Store) Workers(ctx context.Context, cid string) ([]Worker, error) {
|
||||
rows, err := s.DB.QueryContext(ctx, `SELECT `+workerCols+` FROM workers WHERE customer_id=? ORDER BY created_at`, cid)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []Worker
|
||||
for rows.Next() {
|
||||
w, err := scanWorker(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, w)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (s *Store) WorkerCount(ctx context.Context, cid string, runningOnly bool) (int, error) {
|
||||
q := `SELECT count(*) FROM workers WHERE customer_id=?`
|
||||
if runningOnly {
|
||||
q += ` AND status='running'`
|
||||
}
|
||||
var n int
|
||||
err := s.DB.QueryRowContext(ctx, q, cid).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
func (s *Store) TotalWorkerCount(ctx context.Context) (int, error) {
|
||||
var n int
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT count(*) FROM workers`).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (s *Store) RunningWorkerCount(ctx context.Context) (int, error) {
|
||||
var n int
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT count(*) FROM workers WHERE status='running'`).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (s *Store) RunningWorkers(ctx context.Context) ([]Worker, error) {
|
||||
rows, err := s.DB.QueryContext(ctx, `SELECT `+workerCols+` FROM workers WHERE status='running' ORDER BY created_at`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []Worker
|
||||
for rows.Next() {
|
||||
w, err := scanWorker(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, w)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Store) ClaimWorkerStart(ctx context.Context, cid, wid string) (bool, error) {
|
||||
res, err := s.DB.ExecContext(ctx, `UPDATE workers SET status='starting',last_error='',updated_at=? WHERE id=? AND customer_id=? AND status IN ('stopped','error')`, time.Now().UTC().UnixMilli(), wid, cid)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
n, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return n == 1, nil
|
||||
}
|
||||
func (s *Store) RecoverStartingWorkers(ctx context.Context) error {
|
||||
_, err := s.DB.ExecContext(ctx, `UPDATE workers SET status='stopped',last_error='recovered after Customer Service restart',updated_at=? WHERE status='starting'`, time.Now().UTC().UnixMilli())
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) SetWorkerRuntime(ctx context.Context, wid, status, containerID, lastErr string) error {
|
||||
_, err := s.DB.ExecContext(ctx, `UPDATE workers SET status=?,docker_container_id=?,last_error=?,updated_at=? WHERE id=?`, status, containerID, lastErr, time.Now().UTC().UnixMilli(), wid)
|
||||
return err
|
||||
}
|
||||
func (s *Store) SetWorkerClient(ctx context.Context, wid, clientID string) error {
|
||||
_, err := s.DB.ExecContext(ctx, `UPDATE workers SET worker_client_id=?,updated_at=? WHERE id=?`, clientID, time.Now().UTC().UnixMilli(), wid)
|
||||
return err
|
||||
}
|
||||
func (s *Store) UpdateWorkerConfig(ctx context.Context, cid, wid, taskID, beaconPath string) error {
|
||||
_, err := s.DB.ExecContext(ctx, `UPDATE workers SET task_id=?,beacon_path=?,updated_at=? WHERE id=? AND customer_id=?`, taskID, beaconPath, time.Now().UTC().UnixMilli(), wid, cid)
|
||||
return err
|
||||
}
|
||||
func (s *Store) DeleteWorker(ctx context.Context, cid, wid string) error {
|
||||
_, err := s.DB.ExecContext(ctx, `DELETE FROM workers WHERE id=? AND customer_id=?`, wid, cid)
|
||||
return err
|
||||
}
|
||||
func (s *Store) MarkWorkerCharged(ctx context.Context, wid string, when time.Time) error {
|
||||
_, err := s.DB.ExecContext(ctx, `UPDATE workers SET last_charge_at=?,updated_at=? WHERE id=?`, when.UTC().UnixMilli(), time.Now().UTC().UnixMilli(), wid)
|
||||
return err
|
||||
}
|
||||
|
||||
// ChargeWorkerMinute debits one prepaid minute atomically. It never allows a
|
||||
// negative balance, so a billing loop can stop the worker as soon as funding is
|
||||
// exhausted.
|
||||
func (s *Store) ChargeWorkerMinute(ctx context.Context, w Worker, minute time.Time) (bool, error) {
|
||||
tx, err := s.DB.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
var bal sql.NullInt64
|
||||
if err := tx.QueryRowContext(ctx, `SELECT sum(delta_micros) FROM credit_ledger WHERE customer_id=?`, w.CustomerID).Scan(&bal); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if bal.Int64 < w.RateMicrosPerMinute {
|
||||
return false, nil
|
||||
}
|
||||
ref := fmt.Sprintf("worker:%s:%d", w.ID, minute.UTC().UnixMilli())
|
||||
if _, err := tx.ExecContext(ctx, `INSERT OR IGNORE INTO credit_ledger(customer_id,delta_micros,reason,reference,created_at) VALUES(?,?,?,?,?)`, w.CustomerID, -w.RateMicrosPerMinute, "worker_minute", ref, time.Now().UTC().UnixMilli()); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if _, err := tx.ExecContext(ctx, `UPDATE workers SET last_charge_at=?,updated_at=? WHERE id=?`, minute.UTC().UnixMilli(), time.Now().UTC().UnixMilli(), w.ID); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *Store) RefundWorkerStartMinute(ctx context.Context, w Worker, chargedAt time.Time, detail string) error {
|
||||
ref := fmt.Sprintf("worker_start_refund:%s:%d", w.ID, chargedAt.UTC().UnixMilli())
|
||||
reason := "worker_start_refund"
|
||||
if strings.TrimSpace(detail) != "" {
|
||||
reason += ":" + strings.TrimSpace(detail)
|
||||
}
|
||||
_, err := s.DB.ExecContext(ctx, `INSERT OR IGNORE INTO credit_ledger(customer_id,delta_micros,reason,reference,created_at) VALUES(?,?,?,?,?)`, w.CustomerID, w.RateMicrosPerMinute, reason, ref, time.Now().UTC().UnixMilli())
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) UpsertPayPalOrder(ctx context.Context, orderID, cid, pkg string, cents int64, currency string, credits int64, status string) error {
|
||||
now := time.Now().UTC().UnixMilli()
|
||||
_, err := s.DB.ExecContext(ctx, `INSERT INTO paypal_orders(order_id,customer_id,package_id,amount_cents,currency,credits_micros,status,created_at,updated_at) VALUES(?,?,?,?,?,?,?,?,?) ON CONFLICT(order_id) DO UPDATE SET status=excluded.status,updated_at=excluded.updated_at`, orderID, cid, pkg, cents, currency, credits, status, now, now)
|
||||
return err
|
||||
}
|
||||
|
||||
type PayPalOrder struct {
|
||||
OrderID, CustomerID, PackageID, Currency, Status, CaptureID string
|
||||
AmountCents, CreditsMicros int64
|
||||
}
|
||||
|
||||
func (s *Store) PayPalOrder(ctx context.Context, orderID string) (PayPalOrder, error) {
|
||||
var o PayPalOrder
|
||||
err := s.DB.QueryRowContext(ctx, `SELECT order_id,customer_id,package_id,amount_cents,currency,credits_micros,status,capture_id FROM paypal_orders WHERE order_id=?`, orderID).Scan(&o.OrderID, &o.CustomerID, &o.PackageID, &o.AmountCents, &o.Currency, &o.CreditsMicros, &o.Status, &o.CaptureID)
|
||||
return o, err
|
||||
}
|
||||
func (s *Store) CompletePayPalOrder(ctx context.Context, orderID, captureID string) error {
|
||||
tx, err := s.DB.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
var o PayPalOrder
|
||||
if err := tx.QueryRowContext(ctx, `SELECT order_id,customer_id,package_id,amount_cents,currency,credits_micros,status,capture_id FROM paypal_orders WHERE order_id=?`, orderID).Scan(&o.OrderID, &o.CustomerID, &o.PackageID, &o.AmountCents, &o.Currency, &o.CreditsMicros, &o.Status, &o.CaptureID); err != nil {
|
||||
return err
|
||||
}
|
||||
ref := "paypal:" + orderID
|
||||
if _, err := tx.ExecContext(ctx, `INSERT OR IGNORE INTO credit_ledger(customer_id,delta_micros,reason,reference,created_at) VALUES(?,?,?,?,?)`, o.CustomerID, o.CreditsMicros, "paypal_topup", ref, time.Now().UTC().UnixMilli()); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := tx.ExecContext(ctx, `UPDATE paypal_orders SET status='COMPLETED',capture_id=?,updated_at=? WHERE order_id=?`, captureID, time.Now().UTC().UnixMilli(), orderID); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
Reference in New Issue
Block a user