From 1a3f2484a8cd57bec17b937b45455df3aa071b65 Mon Sep 17 00:00:00 2001 From: Otavio Salvador Date: Fri, 28 Aug 2026 13:17:27 -0300 Subject: [PATCH] Rewrite the SDK for v2.0.0 The module path becomes 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 its update channel breaks. Three defects are fatal there. The listener called log.Fatal on every error path, including inside Handler.Cancel, so a transient socket error took the host down. Every client method ended in a single-value type assertion on a value that processRequest sets to nil on error, so an unreachable agent panicked the caller. And processRequest discarded the HTTP response, so no status code was ever read. The first two composed. Probe always sent a body, so an empty custom server sent {"custom_server": ""}, which agent 2.1.6 answers with a 500 - and keeps answering with a 500 until it restarts, because a parked agent never reaches the entry point that resets the address. The 500 body then failed to decode, so Probe panicked rather than returning. What the rewrite does: - Nothing exits the process, and no method panics. Every failure is returned to the caller, which decides what it means. - Every call reads the HTTP status code and reports an unexpected one as a *StatusError carrying the code and the body, so "the agent answered 500" is distinguishable from "the reply did not decode". - Probe sends no body when no custom server is given, as the agent's own Rust and Python SDKs do. - ProbeResponse is a type rather than interface{}, and carries all four replies: updating, no_update, try_again(N), and busy with the agent's own state name. Busy is not a failure, and one of the busy state names is "error", so the two must stay apart. - LocalInstall, RemoteInstall and AbortDownload report the agent's 406 refusal as a result rather than as an error. - Calls to one agent are serialised, through a turn held per base URL, so a second Client cannot defeat it. - Configuration is a constructor parameter. NewClient takes the base URL, a timeout and a hold, NewStateChange takes the socket path, and the UH_LISTENER_TEST environment variable is deleted rather than kept as a fallback: a fallback leaves a process-global seam for anyone who does not pass the parameter. - The listener binds without unlinking, so a path that is already there fails with EADDRINUSE and the caller decides what that means. - 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 rather than one at a time. - A callback takes a context and returns an error. Errors, and panics, reach the sink OnError registers; so does a veto the listener could not write, whatever the callback did with that error. A single connection never stops the listener, because a listener that stops leaves the agent blocked. - TriggerInstalled reports whether the agent's trigger script is there. Without it the agent never consults the socket, so it installs whatever the server offers on every poll, unasked. What to do about that is the consumer's decision, so the package only reports it. - The package writes nothing to stdout and nothing through log. - gorequest is gone, with the goproxy and pkg/errors dependencies it pulled in for two POSTs to localhost. The package needs the standard library alone. Why a context bounds the wait and never the request. A call that ended on its context used to release the turn at once, so the next call could reach an agent that was still working on the first one, and the port's own budgets reach that state: a probe was measured blocking 15.569 s against a 10 s local_install budget. That matters because of how the agent handles a request. Its machine selects between a sleep, a waker and await_communication, and await_communication awaits the whole handler. When the sleep or the waker wins the race, select drops the handler mid-flight, the reply channel is dropped without a send, and the task waiting on it reaches unreachable!("Unexpected response: Err(RecvError)") in states/machine/address.rs. Every request still in flight there widens that window, and the communication channel is bounded at 10, which is where the "ten concurrent requests" measurement comes from. So the round trip now runs to its end on a goroutine that holds the turn, while the caller returns at its deadline; a call that then gives up waiting for its turn reports ErrCallOutstanding. An abandoned request is not itself a hazard for the agent - it drops the reply with responder.send(...).ok()? and carries on. The hazard is the second request. The hold is what guarantees the turn comes back. The agent can stay alive and answer nothing at all, which is exactly what it does after the panic above, and neither its answer nor the death of its process would then release the turn. So the caller states how long the client keeps it after giving up. Long is safer than short, and the value is the caller's, because only the caller knows when to stop assuming its agent is working. The module asked for go 1.15, which predates every language feature this rewrite uses. It now asks for 1.26, the release its intended consumer builds with, and modernize -fix has run over the tree. Both files have tests, where the repository had none on either branch. The four that pin the turn - the serialising, the abandoned call, the hold and the shared turn - each fail against the behaviour they replace. The CI jobs drove an OpenAPI mock and UH_LISTENER_TEST through the example binaries; they run that suite instead, under the race detector, against the version go.mod declares and against the newest release. The README states the agent version the release is validated against, 2.1.6, and lists what v2.0.0 breaks. Claude-Session: https://claude.ai/code/session_01Rk8Pd5KHDPTSifYYmoxRan --- .github/workflows/ci.yml | 69 +-- .github/workflows/golangci-lint.yml | 9 +- README.md | 142 +++++- api.go | 721 +++++++++++++++++++++----- api_test.go | 752 ++++++++++++++++++++++++++++ examples/api/main.go | 77 ++- examples/listener/main.go | 55 +- go.mod | 13 +- go.sum | 30 -- listener.go | 345 ++++++++++--- listener_test.go | 560 +++++++++++++++++++++ 11 files changed, 2442 insertions(+), 331 deletions(-) create mode 100644 api_test.go delete mode 100644 go.sum create mode 100644 listener_test.go 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) + } +}