diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c3ce31b..6f9ff7a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,57 +7,38 @@ on: pull_request: jobs: - build: + test: runs-on: ubuntu-latest strategy: matrix: - go: [ '1.15' ] - name: Go ${{ matrix.go }} sample + # The version the module declares, and the newest release. They are the + # same toolchain until Go moves on, and no version is named here, so the + # floor is stated once, in go.mod. + include: + - name: declared + versionFile: go.mod + - name: latest + version: stable + name: Go ${{ matrix.name }} steps: - - uses: actions/checkout@v2 + - uses: actions/checkout@v4 - name: Setup go - uses: actions/setup-go@v2 + uses: actions/setup-go@v5 with: - go-version: ${{ matrix.go }} + go-version: ${{ matrix.version }} + go-version-file: ${{ matrix.versionFile }} - - name: Install Dependencies - run: | - sudo apt-get update; sudo apt-get install socat - - name: Run API test - run: | - wget https://raw.githubusercontent.com/UpdateHub/updatehub/master/doc/agent-http.yaml \ - -O agent-http.yaml - docker run \ - --rm \ - --detach=true \ - --name agent-sdk-go-mock \ - -p 8080:8000 \ - -v $PWD/agent-http.yaml:/api.yaml \ - danielgtaylor/apisprout@sha256:6c07143937e57095d8478efc8ab7eab52b44e67c7673285f8c0a2bf4a7b137ad \ - /api.yaml --validate-request - go run examples/api/main.go - - name: Run listener test - run: | - export UH_LISTENER_TEST=updatehub-statechange.sock - go run examples/listener/main.go & - - while [ ! -S "$UH_LISTENER_TEST" ]; do - sleep 1 - done + - name: Build + run: go build ./... - if [[ "$(echo "download" | socat - UNIX-CONNECT:updatehub-statechange.sock)" != "cancel" ]]; then - echo "Unexpected download response" + - name: Check the formatting + run: | + unformatted="$(gofmt -l .)" + if [ -n "$unformatted" ]; then + echo "gofmt reports:" + echo "$unformatted" exit 1 fi - if [[ "$(echo "install" | socat - UNIX-CONNECT:updatehub-statechange.sock)" != "" ]]; then - echo "Unexpected install response" - exit 2 - fi - if [[ "$(echo "error" | socat - UNIX-CONNECT:updatehub-statechange.sock)" != "" ]]; then - echo "Unexpected error response" - exit 3 - fi - if [[ "$(echo "reboot" | socat - UNIX-CONNECT:updatehub-statechange.sock)" != "" ]]; then - echo "Unexpected reboot response" - exit 4 - fi + + - name: Test + run: go test -race ./... diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 48b0c1f..632aa3f 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -9,8 +9,11 @@ jobs: name: lint runs-on: ubuntu-latest steps: - - uses: actions/checkout@v2 + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: 'stable' - name: golangci-lint - uses: golangci/golangci-lint-action@v2 + uses: golangci/golangci-lint-action@v8 with: - version: v1.29 + version: latest diff --git a/README.md b/README.md index 2a22bbb..8c82bea 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,143 @@ # updatehub agent SDK for Go -[![godoc](https://godoc.org/github.com/UpdateHub/agent-sdk-go?status.svg)](https://godoc.org/github.com/UpdateHub/agent-sdk-go) +[![godoc](https://pkg.go.dev/badge/github.com/UpdateHub/agent-sdk-go/v2.svg)](https://pkg.go.dev/github.com/UpdateHub/agent-sdk-go/v2) + +A Go client for the [UpdateHub](https://updatehub.io) agent's local API. It +speaks the agent's two transports: + +- `Client` calls the HTTP API the agent serves on `localhost:8080`. +- `StateChange` serves the Unix socket the agent connects to before every state + transition, so your program can veto one. + +This release is validated against UpdateHub agent **2.1.6**. + +## Install + +```sh +go get github.com/UpdateHub/agent-sdk-go/v2 +``` + +It needs Go 1.26 or later, and nothing outside the standard library. + +## Ask the agent to search for an update + +```go +client, err := updatehub.NewClient(updatehub.DefaultBaseURL, 30*time.Second, 5*time.Minute) +if err != nil { + return err +} + +response, err := client.Probe(ctx, "") +if err != nil { + return err +} + +switch response.Outcome { +case updatehub.ProbeUpdating: + // an update is available +case updatehub.ProbeNoUpdate: + // the server offers nothing new +case updatehub.ProbeTryAgain: + // the server asked for a back-off of response.TryAgainIn +case updatehub.ProbeBusy: + // the agent did not probe; it is in the response.BusyState state +} +``` + +Pass an empty custom server to probe the address the agent is configured with. +`ProbeBusy` is not a failure: the agent answers it, without reaching the server, +whenever it is in a state that does not accept a probe. One of those states is +named `error`, so read `Outcome` rather than the state name alone. + +**One turn per agent.** A `Client` serialises its calls, because agent 2.1.6 +panics a worker thread when several requests arrive together and then answers +nothing more until it restarts. The turn is held per agent rather than per +`Client`, so a second client built for the same base URL waits rather than +arriving beside the first — but the key is the base URL as written, so spell the +address the same way everywhere. + +A context bounds the wait, never the request. When a caller gives up, the +request keeps the turn until the agent answers it, until the connection breaks, +or until the hold runs out — so a call that gave up cannot leave a second one +arriving beside it. A call that then gives up waiting for its turn reports +`ErrCallOutstanding`, which tells you the agent is busy with an earlier request +of yours rather than slow to answer this one. + +The hold is the third argument to `NewClient`, and long is safer than short: the +agent may still be working on the abandoned request. It exists because an agent +can also stay alive and answer nothing, and something must guarantee the turn +comes back from that. + +## Veto a state transition + +```go +listener := updatehub.NewStateChange(updatehub.DefaultSocketPath) + +listener.OnError(func(err error) { + // one connection failed; the listener keeps serving +}) + +listener.OnState(updatehub.StateDownload, func(ctx context.Context, handler *updatehub.Handler) error { + return handler.Cancel() +}) + +if err := listener.Bind(); err != nil { + if errors.Is(err, syscall.EADDRINUSE) { + // somebody else holds the socket, or a dead process left the path + } + + return err +} +defer listener.Close() + +return listener.Serve(ctx) +``` + +A connection that closes with nothing written lets the agent proceed, so an +unbound socket is not a degraded veto: it is consent. The same is true when the +agent's trigger script is missing, because the agent then never consults the +socket at all. `TriggerInstalled(updatehub.TriggerPath)` reports whether the script is +there; what to do about its absence is your program's decision, not this +package's. + +## What changed in v2.0.0 + +The module path is now `github.com/UpdateHub/agent-sdk-go/v2`, and the API is +not compatible with what came before it. The rewrite is driven by one intended +consumer: a single-process firmware daemon on a device that has no way back if +the update channel breaks. + +- **Nothing exits the process.** `log.Fatal` is gone from every error path. +- **No method panics.** Each one used to end in a single-value type assertion on + a value that is nil on the error path, so an unreachable agent crashed the + caller. +- **Every call reads the HTTP status code**, and reports an unexpected one as a + `*StatusError` carrying the code and the body. +- **`Probe` sends no body when no custom server is given.** It used to send + `{"custom_server": ""}`, which agent 2.1.6 answers with a 500 — and keeps + answering with a 500 until it restarts. +- **`ProbeResponse` is a type**, not `interface{}`, and it carries all four + replies the agent can send. +- **`LocalInstall`, `RemoteInstall` and `AbortDownload` report a refusal** as a + result rather than as an error. +- **Calls to one agent are serialised**, through a turn held per base URL, and + the turn lasts until the request really ends rather than until the caller + stops waiting for it. +- **Configuration is a constructor parameter.** The base URL and the timeout go + to `NewClient`, the socket path goes to `NewStateChange`, and the + `UH_LISTENER_TEST` environment variable is deleted. +- **The listener binds without unlinking**, so a path that is already there + fails with `EADDRINUSE` instead of being taken silently. +- **The listener has a lifetime.** `Bind`, `Serve(ctx)` and `Close` replace a + `Listen` that blocked for ever, and it serves each connection in its own + goroutine. +- **A callback takes a context and returns an error**, and errors reach the sink + `OnError` registers. A callback that panics is reported there too, rather than + taking down a process that could not have recovered it. +- **The package writes nothing to stdout and nothing through `log`.** +- **`gorequest` is gone**, with the `goproxy` and `pkg/errors` dependencies it + pulled in. The package now needs the standard library alone. + +## License + +Licensed under the Apache License, Version 2.0. See [LICENSE](LICENSE). diff --git a/api.go b/api.go index f297ce9..58735af 100644 --- a/api.go +++ b/api.go @@ -1,222 +1,709 @@ +/* +UpdateHub +Copyright (C) 2019 +O.S. Systems Sofware LTDA: contato@ossystems.com.br + +SPDX-License-Identifier: Apache-2.0 +*/ + +// Package updatehub is a client for the UpdateHub agent's local API. +// +// It speaks two transports. Client calls the agent's HTTP API, by default on +// localhost:8080. StateChange serves the state-change socket the agent connects +// to before every transition, so a consumer can veto one. +// +// The package is validated against UpdateHub agent 2.1.6. +// +// Two properties hold across the whole package, because the intended consumer +// is a single-process firmware daemon on a device with no way back. No call +// panics, and no call writes to stdout or through the log package: every +// failure is returned to the caller, which decides what it means. package updatehub import ( + "bytes" + "context" "encoding/json" + "errors" "fmt" + "io" + "net/http" "net/url" - - "github.com/parnurzeal/gorequest" + "strings" + "sync" + "time" ) -type ProbeResponse interface{} +// DefaultBaseURL is where the agent serves its HTTP API out of the box. +const DefaultBaseURL = "http://localhost:8080" + +// maxResponseBytes bounds a reply, so a broken peer cannot exhaust the +// consumer's memory. A larger reply is an error of its own rather than a body +// cut short, which would reach the caller as a decode failure and read as an +// agent that answered nonsense. +const maxResponseBytes = 1 << 20 + +// State is a state of the UpdateHub agent's machine. The agent reports one in +// three places: the /info reply, the reply to an install request, and the reply +// to a probe it is too busy to serve. +type State string +// The states of agent 2.1.6, as its own State::name reports them. const ( - Updating = "updating" - NoUpdate = "no_update" - TryAgain = "try_again" + StateEntryPoint State = "entry_point" + StatePark State = "park" + StatePoll State = "poll" + StateProbe State = "probe" + StateValidation State = "validation" + StateDownload State = "download" + StateDirectDownload State = "direct_download" + StateInstall State = "install" + StateReboot State = "reboot" + StateError State = "error" + StatePrepareLocalInstall State = "prepare_local_install" ) -type APIState string +// ProbeOutcome discriminates the replies a probe can get. +type ProbeOutcome string const ( - Park = "park" - EntryPoint = "entry_point" - Poll = "poll" - Validation = "validation" - Download = "download" - Install = "install" - Reboot = "reboot" - DirectDownload = "direct_download" - PrepareLocalInstall = "prepare_local_install" - Error = "error" + // ProbeUpdating reports that an update is available. + ProbeUpdating ProbeOutcome = "updating" + // ProbeNoUpdate reports that the server offers nothing new. + ProbeNoUpdate ProbeOutcome = "no_update" + // ProbeTryAgain reports that the server asked for a back-off. + ProbeTryAgain ProbeOutcome = "try_again" + // ProbeBusy reports that the agent did not probe at all, because it was in + // a state that does not accept one. Any bare string other than the two + // above is read as this outcome, so a state a later agent adds arrives as a + // busy state rather than as a decode failure. + ProbeBusy ProbeOutcome = "busy" ) -type Client struct { +// ProbeResponse is the result of a probe. +// +// ProbeBusy is not a failure. The agent answers it, without reaching the +// server, whenever it is in a state that is not preemptive — and one of those +// states is named "error", which is why a caller must read Outcome rather than +// the state name alone. +type ProbeResponse struct { + // Outcome is which of the four replies the agent sent. + Outcome ProbeOutcome + // TryAgainIn is the back-off the server asked for. It is set only when + // Outcome is ProbeTryAgain. The agent states it in seconds. + TryAgainIn time.Duration + // BusyState is the state the agent was in when it refused the probe. It is + // set only when Outcome is ProbeBusy. + BusyState State +} + +// StateResponse is the agent's answer to an install request. The agent replies +// with the state it was in when the request arrived, and refuses the request +// when that state cannot start an installation. +type StateResponse struct { + // Accepted reports whether the agent took the request. + Accepted bool + // State is the state the agent was in when the request arrived. + State State +} + +// AbortDownloadResponse is the agent's answer to an abort request. +type AbortDownloadResponse struct { + // Accepted reports whether there was a download to abort. + Accepted bool + // Message is the agent's own wording, for the accepted case and for the + // refused one. + Message string +} + +// ErrCallOutstanding reports that a call gave up before it reached the agent, +// because an earlier call is still with the agent. It tells a caller that the +// agent is busy with its own earlier request, rather than slow to answer this +// one. +var ErrCallOutstanding = errors.New("updatehub: an earlier call is still with the agent") + +// StatusError reports an HTTP status the agent's API is not documented to send +// for the call that got it. It carries the status code and the body, because a +// consumer must be able to tell "the agent answered 500" from "the reply did +// not decode". +type StatusError struct { + Method string + Path string + StatusCode int + Body string +} + +func (e *StatusError) Error() string { + if e.Body == "" { + return fmt.Sprintf("updatehub: %s %s: unexpected status %d", e.Method, e.Path, e.StatusCode) + } + + return fmt.Sprintf("updatehub: %s %s: unexpected status %d: %s", e.Method, e.Path, e.StatusCode, e.Body) +} + +// Settings is the agent's configuration, as /info reports it. +type Settings struct { + Firmware FirmwareSettings `json:"firmware"` + Network Network `json:"network"` + Polling Polling `json:"polling"` + Storage Storage `json:"storage"` + Update Update `json:"update"` } -type Metadata struct { +// FirmwareSettings names where the agent reads the firmware metadata. +type FirmwareSettings struct { Metadata string `json:"metadata"` } +// Network is the agent's server address and the address its API listens on. type Network struct { ServerAddress string `json:"server_address"` ListenSocket string `json:"listen_socket"` } +// Polling is the agent's automatic poll configuration. Interval carries the +// agent's own wording, such as "1h". type Polling struct { Interval string `json:"interval"` Enabled bool `json:"enabled"` } +// Storage tells where the agent keeps its runtime settings. type Storage struct { ReadOnly bool `json:"read_only"` RuntimeSettings string `json:"runtime_settings"` } +// Update is the agent's download directory and the install modes it supports. type Update struct { DownloadDir string `json:"download_dir"` SupportedInstallModes []string `json:"supported_install_modes"` } -type ServerAddress struct { - Custom string `json:"custom"` -} +// MetadataStrings holds a device identity or attribute value. The agent writes +// a single value as a bare string and several as an array, so this type accepts +// both and always presents a slice. +type MetadataStrings []string + +// UnmarshalJSON accepts a bare string as well as an array of strings. +func (m *MetadataStrings) UnmarshalJSON(data []byte) error { + var single string + if err := json.Unmarshal(data, &single); err == nil { + *m = MetadataStrings{single} + return nil + } -type DeviceAttributes struct { - Attr1 string `json:"attr1"` - Attr2 string `json:"attr2"` -} + var many []string + if err := json.Unmarshal(data, &many); err != nil { + return err + } + *m = many -type DeviceIdentity struct { - ID1 string `json:"id1"` - ID2 string `json:"id2"` + return nil } -type Firmware struct { - DeviceAttributes DeviceAttributes `json:"device_attributes"` - DeviceIdentity DeviceIdentity `json:"device_identity"` - Hardware string `json:"hardware"` - PubKey string `json:"pub_key"` - Version string `json:"version"` +// MarshalJSON writes a single value as a bare string, as the agent does. +func (m MetadataStrings) MarshalJSON() ([]byte, error) { + if len(m) == 1 { + return json.Marshal(m[0]) + } + + return json.Marshal([]string(m)) } -type UpdatePackage struct { - AppliedPackageUid string `json:"applied_package_uid"` - UpdgradeToInstallation string `json:"upgrade_to_installation"` +// MetadataValue maps a device identity or attribute name to its values. +type MetadataValue map[string]MetadataStrings + +// FirmwareMetadata is the metadata the agent loaded from the running firmware. +type FirmwareMetadata struct { + ProductUID string `json:"product_uid"` + Version string `json:"version"` + Hardware string `json:"hardware"` + PubKey string `json:"pub_key"` + DeviceIdentity MetadataValue `json:"device_identity"` + DeviceAttributes MetadataValue `json:"device_attributes"` } -type Settings struct { - Firmware Metadata `json:"firmware"` - Network Network `json:"network"` - Polling Polling `json:"polling"` - Storage Storage `json:"storage"` - Update Update `json:"update"` +// ServerAddress is the address the agent polls: the address a probe set, or +// empty for the configured default. +type ServerAddress string + +// IsDefault reports whether the agent polls its configured address. +func (s ServerAddress) IsDefault() bool { return s == "" } + +// UnmarshalJSON accepts both shapes the agent sends: the bare string "default", +// and {"custom": "
"}. +func (s *ServerAddress) UnmarshalJSON(data []byte) error { + var name string + if err := json.Unmarshal(data, &name); err == nil { + if name != "default" { + return fmt.Errorf("updatehub: unknown server address %q", name) + } + *s = "" + + return nil + } + + var custom struct { + Custom *string `json:"custom"` + } + if err := json.Unmarshal(data, &custom); err != nil { + return err + } + if custom.Custom == nil { + return fmt.Errorf("updatehub: unknown server address %s", data) + } + *s = ServerAddress(*custom.Custom) + + return nil } -type RuntimeSettings struct { - Path string `json:"path"` - Persistent bool `json:"persistent"` - Polling PollingLog `json:"polling"` - Update UpdatePackage `json:"update"` +// MarshalJSON writes back the shape the agent sends. +func (s ServerAddress) MarshalJSON() ([]byte, error) { + if s.IsDefault() { + return json.Marshal("default") + } + + return json.Marshal(struct { + Custom string `json:"custom"` + }{Custom: string(s)}) } -type PollingLog struct { +// RuntimePolling is what the agent remembers about its own polling. +type RuntimePolling struct { + // Last is the timestamp of the last poll, in the agent's own wording. Last string `json:"last"` + Retries int `json:"retries"` Now bool `json:"now"` - Retries int64 `json:"retries"` ServerAddress ServerAddress `json:"server_address"` } +// RuntimeUpdate is what the agent remembers about the update in flight. +type RuntimeUpdate struct { + UpgradeToInstallation string `json:"upgrade_to_installation,omitempty"` + AppliedPackageUID string `json:"applied_package_uid,omitempty"` +} + +// RuntimeSettings is the agent's persisted state. +type RuntimeSettings struct { + Polling RuntimePolling `json:"polling"` + Update RuntimeUpdate `json:"update"` + Path string `json:"path"` + Persistent bool `json:"persistent"` +} + +// AgentInfo is the /info reply. Version is the agent's running version. type AgentInfo struct { - Config Settings `json:"config"` - Firmware Firmware `json:"firmware"` - RuntimeSettings RuntimeSettings `json:"runtime_settings"` - State APIState `json:"state"` - Version string `json:"version"` + State State `json:"state"` + Version string `json:"version"` + Config Settings `json:"config"` + Firmware FirmwareMetadata `json:"firmware"` + RuntimeSettings RuntimeSettings `json:"runtime_settings"` } -type Entry struct { - Data interface{} `json:"data"` - Level string `json:"level"` - Message string `json:"message"` - Time string `json:"time"` +// LogEntry is one entry of the agent's in-memory log. +type LogEntry struct { + Level string `json:"level"` + Message string `json:"message"` + Time string `json:"time"` + Data map[string]string `json:"data"` } +// Log is the /log reply. FirstIndex is the absolute index of the first entry +// among all the entries the agent ever recorded, so a reader can resume from +// where it stopped. type Log struct { - Entries []Entry `json:"entries"` + Entries []LogEntry `json:"entries"` + FirstIndex int `json:"first_index"` } -// NewClient instantiates a new updatehub agent client -func NewClient() *Client { - return &Client{} +// Client calls the UpdateHub agent's local HTTP API. +// +// A Client serialises its calls: one request reaches the agent at a time, and +// the others wait. This is not tidiness. Agent 2.1.6 panics a worker thread +// when several requests arrive together, and its API then answers nothing more +// until the agent restarts. +// +// The turn is held per agent rather than per Client, so a second Client built +// for the same base URL waits for the first one's call instead of arriving +// beside it. The turn is keyed on the base URL as given, so two spellings of +// one agent — localhost and 127.0.0.1 — are two turns, and a consumer that +// builds more than one Client should spell the address the same way. +// +// A context bounds the wait, never the request. When a caller gives up, the +// request keeps the turn until the agent answers it, until the connection +// breaks, or until the hold runs out, and only then does the next call start. +// The agent needs that: it hands each request to its state machine, and a +// request still in flight there widens the window in which the machine's own +// timer cancels the handler and panics the task waiting on it. Nothing on the +// wire says when the agent finishes, so the only safe reading of a call that +// gave up is that the agent is still busy — and a call that then gives up +// waiting for its turn reports ErrCallOutstanding, so a consumer can tell the +// two apart. +// +// A Client is safe for concurrent use. +type Client struct { + baseURL string + timeout time.Duration + hold time.Duration + http *http.Client + + // inFlight holds one token, and every Client for this agent holds the same + // channel. Taking the token is the serialising lock, and a call waiting for + // it still honours its context. + inFlight chan struct{} +} + +// agentTurns holds one token channel per agent base URL, so the serialising +// turn survives a consumer that builds a second Client. It never shrinks, and +// it is bounded by the number of distinct agent addresses a process talks to. +var agentTurns sync.Map + +// NewClient returns a client for the agent at baseURL. +// +// baseURL must carry a scheme and a host; pass DefaultBaseURL for the ordinary +// case. +// +// timeout must be positive. It bounds any call whose context carries no +// deadline of its own, because an agent that stops answering must not be able +// to block its caller for ever. +// +// hold must be positive. It is how long the client keeps the agent's turn after +// its caller gave up, before it cancels the request and lets the next call +// through. Long is safer than short: the agent may still be working on the +// abandoned request, and the turn exists to keep a second one away from it. But +// an agent can also stay alive and answer nothing, and the hold is what +// guarantees the turn comes back from that. +func NewClient(baseURL string, timeout, hold time.Duration) (*Client, error) { + if timeout <= 0 { + return nil, fmt.Errorf("updatehub: timeout must be positive, got %v", timeout) + } + if hold <= 0 { + return nil, fmt.Errorf("updatehub: hold must be positive, got %v", hold) + } + + parsed, err := url.Parse(baseURL) + if err != nil { + return nil, fmt.Errorf("updatehub: parse base URL %q: %w", baseURL, err) + } + if parsed.Scheme == "" || parsed.Host == "" { + return nil, fmt.Errorf("updatehub: base URL %q needs a scheme and a host", baseURL) + } + + address := strings.TrimSuffix(parsed.String(), "/") + turn, _ := agentTurns.LoadOrStore(address, make(chan struct{}, 1)) + + return &Client{ + baseURL: address, + timeout: timeout, + hold: hold, + http: &http.Client{}, + inFlight: turn.(chan struct{}), + }, nil } -// Probe server address for update -func (c *Client) Probe(serverAddress string) (*ProbeResponse, error) { - var probe ProbeResponse +// Probe asks the agent to search the server for an update. +// +// Pass an empty customServer to probe the address the agent is configured with. +// The request then carries no body at all, which is what the agent's own SDKs +// send: an empty custom_server field makes agent 2.1.6 resolve the address to +// an empty string, answer 500, and keep answering 500 until it restarts. +func (c *Client) Probe(ctx context.Context, customServer string) (ProbeResponse, error) { + var body any + if customServer != "" { + body = struct { + CustomServer string `json:"custom_server"` + }{CustomServer: customServer} + } - var req struct { - ServerAddress string `json:"custom_server"` + answer, err := c.do(ctx, http.MethodPost, "/probe", body) + if err != nil { + return ProbeResponse{}, err + } + if err := answer.expectOK(); err != nil { + return ProbeResponse{}, err } - req.ServerAddress = serverAddress - response, err := processRequest(string("/probe"), &probe, req, "POST") - return response.(*ProbeResponse), err + return decodeProbeResponse(answer.body) } -// GetInfo get updatehub agent general information -func (c *Client) GetInfo() (*AgentInfo, error) { - response, err := processRequest(string("/info"), &AgentInfo{}, nil, "GET") - return response.(*AgentInfo), err +// GetInfo reads the agent's general information, including the version it runs. +func (c *Client) GetInfo(ctx context.Context) (*AgentInfo, error) { + return getJSON[AgentInfo](ctx, c, "/info") } -// GetLogs get updatehub agent log entries -func (c *Client) GetLogs() (*Log, error) { - response, err := processRequest(string("/log"), &Log{}, nil, "GET") - return response.(*Log), err +// GetLogs reads the agent's in-memory log entries. +func (c *Client) GetLogs(ctx context.Context) (*Log, error) { + return getJSON[Log](ctx, c, "/log") } -// RemoteInstall trigger the installation of a package from a direct URL -func (c *Client) RemoteInstall(serverAddress string) (*APIState, error) { - var state APIState +// getJSON runs the calls that read a JSON document from the agent. +func getJSON[T any](ctx context.Context, c *Client, path string) (*T, error) { + answer, err := c.do(ctx, http.MethodGet, path, nil) + if err != nil { + return nil, err + } + if err := answer.expectOK(); err != nil { + return nil, err + } + + document := new(T) + if err := decodeJSON(answer.body, document); err != nil { + return nil, err + } + + return document, nil +} - var req struct { +// RemoteInstall asks the agent to install the package at packageURL. +func (c *Client) RemoteInstall(ctx context.Context, packageURL string) (StateResponse, error) { + body := struct { URL string `json:"url"` + }{URL: packageURL} + + return c.requestInstall(ctx, "/remote_install", body) +} + +// LocalInstall asks the agent to install the package already at filePath. +func (c *Client) LocalInstall(ctx context.Context, filePath string) (StateResponse, error) { + body := struct { + File string `json:"file"` + }{File: filePath} + + return c.requestInstall(ctx, "/local_install", body) +} + +// AbortDownload asks the agent to abort the download in flight. +func (c *Client) AbortDownload(ctx context.Context) (AbortDownloadResponse, error) { + const path = "/update/download/abort" + + answer, err := c.do(ctx, http.MethodPost, path, nil) + if err != nil { + return AbortDownloadResponse{}, err + } + + accepted, err := answer.accepted() + if err != nil { + return AbortDownloadResponse{}, err + } + + if !accepted { + var refused struct { + Error string `json:"error"` + } + if err := decodeJSON(answer.body, &refused); err != nil { + return AbortDownloadResponse{}, err + } + + return AbortDownloadResponse{Message: refused.Error}, nil + } + + var taken struct { + Message string `json:"message"` + } + if err := decodeJSON(answer.body, &taken); err != nil { + return AbortDownloadResponse{}, err } - req.URL = serverAddress - response, err := processRequest(string("/remote_install"), &state, req, "POST") - return response.(*APIState), err + return AbortDownloadResponse{Accepted: true, Message: taken.Message}, nil } -// LocalInstall trigger the installation of a local package -func (c *Client) LocalInstall(filePath string) (*APIState, error) { - var state APIState +// requestInstall runs the two install calls, which share a reply shape: the +// agent's state, at 200 when it took the request and at 406 when it refused. +func (c *Client) requestInstall(ctx context.Context, path string, body any) (StateResponse, error) { + answer, err := c.do(ctx, http.MethodPost, path, body) + if err != nil { + return StateResponse{}, err + } + + accepted, err := answer.accepted() + if err != nil { + return StateResponse{}, err + } - var req struct { - FilePath string `json:"file"` + var state State + if err := decodeJSON(answer.body, &state); err != nil { + return StateResponse{}, err } - req.FilePath = filePath - response, err := processRequest(string("/local_install"), &state, req, "POST") - return response.(*APIState), err + return StateResponse{Accepted: accepted, State: state}, nil +} + +// reply is one answer from the agent, together with the call that asked for +// it, so a caller reads the status and builds an error without repeating the +// method and the path a third time. +type reply struct { + method string + path string + status int + body []byte +} + +// accepted reports whether the agent took the request. The install and the +// abort calls answer 406 when the agent's state cannot serve them, which is a +// refusal to report rather than a failure. +func (r *reply) accepted() (bool, error) { + switch r.status { + case http.StatusOK: + return true, nil + case http.StatusNotAcceptable: + return false, nil + default: + return false, r.statusError() + } } -func (c *Client) AbortDownload() (*APIState, error) { - var state APIState +// expectOK returns an error for any status but 200. +func (r *reply) expectOK() error { + if r.status == http.StatusOK { + return nil + } - response, err := processRequest(string("/update/download/abort"), &state, nil, "POST") - return response.(*APIState), err + return r.statusError() } -func processRequest(url string, responseStruct interface{}, req interface{}, method string) (interface{}, error) { - var body []byte - var errs []error +func (r *reply) statusError() error { + return &StatusError{ + Method: r.method, + Path: r.path, + StatusCode: r.status, + Body: truncate(r.body), + } +} - switch method { - case "GET": - _, body, errs = gorequest.New().Get(buildURL(url)).EndStruct(&responseStruct) - case "POST": - _, body, errs = gorequest.New().Post(buildURL(url)).Send(req).EndStruct(&responseStruct) +// do sends one request. It returns an error only when the exchange itself +// failed; every status the agent sends reaches the caller. +func (c *Client) do(ctx context.Context, method, path string, body any) (*reply, error) { + if _, ok := ctx.Deadline(); !ok { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, c.timeout) + defer cancel() } - if len(errs) > 0 { - return nil, errs[0] + var payload io.Reader + if body != nil { + encoded, err := json.Marshal(body) + if err != nil { + return nil, fmt.Errorf("updatehub: encode %s request: %w", path, err) + } + payload = bytes.NewReader(encoded) } - err := json.Unmarshal([]byte(body), &responseStruct) + // The request outlives the caller's context on purpose, and c.hold is what + // ends it when the agent answers neither the caller nor anyone after it. + requestCtx, endRequest := context.WithCancel(context.WithoutCancel(ctx)) + + request, err := http.NewRequestWithContext(requestCtx, method, c.baseURL+path, payload) if err != nil { - return nil, err + endRequest() + + return nil, fmt.Errorf("updatehub: build %s request: %w", path, err) } + if body != nil { + request.Header.Set("Content-Type", "application/json") + } + + select { + case c.inFlight <- struct{}{}: + case <-ctx.Done(): + endRequest() + + return nil, fmt.Errorf("updatehub: %s %s: %w: %w", method, path, ErrCallOutstanding, ctx.Err()) + } + + done := make(chan roundTripResult, 1) + + go func() { + defer endRequest() + defer func() { <-c.inFlight }() + + answer, err := c.roundTrip(request, method, path) + done <- roundTripResult{answer: answer, err: err} + }() + + select { + case result := <-done: + return result.answer, result.err + case <-ctx.Done(): + time.AfterFunc(c.hold, endRequest) - return responseStruct, nil + return nil, fmt.Errorf("updatehub: %s %s: %w", method, path, ctx.Err()) + } +} + +// roundTripResult carries what an exchange produced back to the caller that +// may already have stopped waiting for it. +type roundTripResult struct { + answer *reply + err error } -func buildURL(path string) string { - u, err := url.Parse("localhost:8080") +// roundTrip runs one exchange to its end, even when the caller gave up. +func (c *Client) roundTrip(request *http.Request, method, path string) (*reply, error) { + response, err := c.http.Do(request) + if err != nil { + return nil, fmt.Errorf("updatehub: %s %s: %w", method, path, err) + } + defer func() { _ = response.Body.Close() }() + + answer, err := io.ReadAll(io.LimitReader(response.Body, maxResponseBytes+1)) if err != nil { - panic(err) + return nil, fmt.Errorf("updatehub: read %s reply: %w", path, err) + } + if len(answer) > maxResponseBytes { + return nil, fmt.Errorf("updatehub: %s %s: the reply is larger than %d bytes", method, path, maxResponseBytes) + } + + return &reply{method: method, path: path, status: response.StatusCode, body: answer}, nil +} + +func decodeJSON(body []byte, target any) error { + if err := json.Unmarshal(body, target); err != nil { + return fmt.Errorf("updatehub: decode reply %q: %w", truncate(body), err) + } + + return nil +} + +// decodeProbeResponse reads the four shapes a probe reply takes. Three are bare +// JSON strings and one is an object, because the agent serialises a Rust enum +// whose unit variants have no payload. +func decodeProbeResponse(body []byte) (ProbeResponse, error) { + var name string + if err := json.Unmarshal(body, &name); err == nil { + switch ProbeOutcome(name) { + case ProbeUpdating: + return ProbeResponse{Outcome: ProbeUpdating}, nil + case ProbeNoUpdate: + return ProbeResponse{Outcome: ProbeNoUpdate}, nil + default: + return ProbeResponse{Outcome: ProbeBusy, BusyState: State(name)}, nil + } + } + + var delayed struct { + TryAgain *int64 `json:"try_again"` + } + if err := decodeJSON(body, &delayed); err != nil { + return ProbeResponse{}, err + } + if delayed.TryAgain == nil { + return ProbeResponse{}, fmt.Errorf("updatehub: unknown probe reply %q", truncate(body)) + } + + return ProbeResponse{ + Outcome: ProbeTryAgain, + TryAgainIn: time.Duration(*delayed.TryAgain) * time.Second, + }, nil +} + +func truncate(body []byte) string { + const limit = 256 + + trimmed := bytes.TrimSpace(body) + if len(trimmed) <= limit { + return string(trimmed) } - return fmt.Sprintf("http://%s%s", u, path) + return string(trimmed[:limit]) + "…" } diff --git a/api_test.go b/api_test.go new file mode 100644 index 0000000..c239038 --- /dev/null +++ b/api_test.go @@ -0,0 +1,752 @@ +/* +UpdateHub +Copyright (C) 2019 +O.S. Systems Sofware LTDA: contato@ossystems.com.br + +SPDX-License-Identifier: Apache-2.0 +*/ + +package updatehub + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" +) + +func newTestClient(t *testing.T, handler http.HandlerFunc) *Client { + t.Helper() + + return newTestClientWithTimeout(t, 5*time.Second, handler) +} + +// newTestClientWithTimeout starts a test agent and returns a client for it. The +// server closes when the test ends, which is after the test's own deferred +// calls, so a handler that blocks is safe to release with a defer. +func newTestClientWithTimeout(t *testing.T, timeout time.Duration, handler http.HandlerFunc) *Client { + t.Helper() + + // The hold outlasts any test, so a test that wants a turn back asks for it. + return newTestClientWithHold(t, timeout, time.Minute, handler) +} + +func newTestClientWithHold(t *testing.T, timeout, hold time.Duration, handler http.HandlerFunc) *Client { + t.Helper() + + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + + client, err := NewClient(server.URL, timeout, hold) + if err != nil { + t.Fatalf("NewClient: %v", err) + } + + return client +} + +// trackPeak counts one request into the test agent and records the highest +// number it ever served at once. The returned call ends the request. +func trackPeak(inFlight, peak *atomic.Int32) func() { + current := inFlight.Add(1) + for { + seen := peak.Load() + if current <= seen || peak.CompareAndSwap(seen, current) { + break + } + } + + return func() { inFlight.Add(-1) } +} + +func TestNewClientRejectsUnusableConfiguration(t *testing.T) { + cases := []struct { + name string + baseURL string + timeout time.Duration + hold time.Duration + }{ + {"empty base URL", "", time.Second, time.Minute}, + {"no scheme", "localhost:8080", time.Second, time.Minute}, + {"no host", "http://", time.Second, time.Minute}, + {"zero timeout", DefaultBaseURL, 0, time.Minute}, + {"negative timeout", DefaultBaseURL, -time.Second, time.Minute}, + {"zero hold", DefaultBaseURL, time.Second, 0}, + {"negative hold", DefaultBaseURL, time.Second, -time.Minute}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if _, err := NewClient(tc.baseURL, tc.timeout, tc.hold); err == nil { + t.Fatalf("NewClient(%q, %v, %v) returned no error", tc.baseURL, tc.timeout, tc.hold) + } + }) + } +} + +func TestProbeSendsNoBodyWithoutCustomServer(t *testing.T) { + var ( + gotBody []byte + gotContentType string + ) + + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + if err != nil { + t.Errorf("read body: %v", err) + } + gotBody = body + gotContentType = r.Header.Get("Content-Type") + + if r.Method != http.MethodPost { + t.Errorf("method = %q, want POST", r.Method) + } + if r.URL.Path != "/probe" { + t.Errorf("path = %q, want /probe", r.URL.Path) + } + + _, _ = w.Write([]byte(`"no_update"`)) + }) + + if _, err := client.Probe(context.Background(), ""); err != nil { + t.Fatalf("Probe: %v", err) + } + + if len(gotBody) != 0 { + t.Errorf("request body = %q, want empty", gotBody) + } + if gotContentType != "" { + t.Errorf("Content-Type = %q, want none", gotContentType) + } +} + +func TestProbeSendsCustomServerWhenGiven(t *testing.T) { + var gotBody []byte + + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + gotBody, _ = io.ReadAll(r.Body) + _, _ = w.Write([]byte(`"updating"`)) + }) + + if _, err := client.Probe(context.Background(), "http://example.com:8080"); err != nil { + t.Fatalf("Probe: %v", err) + } + + var req struct { + CustomServer string `json:"custom_server"` + } + if err := json.Unmarshal(gotBody, &req); err != nil { + t.Fatalf("unmarshal request body %q: %v", gotBody, err) + } + if req.CustomServer != "http://example.com:8080" { + t.Errorf("custom_server = %q, want %q", req.CustomServer, "http://example.com:8080") + } +} + +func TestProbeDecodesEveryReplyShape(t *testing.T) { + cases := []struct { + name string + body string + want ProbeResponse + }{ + {"update available", `"updating"`, ProbeResponse{Outcome: ProbeUpdating}}, + {"no update", `"no_update"`, ProbeResponse{Outcome: ProbeNoUpdate}}, + {"back-off", `{"try_again":3600}`, ProbeResponse{Outcome: ProbeTryAgain, TryAgainIn: time.Hour}}, + {"busy downloading", `"download"`, ProbeResponse{Outcome: ProbeBusy, BusyState: StateDownload}}, + {"busy in the error state, which is not a failure", `"error"`, ProbeResponse{Outcome: ProbeBusy, BusyState: StateError}}, + {"busy installing", `"install"`, ProbeResponse{Outcome: ProbeBusy, BusyState: StateInstall}}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(tc.body)) + }) + + got, err := client.Probe(context.Background(), "") + if err != nil { + t.Fatalf("Probe: %v", err) + } + if got != tc.want { + t.Errorf("Probe() = %+v, want %+v", got, tc.want) + } + }) + } +} + +func TestProbeReturnsTheStatusCodeOnFailure(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte("Unhandled rejection: Client(UrlParse(RelativeUrlWithoutBase))")) + }) + + _, err := client.Probe(context.Background(), "") + if err == nil { + t.Fatal("Probe returned no error for a 500") + } + + var statusErr *StatusError + if !errors.As(err, &statusErr) { + t.Fatalf("error %v is not a *StatusError", err) + } + if statusErr.StatusCode != http.StatusInternalServerError { + t.Errorf("StatusCode = %d, want 500", statusErr.StatusCode) + } + if statusErr.Body == "" { + t.Error("StatusError carries no body") + } +} + +func TestProbeReturnsAnErrorForAnUndecodableBody(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("not json")) + }) + + if _, err := client.Probe(context.Background(), ""); err == nil { + t.Fatal("Probe returned no error for a body that is not JSON") + } +} + +func TestAReplyTooLargeToTrustIsAnErrorOfItsOwn(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write(append([]byte(`"`), bytes.Repeat([]byte("x"), (1<<20)+64)...)) + }) + + _, err := client.Probe(context.Background(), "") + if err == nil { + t.Fatal("Probe returned no error for an oversized reply") + } + if !strings.Contains(err.Error(), "larger than") { + t.Errorf("error = %v, want it to name the size limit rather than a decode failure", err) + } +} + +func TestEveryMethodReturnsAnErrorWhenTheAgentIsDown(t *testing.T) { + server := httptest.NewServer(http.NotFoundHandler()) + url := server.URL + server.Close() + + client, err := NewClient(url, 2*time.Second, time.Minute) + if err != nil { + t.Fatalf("NewClient: %v", err) + } + + ctx := context.Background() + calls := map[string]func() error{ + "Probe": func() error { _, err := client.Probe(ctx, ""); return err }, + "GetInfo": func() error { _, err := client.GetInfo(ctx); return err }, + "GetLogs": func() error { _, err := client.GetLogs(ctx); return err }, + "LocalInstall": func() error { _, err := client.LocalInstall(ctx, "/tmp/x.uhupkg"); return err }, + "RemoteInstall": func() error { _, err := client.RemoteInstall(ctx, "https://example.com/x.uhupkg"); return err }, + "AbortDownload": func() error { _, err := client.AbortDownload(ctx); return err }, + } + + for name, call := range calls { + t.Run(name, func(t *testing.T) { + defer func() { + if r := recover(); r != nil { + t.Fatalf("%s panicked: %v", name, r) + } + }() + + if err := call(); err == nil { + t.Fatalf("%s returned no error with no agent listening", name) + } + }) + } +} + +func TestClientSerialisesEveryCall(t *testing.T) { + var ( + inFlight atomic.Int32 + peak atomic.Int32 + ) + + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + defer trackPeak(&inFlight, &peak)() + + time.Sleep(10 * time.Millisecond) + _, _ = w.Write([]byte(`"no_update"`)) + }) + + var wg sync.WaitGroup + for range 10 { + wg.Go(func() { + if _, err := client.Probe(context.Background(), ""); err != nil { + t.Errorf("Probe: %v", err) + } + }) + } + wg.Wait() + + if got := peak.Load(); got != 1 { + t.Errorf("peak concurrent requests = %d, want 1", got) + } +} + +func TestAnAbandonedCallKeepsTheNextOneOutOfTheAgent(t *testing.T) { + var ( + inFlight atomic.Int32 + peak atomic.Int32 + first = make(chan struct{}) + served atomic.Int32 + ) + + // The first request outlives the caller that sent it; every later one runs + // straight through. + release := sync.OnceFunc(func() { close(first) }) + + defer release() + + client := newTestClientWithTimeout(t, 100*time.Millisecond, func(w http.ResponseWriter, _ *http.Request) { + defer trackPeak(&inFlight, &peak)() + + if served.Add(1) == 1 { + <-first + } + + _, _ = w.Write([]byte(`"no_update"`)) + }) + + if _, err := client.Probe(context.Background(), ""); err == nil { + t.Fatal("the first Probe returned no error, want its deadline to bound it") + } + + // The agent is still holding the abandoned request. The second call must + // wait for it rather than arrive beside it. + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + answered := make(chan error, 1) + go func() { + _, err := client.Probe(ctx, "") + answered <- err + }() + + select { + case <-answered: + t.Fatal("the second Probe ran while the first was still with the agent") + case <-time.After(200 * time.Millisecond): + } + + release() + + select { + case err := <-answered: + if err != nil { + t.Fatalf("the second Probe: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("the second Probe never ran") + } + + if got := peak.Load(); got != 1 { + t.Errorf("peak concurrent requests = %d, want 1", got) + } +} + +func TestACallThatNeverReachesTheAgentSaysSo(t *testing.T) { + release := make(chan struct{}) + defer close(release) + + client := newTestClientWithTimeout(t, 100*time.Millisecond, func(w http.ResponseWriter, _ *http.Request) { + <-release + _, _ = w.Write([]byte(`"no_update"`)) + }) + + first, err := client.Probe(context.Background(), "") + if err == nil { + t.Fatalf("the first Probe returned %+v, want its deadline to bound it", first) + } + if errors.Is(err, ErrCallOutstanding) { + t.Error("the first Probe reported ErrCallOutstanding; it did reach the agent") + } + + _, err = client.Probe(context.Background(), "") + if !errors.Is(err, ErrCallOutstanding) { + t.Errorf("the second Probe returned %v, want ErrCallOutstanding", err) + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Errorf("the second Probe returned %v, want it to carry the deadline too", err) + } +} + +func TestTheHoldGivesTheTurnBackWhenTheAgentNeverAnswers(t *testing.T) { + var ( + first = make(chan struct{}) + served atomic.Int32 + ) + + release := sync.OnceFunc(func() { close(first) }) + defer release() + + client := newTestClientWithHold(t, 50*time.Millisecond, 100*time.Millisecond, + func(w http.ResponseWriter, _ *http.Request) { + // Only the first request goes unanswered. A later one must not + // queue behind it inside the handler. + if served.Add(1) == 1 { + <-first + } + + _, _ = w.Write([]byte(`"no_update"`)) + }) + + if _, err := client.Probe(context.Background(), ""); err == nil { + t.Fatal("the first Probe returned no error, want its deadline to bound it") + } + + // The agent never answered the abandoned request, so only the hold can + // bring the turn back. + time.Sleep(300 * time.Millisecond) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + if _, err := client.Probe(ctx, ""); err != nil { + t.Fatalf("the second Probe: %v", err) + } +} + +func TestTwoClientsForOneAgentShareTheTurn(t *testing.T) { + var ( + inFlight atomic.Int32 + peak atomic.Int32 + ) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + defer trackPeak(&inFlight, &peak)() + + time.Sleep(20 * time.Millisecond) + _, _ = w.Write([]byte(`"no_update"`)) + })) + defer server.Close() + + var clients []*Client + for range 2 { + client, err := NewClient(server.URL, 5*time.Second, time.Minute) + if err != nil { + t.Fatalf("NewClient: %v", err) + } + clients = append(clients, client) + } + + if clients[0] == clients[1] { + t.Fatal("NewClient returned the same client twice; the test proves nothing") + } + + var wg sync.WaitGroup + for _, client := range clients { + wg.Go(func() { + if _, err := client.Probe(context.Background(), ""); err != nil { + t.Errorf("Probe: %v", err) + } + }) + } + wg.Wait() + + if got := peak.Load(); got != 1 { + t.Errorf("peak concurrent requests = %d, want 1; two clients for one agent must share the turn", got) + } +} + +func TestTheTimeoutBoundsACall(t *testing.T) { + release := make(chan struct{}) + defer close(release) + + client := newTestClientWithTimeout(t, 50*time.Millisecond, func(w http.ResponseWriter, _ *http.Request) { + <-release + _, _ = w.Write([]byte(`"no_update"`)) + }) + + start := time.Now() + if _, err := client.Probe(context.Background(), ""); err == nil { + t.Fatal("Probe returned no error for a hung agent") + } + if elapsed := time.Since(start); elapsed > 2*time.Second { + t.Errorf("Probe blocked for %v, want the 50ms timeout to bound it", elapsed) + } +} + +func TestACallerDeadlineOverridesTheClientTimeout(t *testing.T) { + release := make(chan struct{}) + defer close(release) + + client := newTestClientWithTimeout(t, time.Hour, func(w http.ResponseWriter, _ *http.Request) { + <-release + _, _ = w.Write([]byte(`"no_update"`)) + }) + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + + start := time.Now() + if _, err := client.Probe(ctx, ""); err == nil { + t.Fatal("Probe returned no error for a hung agent") + } + if elapsed := time.Since(start); elapsed > 2*time.Second { + t.Errorf("Probe blocked for %v, want the caller deadline to bound it", elapsed) + } +} + +func TestLocalInstallReadsTheRequestAndTheReply(t *testing.T) { + var gotBody []byte + + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/local_install" { + t.Errorf("path = %q, want /local_install", r.URL.Path) + } + gotBody, _ = io.ReadAll(r.Body) + _, _ = w.Write([]byte(`"probe"`)) + }) + + got, err := client.LocalInstall(context.Background(), "/tmp/update.uhupkg") + if err != nil { + t.Fatalf("LocalInstall: %v", err) + } + if !got.Accepted { + t.Error("Accepted = false, want true for a 200") + } + if got.State != StateProbe { + t.Errorf("State = %q, want %q", got.State, StateProbe) + } + + var req struct { + File string `json:"file"` + } + if err := json.Unmarshal(gotBody, &req); err != nil { + t.Fatalf("unmarshal request body %q: %v", gotBody, err) + } + if req.File != "/tmp/update.uhupkg" { + t.Errorf("file = %q, want %q", req.File, "/tmp/update.uhupkg") + } +} + +func TestLocalInstallReportsARefusalWithoutAnError(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNotAcceptable) + _, _ = w.Write([]byte(`"install"`)) + }) + + got, err := client.LocalInstall(context.Background(), "/tmp/update.uhupkg") + if err != nil { + t.Fatalf("LocalInstall returned an error for a 406: %v", err) + } + if got.Accepted { + t.Error("Accepted = true, want false for a 406") + } + if got.State != StateInstall { + t.Errorf("State = %q, want %q", got.State, StateInstall) + } +} + +func TestLocalInstallReturnsTheStatusCodeOnAnUnexpectedFailure(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte("Request body deserialize error")) + }) + + _, err := client.LocalInstall(context.Background(), "/tmp/update.uhupkg") + if err == nil { + t.Fatal("LocalInstall returned no error for a 400") + } + + var statusErr *StatusError + if !errors.As(err, &statusErr) { + t.Fatalf("error %v is not a *StatusError", err) + } + if statusErr.StatusCode != http.StatusBadRequest { + t.Errorf("StatusCode = %d, want 400", statusErr.StatusCode) + } +} + +func TestRemoteInstallSendsTheURL(t *testing.T) { + var gotBody []byte + + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/remote_install" { + t.Errorf("path = %q, want /remote_install", r.URL.Path) + } + gotBody, _ = io.ReadAll(r.Body) + _, _ = w.Write([]byte(`"entry_point"`)) + }) + + got, err := client.RemoteInstall(context.Background(), "https://example.com/update.uhupkg") + if err != nil { + t.Fatalf("RemoteInstall: %v", err) + } + if !got.Accepted || got.State != StateEntryPoint { + t.Errorf("RemoteInstall() = %+v, want an accepted entry_point", got) + } + + var req struct { + URL string `json:"url"` + } + if err := json.Unmarshal(gotBody, &req); err != nil { + t.Fatalf("unmarshal request body %q: %v", gotBody, err) + } + if req.URL != "https://example.com/update.uhupkg" { + t.Errorf("url = %q, want %q", req.URL, "https://example.com/update.uhupkg") + } +} + +func TestAbortDownloadReportsBothReplies(t *testing.T) { + t.Run("accepted", func(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/update/download/abort" { + t.Errorf("path = %q, want /update/download/abort", r.URL.Path) + } + _, _ = w.Write([]byte(`{"message":"request accepted, download aborted"}`)) + }) + + got, err := client.AbortDownload(context.Background()) + if err != nil { + t.Fatalf("AbortDownload: %v", err) + } + if !got.Accepted { + t.Error("Accepted = false, want true") + } + if got.Message != "request accepted, download aborted" { + t.Errorf("Message = %q", got.Message) + } + }) + + t.Run("refused", func(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNotAcceptable) + _, _ = w.Write([]byte(`{"error":"there is no download to be aborted"}`)) + }) + + got, err := client.AbortDownload(context.Background()) + if err != nil { + t.Fatalf("AbortDownload returned an error for a 406: %v", err) + } + if got.Accepted { + t.Error("Accepted = true, want false") + } + if got.Message != "there is no download to be aborted" { + t.Errorf("Message = %q", got.Message) + } + }) +} + +func TestGetInfoReadsTheAgentVersion(t *testing.T) { + const body = `{ + "state": "park", + "version": "2.1.6", + "config": { + "firmware": {"metadata": "/usr/share/updatehub"}, + "network": {"server_address": "https://api.updatehub.io", "listen_socket": "localhost:8080"}, + "polling": {"interval": "1h", "enabled": true}, + "storage": {"read_only": false, "runtime_settings": "/var/lib/updatehub/runtime_settings.conf"}, + "update": {"download_dir": "/tmp", "supported_install_modes": ["copy", "raw"]} + }, + "firmware": { + "product_uid": "0123", + "version": "8.0.1", + "hardware": "ema40i", + "pub_key": null, + "device_identity": {"id1": "value", "id2": ["a", "b"]}, + "device_attributes": {} + }, + "runtime_settings": { + "polling": { + "last": "2026-08-27T12:00:00Z", + "retries": 0, + "now": false, + "server_address": "default" + }, + "update": {}, + "path": "/var/lib/updatehub/runtime_settings.conf", + "persistent": true + } + }` + + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + t.Errorf("method = %q, want GET", r.Method) + } + if r.URL.Path != "/info" { + t.Errorf("path = %q, want /info", r.URL.Path) + } + _, _ = w.Write([]byte(body)) + }) + + info, err := client.GetInfo(context.Background()) + if err != nil { + t.Fatalf("GetInfo: %v", err) + } + if info.Version != "2.1.6" { + t.Errorf("Version = %q, want 2.1.6", info.Version) + } + if info.State != StatePark { + t.Errorf("State = %q, want park", info.State) + } + if got := info.Config.Polling.Interval; got != "1h" { + t.Errorf("Config.Polling.Interval = %q, want 1h", got) + } + if got := info.Firmware.DeviceIdentity["id1"]; len(got) != 1 || got[0] != "value" { + t.Errorf("DeviceIdentity[id1] = %q, want [value]", got) + } + if got := info.Firmware.DeviceIdentity["id2"]; len(got) != 2 { + t.Errorf("DeviceIdentity[id2] = %q, want two values", got) + } + if !info.RuntimeSettings.Polling.ServerAddress.IsDefault() { + t.Error("ServerAddress.IsDefault() = false, want true") + } + if got := info.RuntimeSettings.Polling.ServerAddress; got != "" { + t.Errorf("ServerAddress = %q, want empty for the default", got) + } +} + +func TestGetInfoReadsACustomServerAddress(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(`{"version":"2.1.6","runtime_settings":{"polling":{"server_address":{"custom":"https://example.com"}}}}`)) + }) + + info, err := client.GetInfo(context.Background()) + if err != nil { + t.Fatalf("GetInfo: %v", err) + } + + address := info.RuntimeSettings.Polling.ServerAddress + if address.IsDefault() { + t.Fatal("IsDefault() = true, want false") + } + if address != "https://example.com" { + t.Errorf("ServerAddress = %q, want https://example.com", address) + } +} + +func TestGetLogsDecodesTheEntries(t *testing.T) { + client := newTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/log" { + t.Errorf("path = %q, want /log", r.URL.Path) + } + _, _ = w.Write([]byte(`{"entries":[{"level":"info","message":"probing","time":"2026-08-27T12:00:00Z","data":{"k":"v"}}],"first_index":7}`)) + }) + + logs, err := client.GetLogs(context.Background()) + if err != nil { + t.Fatalf("GetLogs: %v", err) + } + if len(logs.Entries) != 1 { + t.Fatalf("len(Entries) = %d, want 1", len(logs.Entries)) + } + if logs.Entries[0].Message != "probing" { + t.Errorf("Message = %q, want probing", logs.Entries[0].Message) + } + if logs.Entries[0].Data["k"] != "v" { + t.Errorf("Data = %v, want k=v", logs.Entries[0].Data) + } + if logs.FirstIndex != 7 { + t.Errorf("FirstIndex = %d, want 7", logs.FirstIndex) + } +} diff --git a/examples/api/main.go b/examples/api/main.go index f34711c..c544243 100644 --- a/examples/api/main.go +++ b/examples/api/main.go @@ -7,90 +7,79 @@ SPDX-License-Identifier: Apache-2.0 package main import ( + "context" "encoding/json" "fmt" "log" + "time" - updatehub "github.com/UpdateHub/agent-sdk-go" + updatehub "github.com/UpdateHub/agent-sdk-go/v2" ) func main() { - client := updatehub.NewClient() - - logs, err := client.GetLogs() - if err != nil { - log.Fatal(err) - } - - resp, err := json.Marshal(logs) - if err != nil { - log.Fatal(err) - } - fmt.Println(string(resp) + "\n") - - info, err := client.GetInfo() - if err != nil { - log.Fatal(err) - } - - resp, err = json.Marshal(info) + // The hold is how long the client keeps the agent's turn after a call gives + // up, so that the next call cannot arrive beside a request the agent may + // still be working on. + client, err := updatehub.NewClient(updatehub.DefaultBaseURL, 30*time.Second, 5*time.Minute) if err != nil { log.Fatal(err) } - fmt.Println(string(resp) + "\n") - probe, err := client.Probe("") - if err != nil { - log.Fatal(err) - } + ctx := context.Background() - resp, err = json.Marshal(probe) + logs, err := client.GetLogs(ctx) if err != nil { log.Fatal(err) } - fmt.Println(string(resp) + "\n") + dump(logs) - probeCustom, err := client.Probe("http://www.example.com:8080") + info, err := client.GetInfo(ctx) if err != nil { log.Fatal(err) } + fmt.Println("agent version:", info.Version) - resp, err = json.Marshal(probeCustom) + // An empty custom server probes the address the agent is configured with. + probe, err := client.Probe(ctx, "") if err != nil { log.Fatal(err) } - fmt.Println(string(resp) + "\n") - remoteInstall, err := client.RemoteInstall("https://foo.bar/update.uhu") - if err != nil { - log.Fatal(err) + switch probe.Outcome { + case updatehub.ProbeUpdating: + fmt.Println("an update is available") + case updatehub.ProbeNoUpdate: + fmt.Println("no update is available") + case updatehub.ProbeTryAgain: + fmt.Println("the server asked for a back-off of", probe.TryAgainIn) + case updatehub.ProbeBusy: + fmt.Println("the agent did not probe; it is in the", probe.BusyState, "state") } - resp, err = json.Marshal(remoteInstall) + remoteInstall, err := client.RemoteInstall(ctx, "https://foo.bar/update.uhupkg") if err != nil { log.Fatal(err) } - fmt.Println(string(resp) + "\n") + dump(remoteInstall) - localInstall, err := client.LocalInstall("/tmp/update.uhu") + localInstall, err := client.LocalInstall(ctx, "/tmp/update.uhupkg") if err != nil { log.Fatal(err) } + dump(localInstall) - resp, err = json.Marshal(localInstall) + abortDownload, err := client.AbortDownload(ctx) if err != nil { log.Fatal(err) } - fmt.Println(string(resp) + "\n") + dump(abortDownload) +} - abortDownload, err := client.AbortDownload() +func dump(value any) { + encoded, err := json.Marshal(value) if err != nil { log.Fatal(err) } - resp, err = json.Marshal(abortDownload) - if err != nil { - log.Fatal(err) - } - fmt.Println(string(resp) + "\n") + fmt.Println(string(encoded)) } diff --git a/examples/listener/main.go b/examples/listener/main.go index 90ee2c5..93ccdb8 100644 --- a/examples/listener/main.go +++ b/examples/listener/main.go @@ -7,27 +7,62 @@ SPDX-License-Identifier: Apache-2.0 package main import ( + "context" + "errors" "fmt" "log" + "os" + "os/signal" + "syscall" - updatehub "github.com/UpdateHub/agent-sdk-go" + updatehub "github.com/UpdateHub/agent-sdk-go/v2" ) func main() { - listener := updatehub.NewStateChange() + socketPath := updatehub.DefaultSocketPath + if len(os.Args) > 1 { + socketPath = os.Args[1] + } - listener.OnState(updatehub.StateDownload, func(handler *updatehub.Handler) { - fmt.Println("function called when starting the Download state; it will cancel the transition") - handler.Cancel() + installed, err := updatehub.TriggerInstalled(updatehub.TriggerPath) + if err != nil { + log.Fatal(err) + } + if !installed { + fmt.Println("the trigger script is missing; the agent will not consult this socket") + } + + listener := updatehub.NewStateChange(socketPath) + + listener.OnError(func(err error) { + fmt.Println("state change error:", err) }) - listener.OnState(updatehub.StateInstall, func(handler *updatehub.Handler) { - fmt.Println("function called when starting the Install state") - handler.Proceed() + listener.OnState(updatehub.StateDownload, func(_ context.Context, handler *updatehub.Handler) error { + fmt.Println("the agent is about to download; cancelling the transition") + + return handler.Cancel() }) - err := listener.Listen() - if err != nil { + listener.OnState(updatehub.StateInstall, func(_ context.Context, handler *updatehub.Handler) error { + fmt.Println("the agent is about to install; letting it proceed") + + return handler.Proceed() + }) + + if err := listener.Bind(); err != nil { + if errors.Is(err, syscall.EADDRINUSE) { + fmt.Println("another process holds", socketPath) + } + + log.Fatal(err) + } + defer func() { _ = listener.Close() }() + + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + if err := listener.Serve(ctx); err != nil { log.Fatal(err) } } diff --git a/go.mod b/go.mod index 5158f32..d5f8860 100644 --- a/go.mod +++ b/go.mod @@ -1,12 +1,3 @@ -module github.com/UpdateHub/agent-sdk-go +module github.com/UpdateHub/agent-sdk-go/v2 -go 1.15 - -require ( - github.com/elazarl/goproxy v0.0.0-20200809112317-0581fc3aee2d // indirect - github.com/parnurzeal/gorequest v0.2.16 - github.com/pkg/errors v0.9.1 // indirect - github.com/smartystreets/goconvey v1.6.4 // indirect - golang.org/x/net v0.0.0-20200904194848-62affa334b73 // indirect - moul.io/http2curl v1.0.0 // indirect -) +go 1.26 diff --git a/go.sum b/go.sum deleted file mode 100644 index 425e25b..0000000 --- a/go.sum +++ /dev/null @@ -1,30 +0,0 @@ -github.com/elazarl/goproxy v0.0.0-20200809112317-0581fc3aee2d h1:rtM8HsT3NG37YPjz8sYSbUSdElP9lUsQENYzJDZDUBE= -github.com/elazarl/goproxy v0.0.0-20200809112317-0581fc3aee2d/go.mod h1:Ro8st/ElPeALwNFlcTpWmkr6IoMFfkjXAvTHpevnDsM= -github.com/elazarl/goproxy/ext v0.0.0-20190711103511-473e67f1d7d2 h1:dWB6v3RcOy03t/bUadywsbyrQwCqZeNIEX6M1OtSZOM= -github.com/elazarl/goproxy/ext v0.0.0-20190711103511-473e67f1d7d2/go.mod h1:gNh8nYJoAm43RfaxurUnxr+N1PwuFV3ZMl/efxlIlY8= -github.com/gopherjs/gopherjs v0.0.0-20181017120253-0766667cb4d1 h1:EGx4pi6eqNxGaHF6qqu48+N2wcFQ5qg5FXgOdqsJ5d8= -github.com/gopherjs/gopherjs v0.0.0-20181017120253-0766667cb4d1/go.mod h1:wJfORRmW1u3UXTncJ5qlYoELFm8eSnnEO6hX4iZ3EWY= -github.com/jtolds/gls v4.20.0+incompatible h1:xdiiI2gbIgH/gLH7ADydsJ1uDOEzR8yvV7C0MuV77Wo= -github.com/jtolds/gls v4.20.0+incompatible/go.mod h1:QJZ7F/aHp+rZTRtaJ1ow/lLfFfVYBRgL+9YlvaHOwJU= -github.com/parnurzeal/gorequest v0.2.16 h1:T/5x+/4BT+nj+3eSknXmCTnEVGSzFzPGdpqmUVVZXHQ= -github.com/parnurzeal/gorequest v0.2.16/go.mod h1:3Kh2QUMJoqw3icWAecsyzkpY7UzRfDhbRdTjtNwNiUE= -github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= -github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= -github.com/rogpeppe/go-charset v0.0.0-20180617210344-2471d30d28b4/go.mod h1:qgYeAmZ5ZIpBWTGllZSQnw97Dj+woV0toclVaRGI8pc= -github.com/smartystreets/assertions v0.0.0-20180927180507-b2de0cb4f26d h1:zE9ykElWQ6/NYmHa3jpm/yHnI4xSofP+UP6SpjHcSeM= -github.com/smartystreets/assertions v0.0.0-20180927180507-b2de0cb4f26d/go.mod h1:OnSkiWE9lh6wB0YB77sQom3nweQdgAjqCqsofrRNTgc= -github.com/smartystreets/goconvey v1.6.4 h1:fv0U8FUIMPNf1L9lnHLvLhgicrIVChEkdzIKYqbNC9s= -github.com/smartystreets/goconvey v1.6.4/go.mod h1:syvi0/a8iFYH4r/RixwvyeAJjdLS9QV7WQ/tjFTllLA= -golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= -golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= -golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= -golang.org/x/net v0.0.0-20200904194848-62affa334b73 h1:MXfv8rhZWmFeqX3GNZRsd6vOLoaCHjYEX3qkRo3YBUA= -golang.org/x/net v0.0.0-20200904194848-62affa334b73/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA= -golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= -golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20200323222414-85ca7c5b95cd/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= -golang.org/x/tools v0.0.0-20190328211700-ab21143f2384/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= -moul.io/http2curl v1.0.0 h1:6XwpyZOYsgZJrU8exnG87ncVkU1FVCcTRpwzOkTDUi8= -moul.io/http2curl v1.0.0/go.mod h1:f6cULg+e4Md/oW1cYmwW4IWQOVl2lGbmCNGOHvzX2kE= diff --git a/listener.go b/listener.go index e726fbc..097b692 100644 --- a/listener.go +++ b/listener.go @@ -10,125 +10,328 @@ package updatehub import ( "bufio" + "context" + "errors" "fmt" - "log" + "io" "net" "os" "strings" + "sync" + "time" ) -const SDKTriggerFilename string = "/usr/share/updatehub/state-change-callbacks.d/10-updatehub-sdk-statechange-trigger" -const SocketPath string = "/run/updatehub-statechange.sock" +const ( + // DefaultSocketPath is the state-change socket the agent's trigger script + // connects to. + DefaultSocketPath = "/run/updatehub-statechange.sock" -// CallbackFunc the type of the callbacks. -type CallbackFunc func(handler *Handler) + // TriggerPath is the script the agent runs before every state transition. + // The script connects to the state-change socket and hands the agent + // whatever the socket answers. Without it the agent never consults the + // socket at all, so a listener that is bound and correct still vetoes + // nothing. Read TriggerInstalled to find out. + TriggerPath = "/usr/share/updatehub/state-change-callbacks.d/10-updatehub-sdk-statechange-trigger" +) -// StateChange struct that store the callbacks for a state. -type StateChange struct { - Listeners map[string][]CallbackFunc -} +// readBufferSize holds the longest state name the agent sends, with room to +// spare. A longer line still reads correctly; it takes one more fill. +const readBufferSize = 64 -// State Represent the states of UpdateHub Agent can handle. -type State string +// cancelReply is the exact byte string the agent reads as a veto. Anything else +// it reads, including a trailing newline, is an error it also treats as a +// cancel, so the reply carries no newline. +const cancelReply = "cancel" -const ( - StateProbe = "probe" - StateDownload = "download" - StateInstall = "install" - StateReboot = "reboot" - StateError = "error" -) +// ErrAlreadyReplied reports a second answer to one state change. The agent +// reads one reply per connection. +var ErrAlreadyReplied = errors.New("updatehub: the state change was already answered") -/// Handler used to communicate with UpdateHub -/// to call commands on the state callbacks. +// CallbackFunc answers one state change. +// +// The callback must answer through the handler, by calling Cancel to veto the +// transition or Proceed to allow it. The listener closes the connection once +// the callback returns, and a connection that closes with nothing written lets +// the agent proceed — so a callback that returns an error without answering +// still lets the transition through. +// +// ctx is the context Serve runs under. It is cancelled when Serve returns. +// +// The listener runs one callback per connection, in its own goroutine, so a +// slow callback delays only the agent transition it answers. Callbacks for +// different connections can run at the same time. A callback that panics is +// reported to the sink OnError registers, like one that returns an error. +type CallbackFunc func(ctx context.Context, handler *Handler) error + +// Handler answers one state change. It is not safe for concurrent use. type Handler struct { - conn net.Conn + conn net.Conn + state State + replied bool + writeErr error } -// Cancel cancels the current state. -func (h Handler) Cancel() { - _, err := h.conn.Write([]byte("cancel")) - checkErr(err) +// State reports the state the agent is about to enter. +func (h *Handler) State() State { return h.state } + +// Cancel vetoes the transition. The agent then resets its machine and enters +// its entry point again. +func (h *Handler) Cancel() error { + if h.replied { + return ErrAlreadyReplied + } + h.replied = true + + if _, err := h.conn.Write([]byte(cancelReply)); err != nil { + h.writeErr = fmt.Errorf("updatehub: answer the %s state change: %w", h.state, err) + + return h.writeErr + } + + return nil } -// Proceed proceeds to the next state. -func (h Handler) Proceed() {} +// Proceed lets the transition happen. It writes nothing, which is how the agent +// reads consent. +func (h *Handler) Proceed() error { + if h.replied { + return ErrAlreadyReplied + } + h.replied = true + + return nil +} -// NewStateChange instantiates a new StateChange. -func NewStateChange() *StateChange { +// StateChange serves the agent's state-change socket. +// +// The agent connects to the socket before every transition and blocks until it +// gets an answer, so an unbound socket is not a degraded veto: the trigger +// script prints nothing, and the agent reads that as consent. +// +// Register the callbacks first, then Bind, then Serve. Bind and Serve are +// separate because a bind that fails with EADDRINUSE is a decision the caller +// must make — another process holds the socket, or a previous process left the +// path behind — and this package never makes it for them. +// +// A StateChange is safe for concurrent use. +type StateChange struct { + socketPath string + + mu sync.Mutex + callbacks map[State]CallbackFunc + onError func(error) + listener net.Listener + closed bool + + handlers sync.WaitGroup +} + +// NewStateChange returns a listener for the socket at socketPath. Pass +// DefaultSocketPath for the ordinary case. +func NewStateChange(socketPath string) *StateChange { return &StateChange{ - Listeners: make(map[string][]CallbackFunc), + socketPath: socketPath, + callbacks: make(map[State]CallbackFunc), } } -// OnState register the callbacks for a state passed as argument. +// OnState registers the callback for one state. A second registration for the +// same state replaces the first. A state with no callback is answered by +// closing the connection, which lets the agent proceed. func (sc *StateChange) OnState(state State, f CallbackFunc) { - name := strings.Join([]string{string(state)}, "") - sc.Listeners[name] = append(sc.Listeners[name], f) + sc.mu.Lock() + defer sc.mu.Unlock() + + sc.callbacks[state] = f } -/// Listen start the agent to listen for messages on the socket. -func (sc *StateChange) Listen() error { - _, err := os.Stat(SDKTriggerFilename) - if err != nil && os.IsNotExist(err) { - fmt.Println("WARNING: updatehub-sdk-statechange-trigger not found on", SDKTriggerFilename) - } +// OnError registers the sink for the errors one connection raises: a failed +// read, a callback error, or a reply that could not be written. Serve keeps +// serving after each of them, because a listener that stops leaves the agent +// blocked on the next transition. +// +// The sink is called from the goroutine that serves the connection, so it can +// run for several connections at once. A listener with no sink discards these +// errors: this package writes nothing to stdout and nothing through log. +func (sc *StateChange) OnError(f func(error)) { + sc.mu.Lock() + defer sc.mu.Unlock() - ln, err := createListener() - checkErr(err) + sc.onError = f +} - for { - conn, err := ln.Accept() - checkErr(err) +// Bind binds the socket. +// +// It never unlinks the path first. A path that is already there fails with +// EADDRINUSE, which errors.Is reports through the returned error, and the +// caller decides: a live binder means another process holds the agent's only +// veto channel, while a stale path from a dead process is the caller's to +// remove before it binds again. +func (sc *StateChange) Bind() error { + sc.mu.Lock() + defer sc.mu.Unlock() + + if sc.listener != nil { + return fmt.Errorf("updatehub: %q is already bound", sc.socketPath) + } - sc.handleConn(conn) + listener, err := net.Listen("unix", sc.socketPath) + if err != nil { + return fmt.Errorf("updatehub: bind %q: %w", sc.socketPath, err) } + sc.listener = listener + sc.closed = false + + return nil } -func (sc *StateChange) handleConn(c net.Conn) { - buf := bufio.NewReader(c) +// Serve accepts connections and answers them until Close is called or ctx is +// done, and then waits for the callbacks still running. +// +// It returns nil for a stop the caller asked for, which includes a Close that +// arrives before it starts, and an error for a stop it did not ask for: a +// failed Accept, or a Serve on a listener that was never bound. A single +// connection never stops it; those errors reach the sink OnError registers. +func (sc *StateChange) Serve(ctx context.Context) error { + sc.mu.Lock() + listener, closed := sc.listener, sc.closed + sc.mu.Unlock() + + if listener == nil { + if closed { + return nil + } + + return fmt.Errorf("updatehub: serve %q: bind it first", sc.socketPath) + } + + ctx, cancel := context.WithCancel(ctx) + defer cancel() + + defer context.AfterFunc(ctx, func() { _ = sc.Close() })() + + var failure error for { - bytes, err := buf.ReadBytes('\n') + conn, err := listener.Accept() if err != nil { - return + if !errors.Is(err, net.ErrClosed) { + failure = fmt.Errorf("updatehub: accept on %q: %w", sc.socketPath, err) + } + + break } - sc.emit(c, strings.Trim(string(bytes), "\n")) - c.Close() + sc.handlers.Go(func() { + sc.handle(ctx, conn) + }) } + + sc.handlers.Wait() + + return failure } -func (sc *StateChange) emit(c net.Conn, state string) { - for _, f := range sc.Listeners[strings.Join([]string{state}, "")] { - f(&Handler{conn: c}) +// Close stops Serve and releases the socket. It is safe to call more than once, +// and safe to call on a listener that was never bound. A closed listener can be +// bound again. +func (sc *StateChange) Close() error { + sc.mu.Lock() + defer sc.mu.Unlock() + + listener := sc.listener + if listener == nil { + return nil + } + sc.listener = nil + sc.closed = true + + if err := listener.Close(); err != nil { + return fmt.Errorf("updatehub: close %q: %w", sc.socketPath, err) } + + return nil } -func createListener() (net.Listener, error) { - socketEnv := os.Getenv("UH_LISTENER_TEST") - if len(socketEnv) == 0 { - removeFile(SocketPath) +// handle reads one state name and answers it. The agent's trigger script writes +// the name with a newline after it, then reads until the connection closes. +func (sc *StateChange) handle(ctx context.Context, conn net.Conn) { + var state State + + defer func() { _ = conn.Close() }() + + // A peer that connects and writes nothing would park this goroutine, and + // Serve waits for it, so the read gives up when Serve does. + defer context.AfterFunc(ctx, func() { _ = conn.SetReadDeadline(time.Now()) })() + + // A callback runs in a goroutine this package owns, so a panic in one would + // take the whole process down and the consumer could not recover it. It + // becomes an error on the sink instead. Any deferred answer the callback + // registered has already run by now. + defer func() { + if raised := recover(); raised != nil { + sc.report(fmt.Errorf("updatehub: the %q callback panicked: %v", state, raised)) + } + }() - ln, err := net.Listen("unix", SocketPath) - return ln, err + line, err := bufio.NewReaderSize(conn, readBufferSize).ReadString('\n') + if err != nil && !errors.Is(err, io.EOF) { + sc.report(fmt.Errorf("updatehub: read a state name from %q: %w", sc.socketPath, err)) + + return + } + + state = State(strings.TrimSpace(line)) + if state == "" { + return + } + + callback := sc.callbackFor(state) + if callback == nil { + return } - removeFile(socketEnv) - ln, err := net.Listen("unix", socketEnv) - return ln, err + handler := &Handler{conn: conn, state: state} + + switch err := callback(ctx, handler); { + case err != nil: + sc.report(err) + case handler.writeErr != nil: + sc.report(handler.writeErr) + } } -func removeFile(file string) { - _, err := os.Stat(file) - if err == nil && !os.IsNotExist(err) { - err := os.Remove(file) - checkErr(err) +func (sc *StateChange) callbackFor(state State) CallbackFunc { + sc.mu.Lock() + defer sc.mu.Unlock() + + return sc.callbacks[state] +} + +func (sc *StateChange) report(err error) { + sc.mu.Lock() + sink := sc.onError + sc.mu.Unlock() + + if sink != nil { + sink(err) } } -func checkErr(err error) { - if err != nil { - log.Fatal(err) +// TriggerInstalled reports whether the agent's trigger script is at path. Pass +// TriggerPath for the ordinary case. +// +// Without the script the agent never consults the socket, so it installs +// whatever the server offers, on every poll, unasked. What to do about that is +// the caller's decision. +func TriggerInstalled(path string) (bool, error) { + if _, err := os.Stat(path); err != nil { + if errors.Is(err, os.ErrNotExist) { + return false, nil + } + + return false, fmt.Errorf("updatehub: stat %q: %w", path, err) } + + return true, nil } diff --git a/listener_test.go b/listener_test.go new file mode 100644 index 0000000..c9a3775 --- /dev/null +++ b/listener_test.go @@ -0,0 +1,560 @@ +/* +UpdateHub +Copyright (C) 2019 +O.S. Systems Sofware LTDA: contato@ossystems.com.br + +SPDX-License-Identifier: Apache-2.0 +*/ + +package updatehub + +import ( + "context" + "errors" + "io" + "net" + "os" + "path/filepath" + "strings" + "sync" + "syscall" + "testing" + "time" +) + +// socketPath returns a short path under the test's own directory. A Unix socket +// address is bounded to about 100 bytes, and the default temporary directory of +// a test can be longer than that. +func socketPath(t *testing.T) string { + t.Helper() + + dir, err := os.MkdirTemp("", "uh") + if err != nil { + t.Fatalf("MkdirTemp: %v", err) + } + t.Cleanup(func() { _ = os.RemoveAll(dir) }) + + return filepath.Join(dir, "s.sock") +} + +// serve binds the listener and serves it until the test ends. The returned call +// stops it earlier, and waits for Serve to return. +func serve(t *testing.T, sc *StateChange) (stop func()) { + t.Helper() + + if err := sc.Bind(); err != nil { + t.Fatalf("Bind: %v", err) + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + + go func() { done <- sc.Serve(ctx) }() + + stop = sync.OnceFunc(func() { + cancel() + select { + case err := <-done: + if err != nil { + t.Errorf("Serve: %v", err) + } + case <-time.After(5 * time.Second): + t.Error("Serve did not return after the context was cancelled") + } + _ = sc.Close() + }) + t.Cleanup(stop) + + return stop +} + +// ask plays the agent's trigger script: it writes one state name and reads +// whatever the listener answers before the connection closes. +func ask(t *testing.T, path, state string) string { + t.Helper() + + conn, err := net.Dial("unix", path) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = conn.Close() }() + + if err := conn.SetDeadline(time.Now().Add(5 * time.Second)); err != nil { + t.Fatalf("SetDeadline: %v", err) + } + if _, err := conn.Write([]byte(state + "\n")); err != nil { + t.Fatalf("Write: %v", err) + } + + reply, err := io.ReadAll(conn) + if err != nil { + t.Fatalf("ReadAll: %v", err) + } + + return string(reply) +} + +func TestTheListenerCancelsTheStateItWasAskedTo(t *testing.T) { + path := socketPath(t) + sc := NewStateChange(path) + sc.OnState(StateDownload, func(_ context.Context, h *Handler) error { return h.Cancel() }) + sc.OnState(StateInstall, func(_ context.Context, h *Handler) error { return h.Proceed() }) + serve(t, sc) + + if got := ask(t, path, "download"); got != "cancel" { + t.Errorf("download answered %q, want cancel", got) + } + if got := ask(t, path, "install"); got != "" { + t.Errorf("install answered %q, want nothing", got) + } +} + +func TestTheHandlerReportsTheStateItWasCalledFor(t *testing.T) { + path := socketPath(t) + seen := make(chan State, 1) + + sc := NewStateChange(path) + sc.OnState(StateReboot, func(_ context.Context, h *Handler) error { + seen <- h.State() + return h.Proceed() + }) + serve(t, sc) + + ask(t, path, "reboot") + + select { + case got := <-seen: + if got != StateReboot { + t.Errorf("State() = %q, want reboot", got) + } + case <-time.After(5 * time.Second): + t.Fatal("the callback did not run") + } +} + +func TestAStateWithNoCallbackIsAnsweredByClosing(t *testing.T) { + path := socketPath(t) + sc := NewStateChange(path) + serve(t, sc) + + if got := ask(t, path, "probe"); got != "" { + t.Errorf("probe answered %q, want nothing", got) + } +} + +func TestABareNewlineIsAnsweredByClosing(t *testing.T) { + path := socketPath(t) + called := make(chan struct{}, 1) + + sc := NewStateChange(path) + sc.OnState("", func(_ context.Context, h *Handler) error { + called <- struct{}{} + return h.Proceed() + }) + serve(t, sc) + + if got := ask(t, path, ""); got != "" { + t.Errorf("a bare newline answered %q, want nothing", got) + } + + select { + case <-called: + t.Error("a bare newline reached a callback, want no dispatch") + case <-time.After(100 * time.Millisecond): + } +} + +func TestTheListenerServesConnectionsConcurrently(t *testing.T) { + path := socketPath(t) + release := make(chan struct{}) + // The callback blocks, and Serve waits for it, so the release must happen + // before the cleanup that stops Serve. + defer close(release) + + sc := NewStateChange(path) + sc.OnState(StateDownload, func(_ context.Context, h *Handler) error { + <-release + return h.Cancel() + }) + sc.OnState(StateInstall, func(_ context.Context, h *Handler) error { return h.Proceed() }) + serve(t, sc) + + blocked, err := net.Dial("unix", path) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = blocked.Close() }() + if _, err := blocked.Write([]byte("download\n")); err != nil { + t.Fatalf("Write: %v", err) + } + + answered := make(chan string, 1) + go func() { answered <- ask(t, path, "install") }() + + select { + case got := <-answered: + if got != "" { + t.Errorf("install answered %q, want nothing", got) + } + case <-time.After(5 * time.Second): + t.Fatal("a blocked callback stopped the listener from serving a second connection") + } +} + +func TestBindReturnsEADDRINUSEAndKeepsTheLiveListenerServing(t *testing.T) { + path := socketPath(t) + + first := NewStateChange(path) + first.OnState(StateDownload, func(_ context.Context, h *Handler) error { return h.Cancel() }) + serve(t, first) + + second := NewStateChange(path) + err := second.Bind() + if err == nil { + _ = second.Close() + t.Fatal("the second Bind succeeded, want EADDRINUSE") + } + if !errors.Is(err, syscall.EADDRINUSE) { + t.Fatalf("Bind error = %v, want EADDRINUSE", err) + } + + if got := ask(t, path, "download"); got != "cancel" { + t.Errorf("download answered %q after a failed second bind, want cancel", got) + } +} + +func TestBindDoesNotUnlinkAStaleSocket(t *testing.T) { + path := socketPath(t) + + stale, err := net.Listen("unix", path) + if err != nil { + t.Fatalf("Listen: %v", err) + } + if unix, ok := stale.(*net.UnixListener); ok { + unix.SetUnlinkOnClose(false) + } + if err := stale.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + sc := NewStateChange(path) + if err := sc.Bind(); err == nil { + _ = sc.Close() + t.Fatal("Bind removed a socket that was already there, want EADDRINUSE") + } + + if _, err := os.Stat(path); err != nil { + t.Errorf("the socket file is gone: %v", err) + } +} + +func TestServeWithoutBindReturnsAnError(t *testing.T) { + sc := NewStateChange(socketPath(t)) + + if err := sc.Serve(context.Background()); err == nil { + t.Fatal("Serve returned no error with nothing bound") + } +} + +func TestBindTwiceReturnsAnError(t *testing.T) { + sc := NewStateChange(socketPath(t)) + if err := sc.Bind(); err != nil { + t.Fatalf("Bind: %v", err) + } + defer func() { _ = sc.Close() }() + + if err := sc.Bind(); err == nil { + t.Fatal("the second Bind on the same listener returned no error") + } +} + +func TestCloseStopsServe(t *testing.T) { + sc := NewStateChange(socketPath(t)) + if err := sc.Bind(); err != nil { + t.Fatalf("Bind: %v", err) + } + + done := make(chan error, 1) + go func() { done <- sc.Serve(context.Background()) }() + + // Serve may not have reached Accept yet; Close is safe either way. + if err := sc.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + select { + case err := <-done: + if err != nil { + t.Errorf("Serve returned %v after Close, want nil", err) + } + case <-time.After(5 * time.Second): + t.Fatal("Serve did not return after Close") + } + + if err := sc.Close(); err != nil { + t.Errorf("the second Close returned %v, want nil", err) + } +} + +func TestServeAfterCloseStopsWithoutFaulting(t *testing.T) { + sc := NewStateChange(socketPath(t)) + if err := sc.Bind(); err != nil { + t.Fatalf("Bind: %v", err) + } + if err := sc.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + // A Close that wins the race against a starting Serve is an ordinary stop, + // not a fault. Only a listener that was never bound is an error. + if err := sc.Serve(context.Background()); err != nil { + t.Fatalf("Serve returned %v on a closed listener, want nil", err) + } +} + +func TestACloseIsFollowedByABind(t *testing.T) { + path := socketPath(t) + + sc := NewStateChange(path) + if err := sc.Bind(); err != nil { + t.Fatalf("Bind: %v", err) + } + if err := sc.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + sc.OnState(StateDownload, func(_ context.Context, h *Handler) error { return h.Cancel() }) + serve(t, sc) + + if got := ask(t, path, "download"); got != "cancel" { + t.Errorf("download answered %q after a rebind, want cancel", got) + } +} + +func TestACallbackErrorReachesTheSinkAndTheListenerKeepsServing(t *testing.T) { + path := socketPath(t) + failure := errors.New("the participant is not ready") + reported := make(chan error, 4) + + sc := NewStateChange(path) + sc.OnError(func(err error) { reported <- err }) + sc.OnState(StateDownload, func(_ context.Context, _ *Handler) error { return failure }) + sc.OnState(StateInstall, func(_ context.Context, h *Handler) error { return h.Proceed() }) + serve(t, sc) + + ask(t, path, "download") + + select { + case err := <-reported: + if !errors.Is(err, failure) { + t.Errorf("the sink got %v, want %v", err, failure) + } + case <-time.After(5 * time.Second): + t.Fatal("the callback error never reached the sink") + } + + if got := ask(t, path, "install"); got != "" { + t.Errorf("install answered %q after a callback error, want nothing", got) + } +} + +func TestACallbackErrorWithNoSinkDoesNotKillTheProcess(t *testing.T) { + path := socketPath(t) + + sc := NewStateChange(path) + sc.OnState(StateDownload, func(_ context.Context, _ *Handler) error { + return errors.New("the participant is not ready") + }) + serve(t, sc) + + if got := ask(t, path, "download"); got != "" { + t.Errorf("download answered %q, want nothing", got) + } +} + +func TestACallbackPanicReachesTheSinkAndTheListenerKeepsServing(t *testing.T) { + path := socketPath(t) + reported := make(chan error, 4) + + sc := NewStateChange(path) + sc.OnError(func(err error) { reported <- err }) + sc.OnState(StateDownload, func(_ context.Context, _ *Handler) error { + panic("the participant lost its bridge") + }) + sc.OnState(StateInstall, func(_ context.Context, h *Handler) error { return h.Proceed() }) + serve(t, sc) + + ask(t, path, "download") + + select { + case err := <-reported: + if !strings.Contains(err.Error(), "the participant lost its bridge") { + t.Errorf("the sink got %v, want the panic value", err) + } + case <-time.After(5 * time.Second): + t.Fatal("the panic never reached the sink") + } + + if got := ask(t, path, "install"); got != "" { + t.Errorf("install answered %q after a callback panic, want nothing", got) + } +} + +func TestAPanickingCallbackStillAnswersWhatItDeferred(t *testing.T) { + path := socketPath(t) + + sc := NewStateChange(path) + sc.OnState(StateDownload, func(_ context.Context, h *Handler) error { + defer h.Cancel() //nolint:errcheck // the deferred answer is the point + + panic("the participant lost its bridge") + }) + serve(t, sc) + + if got := ask(t, path, "download"); got != "cancel" { + t.Errorf("download answered %q, want the answer the callback deferred", got) + } +} + +func TestAFailedVetoReachesTheSinkEvenWhenTheCallbackDropsIt(t *testing.T) { + path := socketPath(t) + reported := make(chan error, 1) + hungUp := make(chan struct{}) + + sc := NewStateChange(path) + sc.OnError(func(err error) { reported <- err }) + sc.OnState(StateDownload, func(_ context.Context, h *Handler) error { + <-hungUp + + // The peer is gone, so the veto cannot be written. The callback drops + // the error on purpose: the listener must report it anyway. + _ = h.Cancel() + + return nil + }) + serve(t, sc) + + conn, err := net.Dial("unix", path) + if err != nil { + t.Fatalf("Dial: %v", err) + } + if _, err := conn.Write([]byte("download\n")); err != nil { + t.Fatalf("Write: %v", err) + } + // The state name is already in the socket buffer, so the listener still + // reads it and still runs the callback. + if err := conn.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + close(hungUp) + + select { + case err := <-reported: + if err == nil { + t.Error("the sink got a nil error") + } + case <-time.After(5 * time.Second): + t.Fatal("the failed veto never reached the sink") + } + + if got := ask(t, path, "download"); got != "cancel" { + t.Errorf("download answered %q after a lost peer, want cancel", got) + } +} + +func TestReplyingTwiceReturnsAnError(t *testing.T) { + path := socketPath(t) + second := make(chan error, 1) + + sc := NewStateChange(path) + sc.OnState(StateDownload, func(_ context.Context, h *Handler) error { + if err := h.Cancel(); err != nil { + return err + } + second <- h.Proceed() + + return nil + }) + serve(t, sc) + + if got := ask(t, path, "download"); got != "cancel" { + t.Errorf("download answered %q, want cancel", got) + } + + select { + case err := <-second: + if !errors.Is(err, ErrAlreadyReplied) { + t.Errorf("the second reply returned %v, want ErrAlreadyReplied", err) + } + case <-time.After(5 * time.Second): + t.Fatal("the callback did not run") + } +} + +func TestTheCallbackSeesTheServeContext(t *testing.T) { + path := socketPath(t) + seen := make(chan context.Context, 1) + + sc := NewStateChange(path) + sc.OnState(StateDownload, func(ctx context.Context, h *Handler) error { + seen <- ctx + return h.Proceed() + }) + stop := serve(t, sc) + + ask(t, path, "download") + + var callbackCtx context.Context + select { + case callbackCtx = <-seen: + case <-time.After(5 * time.Second): + t.Fatal("the callback did not run") + } + + if callbackCtx.Err() != nil { + t.Errorf("the callback context was already done: %v", callbackCtx.Err()) + } + + stop() + + if callbackCtx.Err() == nil { + t.Error("the callback context outlived Serve") + } +} + +func TestTriggerInstalledReportsPresence(t *testing.T) { + dir := t.TempDir() + + missing := filepath.Join(dir, "absent") + installed, err := TriggerInstalled(missing) + if err != nil { + t.Fatalf("TriggerInstalled: %v", err) + } + if installed { + t.Error("a missing trigger reported as installed") + } + + present := filepath.Join(dir, "present") + if err := os.WriteFile(present, []byte("#!/bin/sh\n"), 0o755); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + installed, err = TriggerInstalled(present) + if err != nil { + t.Fatalf("TriggerInstalled: %v", err) + } + if !installed { + t.Error("an installed trigger reported as missing") + } +} + +func TestTheTriggerAndSocketPathsAreTheAgentsOwn(t *testing.T) { + const want = "/usr/share/updatehub/state-change-callbacks.d/10-updatehub-sdk-statechange-trigger" + + if TriggerPath != want { + t.Errorf("TriggerPath = %q, want %q", TriggerPath, want) + } + if DefaultSocketPath != "/run/updatehub-statechange.sock" { + t.Errorf("DefaultSocketPath = %q", DefaultSocketPath) + } +}