From 35b786277c8513ded0381251f17bfce1b175a17b Mon Sep 17 00:00:00 2001 From: Mateusz Filipowicz Date: Mon, 10 Feb 2025 18:43:20 +0100 Subject: [PATCH] feat: use sysinfo API for getting system information with fallback to old API (#10) * feat: use sysinfo API for getting system information with fallback to old API * chore: apply linting --- unifi/sysinfo.go | 21 ++- unifi/unifi.go | 51 ++++-- unifi/unifi_test.go | 377 ++++++++++++++++++++------------------------ 3 files changed, 228 insertions(+), 221 deletions(-) diff --git a/unifi/sysinfo.go b/unifi/sysinfo.go index c217e8a..d83c2ae 100644 --- a/unifi/sysinfo.go +++ b/unifi/sysinfo.go @@ -5,12 +5,17 @@ import ( "fmt" ) -type sysInfo struct { //nolint: unused - Timezone string `json:"timezone"` - Version string `json:"version"` - PreviousVersion string `json:"previous_version"` - UBNTDeviceType string `json:"ubnt_device_type"` - UDMVersion string `json:"udm_version"` +type SysInfo struct { + Timezone string `json:"timezone"` + Version string `json:"version"` + PreviousVersion string `json:"previous_version"` + Build string `json:"build"` + Name string `json:"name"` + Hostname string `json:"hostname"` + IPAddrs []string `json:"ip_addrs"` + Uptime int64 `json:"uptime"` + UBNTDeviceType string `json:"ubnt_device_type"` + UDMVersion string `json:"udm_version"` /* @@ -66,10 +71,10 @@ type sysInfo struct { //nolint: unused */ } -func (c *Client) sysinfo(ctx context.Context, id string) (*sysInfo, error) { //nolint: unused +func (c *Client) GetSystemInfo(ctx context.Context, id string) (*SysInfo, error) { var respBody struct { Meta Meta `json:"Meta"` - Data []sysInfo `json:"data"` + Data []SysInfo `json:"data"` } err := c.Get(ctx, fmt.Sprintf("s/%s/stat/sysinfo", id), nil, &respBody) diff --git a/unifi/unifi.go b/unifi/unifi.go index 4a4f12d..4be1da3 100644 --- a/unifi/unifi.go +++ b/unifi/unifi.go @@ -68,7 +68,7 @@ type ClientConfig struct { type Client struct { BaseURL *url.URL - ServerInfo *ServerInfo + SysInfo *SysInfo apiPaths *ApiPaths config *ClientConfig http *http.Client @@ -180,13 +180,16 @@ func NewClient(config *ClientConfig) (*Client, error) { if err != nil { return nil, fmt.Errorf("failed creating validator: %w", err) } + if err := v.Validate(config); err != nil { return nil, fmt.Errorf("failed validating config: %w", err) } + u, err := newUnifi(config, v) if err != nil { return nil, fmt.Errorf("failed creating unifi client: %w", err) } + if err = u.determineApiStyle(); err != nil { return u, fmt.Errorf("failed determining API style: %w", err) } @@ -195,10 +198,10 @@ func NewClient(config *ClientConfig) (*Client, error) { return u, fmt.Errorf("failed logging in: %w", err) } - if serverInfo, err := u.GetServerInfo(); err != nil { + if sysInfo, err := u.getSystemInformation(); err != nil { return u, fmt.Errorf("failed getting server info: %w", err) } else { - u.ServerInfo = serverInfo + u.SysInfo = sysInfo } return u, nil } @@ -388,12 +391,7 @@ func (c *Client) determineApiStyle() error { return nil } -// GetServerInfo reads the controller's version and UUID. Only call this if you -// previously called Login and suspect the controller version has changed. -func (c *Client) GetServerInfo() (*ServerInfo, error) { - ctx, cancel := c.createRequestContext() - defer cancel() - +func (c *Client) getOldSysInfo(ctx context.Context) (*SysInfo, error) { var response struct { Data ServerInfo `json:"Meta"` } @@ -402,8 +400,41 @@ func (c *Client) GetServerInfo() (*ServerInfo, error) { if err != nil { return nil, err } + data := response.Data + return &SysInfo{ + Version: data.ServerVersion, + }, nil +} - return &response.Data, nil +// getSystemInformation reads the controller's version and UUID. Only call this if you +// previously called Login and suspect the controller version has changed. +func (c *Client) getSystemInformation() (*SysInfo, error) { + ctx, cancel := c.createRequestContext() + defer cancel() + + var resultingError error + info, err := c.GetSystemInfo(ctx, "default") // get for default site which must exist + if err != nil { + resultingError = err + } else if info == nil || info.Version == "" { + resultingError = errors.New("new API returned empty server info") + } + + if resultingError != nil { + info, err = c.getOldSysInfo(ctx) + if err != nil { + resultingError = errors.Join(resultingError, err) + } else if info == nil || info.Version == "" { + resultingError = errors.Join(resultingError, errors.New("old API returned empty server info")) + } else { + resultingError = nil + } + } + + if resultingError != nil { + return nil, resultingError + } + return info, nil } func marshalRequest(reqBody interface{}) (io.Reader, error) { diff --git a/unifi/unifi_test.go b/unifi/unifi_test.go index 80e6cbc..ef43489 100644 --- a/unifi/unifi_test.go +++ b/unifi/unifi_test.go @@ -5,11 +5,9 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "net/http/httptest" "reflect" - "slices" "strings" "testing" @@ -22,41 +20,26 @@ const ( testUrl = "http://test.url" ) -func verifyContainsInterceptors(a *assert.Assertions, c *Client, interceptors ...interface{}) { - var ( - expectedTypes = make([]reflect.Type, len(interceptors)) - matchingTypes = make([]reflect.Type, len(interceptors)) - ) +// verifyInterceptorPresence checks each expected interceptor type for presence or absence in the client. +func verifyInterceptorPresence(a *assert.Assertions, c *Client, interceptors []interface{}, shouldExist bool) { + expectedTypes := make([]reflect.Type, 0, len(interceptors)) for _, i := range interceptors { expectedTypes = append(expectedTypes, reflect.TypeOf(i)) } - for _, i := range c.interceptors { - actualType := reflect.TypeOf(i) - if slices.Contains(expectedTypes, actualType) { - matchingTypes = append(matchingTypes, actualType) + for _, et := range expectedTypes { + found := false + for _, actual := range c.interceptors { + if reflect.TypeOf(actual) == et { + found = true + break + } } - } - if len(matchingTypes) != len(expectedTypes) { - a.Fail(fmt.Sprintf("interceptors not found; expected: %v, found: %v", expectedTypes, matchingTypes)) - } -} - -func verifyDoesNotContainInterceptors(a *assert.Assertions, c *Client, interceptors ...interface{}) { - var ( - expectedTypes = make([]reflect.Type, 0, len(interceptors)) - matchingTypes = make([]reflect.Type, 0, len(interceptors)) - ) - for _, i := range interceptors { - expectedTypes = append(expectedTypes, reflect.TypeOf(i)) - } - for _, i := range c.interceptors { - actualType := reflect.TypeOf(i) - if slices.Contains(expectedTypes, actualType) { - matchingTypes = append(matchingTypes, actualType) + if shouldExist && !found { + a.Fail(fmt.Sprintf("expected interceptor %v not found", et)) + } + if !shouldExist && found { + a.Fail(fmt.Sprintf("unexpected interceptor %v found", et)) } - } - if len(matchingTypes) != 0 { - a.Fail(fmt.Sprintf("interceptors found; expected to be not present: %v, found: %v", expectedTypes, matchingTypes)) } } @@ -72,8 +55,8 @@ func TestNewClient(t *testing.T) { require.Error(t, err) a.EqualValues(localUrl, c.BaseURL.String()) a.Contains(err.Error(), "connection refused", "an invalid destination should produce a connection error.") - verifyContainsInterceptors(a, c, &CsrfInterceptor{}, &DefaultHeadersInterceptor{}) - verifyDoesNotContainInterceptors(a, c, &ApiKeyAuthInterceptor{}) + verifyInterceptorPresence(a, c, []interface{}{&CsrfInterceptor{}, &DefaultHeadersInterceptor{}}, true) + verifyInterceptorPresence(a, c, []interface{}{&ApiKeyAuthInterceptor{}}, false) } func TestNewClientWithApiKey(t *testing.T) { @@ -90,8 +73,8 @@ func TestNewClientWithApiKey(t *testing.T) { require.Error(t, err) a.EqualValues(localUrl, c.BaseURL.String()) a.Contains(err.Error(), "connection refused", "an invalid destination should produce a connection error.") - verifyContainsInterceptors(a, c, &ApiKeyAuthInterceptor{}, &DefaultHeadersInterceptor{}) - verifyDoesNotContainInterceptors(a, c, &CsrfInterceptor{}) + verifyInterceptorPresence(a, c, []interface{}{&ApiKeyAuthInterceptor{}, &DefaultHeadersInterceptor{}}, true) + verifyInterceptorPresence(a, c, []interface{}{&CsrfInterceptor{}}, false) } func TestCustomizeHttpClient(t *testing.T) { @@ -173,34 +156,54 @@ func NewTestClientWithInterceptor() (*Client, *TestInterceptor) { return c, interceptor } -func TestInterceptors(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given +// runClientGetRequest creates a new test client, performs a GET request, +// asserts that an error occurred, and returns the client and its interceptor. +func runClientGetRequest(t *testing.T, path string, data interface{}) (*Client, *TestInterceptor) { + t.Helper() c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Get(context.Background(), "/", nil, nil) - - // then + err := c.Get(context.Background(), path, data, nil) require.Error(t, err) - a.True(interceptor.IsRequestIntercepted(), "request interceptor not called") - a.False(interceptor.IsResponseIntercepted(), "response interceptor called, but should not because of failed request") + return c, interceptor } -func TestNoSendRequestWhenRequestInterceptorReturnsError(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given +// runClientRequest creates a new test client, performs a request with the given method, +// asserts that an error occurred, and returns the client and its interceptor. +func runClientRequest(t *testing.T, method, path string, body interface{}) (*Client, *TestInterceptor) { + t.Helper() c, interceptor := NewTestClientWithInterceptor() - interceptor.failOnRequest = true - - // when - err := c.Get(context.Background(), "/", nil, nil) - - // then + err := c.Do(context.Background(), method, path, body, nil) require.Error(t, err) - a.Contains(err.Error(), "request interceptor failed") + return c, interceptor +} + +// TestRequestInterceptorBehavior tests the interceptor's behavior in both normal and failing scenarios. +func TestRequestInterceptorBehavior(t *testing.T) { + t.Parallel() + tests := []struct { + name string + failOnRequest bool + expectedErrorSubstring string + expectRequest bool + expectResponse bool + }{ + {"Normal interceptor", false, "", true, false}, + {"Failing interceptor", true, "request interceptor failed", true, false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + c, interceptor := NewTestClientWithInterceptor() + interceptor.failOnRequest = tc.failOnRequest + err := c.Get(context.Background(), "/", nil, nil) + require.Error(t, err) + if tc.expectedErrorSubstring != "" { + require.ErrorContains(t, err, tc.expectedErrorSubstring) + } + assert.Equal(t, tc.expectRequest, interceptor.IsRequestIntercepted()) + assert.Equal(t, tc.expectResponse, interceptor.IsResponseIntercepted()) + }) + } } func TestProperRequestUrl(t *testing.T) { @@ -222,47 +225,34 @@ func TestProperRequestUrl(t *testing.T) { for _, tc := range testCases { t.Run(tc.path, func(t *testing.T) { t.Parallel() - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Get(context.Background(), tc.path, nil, nil) - - // then - require.Error(t, err) + // Use the helper to perform a GET request and capture the interceptor. + _, interceptor := runClientGetRequest(t, tc.path, nil) a.EqualValues(tc.expected, interceptor.request.URL.String()) }) } } -func TestApiKeyAddedToRequest(t *testing.T) { +func TestRequestHeaders(t *testing.T) { t.Parallel() - a := assert.New(t) - // given - c, interceptor := NewTestClientWithInterceptor() + tests := []struct { + name string + header string + expected string + }{ + {"API Key Header", ApiKeyHeader, "test-key"}, + {"Accept Header", AcceptHeader, "application/json"}, + {"Content-Type Header", ContentTypeHeader, "application/json; charset=utf-8"}, + {"User-Agent Header", UserAgentHeader, defaultUserAgent}, + } - // when - err := c.Get(context.Background(), "/", nil, nil) + _, interceptor := runClientGetRequest(t, "/", nil) - // then - require.Error(t, err) - a.EqualValues("test-key", interceptor.RequestHeader(ApiKeyHeader)) -} - -func TestDefaultHeadersAddedToRequest(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Get(context.Background(), "/", nil, nil) - - // then - require.Error(t, err) - a.EqualValues("application/json", interceptor.RequestHeader(AcceptHeader)) - a.EqualValues("application/json; charset=utf-8", interceptor.RequestHeader(ContentTypeHeader)) - a.EqualValues(defaultUserAgent, interceptor.RequestHeader(UserAgentHeader)) + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + assert.EqualValues(t, tc.expected, interceptor.RequestHeader(tc.header)) + }) + } } type TestData struct { @@ -273,20 +263,12 @@ func TestRequestSentWithJson(t *testing.T) { t.Parallel() a := assert.New(t) // given - c, interceptor := NewTestClientWithInterceptor() - data := &TestData{ - Data: "test", - } - - // when - err := c.Get(context.Background(), "/", data, nil) - - // then - require.Error(t, err) - body := &TestData{} - err = json.NewDecoder(interceptor.request.Body).Decode(body) + data := &TestData{Data: "test"} + _, interceptor := runClientGetRequest(t, "/", data) + var body TestData + err := json.NewDecoder(interceptor.request.Body).Decode(&body) require.NoError(t, err) - a.Equal(data, body) + a.Equal(data, &body) } func TestRequestMethod(t *testing.T) { @@ -299,108 +281,24 @@ func TestRequestMethod(t *testing.T) { for _, tc := range testCases { t.Run(tc, func(t *testing.T) { t.Parallel() - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Do(context.Background(), tc, "", nil, nil) - - // then - require.Error(t, err) + _, interceptor := runClientRequest(t, tc, "", nil) a.EqualValues(tc, interceptor.Method()) }) } } -func TestGetRequest(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Get(context.Background(), "/", nil, nil) - - // then - require.Error(t, err) - a.EqualValues(http.MethodGet, interceptor.Method()) -} - -func TestPostRequest(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Post(context.Background(), "/", nil, nil) - - // then - require.Error(t, err) - a.EqualValues(http.MethodPost, interceptor.Method()) -} - -func TestPutRequest(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Put(context.Background(), "/", nil, nil) - - // then - require.Error(t, err) - a.EqualValues(http.MethodPut, interceptor.Method()) -} - -func TestDeleteRequest(t *testing.T) { - t.Parallel() - a := assert.New(t) - // given - c, interceptor := NewTestClientWithInterceptor() - - // when - err := c.Delete(context.Background(), "/", nil, nil) - - // then - require.Error(t, err) - a.EqualValues(http.MethodDelete, interceptor.Method()) -} - -func RunTestServer(path string, requestBody interface{}) *httptest.Server { +func runTestServer(path string) *httptest.Server { return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Add(CsrfHeader, "csrf-token") + // Always set the CSRF header on the response. + w.Header().Set(CsrfHeader, "csrf-token") if !strings.EqualFold(r.URL.Path, path) { - w.WriteHeader(http.StatusNotFound) + http.NotFound(w, r) return } + + // Return a JSON response w.WriteHeader(http.StatusOK) - data, err := io.ReadAll(r.Body) - if err != nil { - w.WriteHeader(http.StatusInternalServerError) - fmt.Printf("error reading body:%v", err) - return - } - err = json.Unmarshal(data, &requestBody) - if err != nil { - w.WriteHeader(http.StatusInternalServerError) - fmt.Printf("error decoding body: %s: %s", string(data), err) - return - } - resp := TestData{ - Data: "test", - } - respData, err := json.Marshal(resp) - if err != nil { - w.WriteHeader(http.StatusInternalServerError) - fmt.Printf("error encoding response: %s", err) - return - } - _, err = w.Write(respData) - if err != nil { - fmt.Printf("error writing response: %s", err) - } + _ = json.NewEncoder(w).Encode(TestData{Data: "test"}) })) } @@ -408,11 +306,7 @@ func TestUnifiIntegrationUserPassInjected(t *testing.T) { t.Parallel() a := assert.New(t) // given - type userPass struct { - Username string `json:"username"` - Password string `json:"password"` - } - srv := RunTestServer(NewStyleAPI.LoginPath, userPass{}) + srv := runTestServer(NewStyleAPI.LoginPath) interceptor := NewTestInterceptor() c, _ := NewClient(&ClientConfig{ URL: srv.URL, @@ -438,7 +332,7 @@ func TestResponseDataHandling(t *testing.T) { reqData := TestData{ Data: "request", } - srv := RunTestServer(NewStyleAPI.ApiPath+"/test", TestData{}) + srv := runTestServer(NewStyleAPI.ApiPath + "/test") c, _ := NewClient(&ClientConfig{ URL: srv.URL, APIKey: "test-key", @@ -458,7 +352,7 @@ func TestCsrfHandling(t *testing.T) { t.Parallel() a := assert.New(t) // given - srv := RunTestServer("", struct{}{}) + srv := runTestServer("") interceptor := NewTestInterceptor() c, _ := NewClient(&ClientConfig{ URL: srv.URL, @@ -622,3 +516,80 @@ func TestValidationModes(t *testing.T) { }) } } + +// Common test server setup for system information tests. +type sysInfoTestCase struct { + name string + newAPIVersion string + oldAPIVersion string + expectedError string + expectedResult string +} + +func setupSysInfoTestServer(tc sysInfoTestCase) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "", "/": + w.WriteHeader(http.StatusOK) + case "/proxy/network/api/s/default/stat/sysinfo": + w.WriteHeader(http.StatusOK) + fmt.Fprintf(w, `{"data": [{"version": "%s"}]}`, tc.newAPIVersion) + case "/proxy/network/status": + w.WriteHeader(http.StatusOK) + fmt.Fprintf(w, `{"Meta": {"server_version": "%s"}}`, tc.oldAPIVersion) + default: + http.NotFound(w, r) + } + })) +} + +func TestGetSystemInformation(t *testing.T) { + t.Parallel() + + testCases := []sysInfoTestCase{ + { + name: "New API Success", + newAPIVersion: "v2-success", + oldAPIVersion: "", + expectedResult: "v2-success", + }, + { + name: "Fallback to Old API", + newAPIVersion: "", + oldAPIVersion: "old-success", + expectedResult: "old-success", + }, + { + name: "Both APIs Failure", + newAPIVersion: "", + oldAPIVersion: "", + expectedError: "new API returned empty server info", + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + a := assert.New(t) + + ts := setupSysInfoTestServer(tc) + defer ts.Close() + + c, _ := NewClient(&ClientConfig{ + URL: ts.URL, + APIKey: "dummy", + VerifySSL: false, + }) + + sysInfo, err := c.getSystemInformation() + + if tc.expectedError != "" { + require.ErrorContains(t, err, tc.expectedError) + a.Nil(sysInfo) + } else { + require.NoError(t, err) + a.Equal(tc.expectedResult, sysInfo.Version) + } + }) + } +}