1055 lines
28 KiB
Go
1055 lines
28 KiB
Go
//go:build linux || freebsd
|
||
// +build linux freebsd
|
||
|
||
package http
|
||
|
||
import (
|
||
"context"
|
||
"net/http"
|
||
"net/url"
|
||
"testing"
|
||
"time"
|
||
|
||
"codeberg.org/petrbalvin/goget/internal/core"
|
||
)
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// Helpers
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func makeResponse(headers map[string]string, statusCode int) *http.Response {
|
||
resp := &http.Response{
|
||
StatusCode: statusCode,
|
||
Header: make(http.Header),
|
||
}
|
||
for k, v := range headers {
|
||
resp.Header.Set(k, v)
|
||
}
|
||
resp.ContentLength = 1000
|
||
return resp
|
||
}
|
||
|
||
func mustParse(raw string) *url.URL {
|
||
u, err := url.Parse(raw)
|
||
if err != nil {
|
||
panic(err)
|
||
}
|
||
return u
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – ParseHTTPURL
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestParseHTTPURL(t *testing.T) {
|
||
tests := []struct {
|
||
name string
|
||
raw string
|
||
want *URLInfo
|
||
wantErr bool
|
||
}{
|
||
{
|
||
name: "simple http",
|
||
raw: "http://example.com/path",
|
||
want: &URLInfo{
|
||
Scheme: "http",
|
||
Host: "example.com",
|
||
Port: "80",
|
||
Path: "/path",
|
||
},
|
||
},
|
||
{
|
||
name: "simple https",
|
||
raw: "https://example.com/path",
|
||
want: &URLInfo{
|
||
Scheme: "https",
|
||
Host: "example.com",
|
||
Port: "443",
|
||
Path: "/path",
|
||
},
|
||
},
|
||
{
|
||
name: "with explicit port",
|
||
raw: "http://example.com:8080/path",
|
||
want: &URLInfo{
|
||
Scheme: "http",
|
||
Host: "example.com",
|
||
Port: "8080",
|
||
Path: "/path",
|
||
},
|
||
},
|
||
{
|
||
name: "with query",
|
||
raw: "https://example.com/path?key=val&a=b",
|
||
want: &URLInfo{
|
||
Scheme: "https",
|
||
Host: "example.com",
|
||
Port: "443",
|
||
Path: "/path",
|
||
Query: "key=val&a=b",
|
||
},
|
||
},
|
||
{
|
||
name: "with fragment",
|
||
raw: "http://example.com/path#section",
|
||
want: &URLInfo{
|
||
Scheme: "http",
|
||
Host: "example.com",
|
||
Port: "80",
|
||
Path: "/path",
|
||
Fragment: "section",
|
||
},
|
||
},
|
||
{
|
||
name: "with userinfo",
|
||
raw: "http://user:pass@example.com/path",
|
||
want: &URLInfo{
|
||
Scheme: "http",
|
||
Host: "example.com",
|
||
Port: "80",
|
||
Path: "/path",
|
||
User: url.UserPassword("user", "pass"),
|
||
},
|
||
},
|
||
{
|
||
name: "root path",
|
||
raw: "https://example.com",
|
||
want: &URLInfo{
|
||
Scheme: "https",
|
||
Host: "example.com",
|
||
Port: "443",
|
||
Path: "",
|
||
},
|
||
},
|
||
{
|
||
name: "user only no password",
|
||
raw: "http://alice@example.com/",
|
||
want: &URLInfo{
|
||
Scheme: "http",
|
||
Host: "example.com",
|
||
Port: "80",
|
||
Path: "/",
|
||
User: url.User("alice"),
|
||
},
|
||
},
|
||
{
|
||
name: "ipv6 host",
|
||
raw: "http://[::1]:9090/api",
|
||
want: &URLInfo{
|
||
Scheme: "http",
|
||
Host: "::1",
|
||
Port: "9090",
|
||
Path: "/api",
|
||
},
|
||
},
|
||
{
|
||
name: "https ipv6 default port",
|
||
raw: "https://[::1]/secure",
|
||
want: &URLInfo{
|
||
Scheme: "https",
|
||
Host: "::1",
|
||
Port: "443",
|
||
Path: "/secure",
|
||
},
|
||
},
|
||
}
|
||
|
||
for _, tc := range tests {
|
||
t.Run(tc.name, func(t *testing.T) {
|
||
u := mustParse(tc.raw)
|
||
got, err := ParseHTTPURL(u)
|
||
if tc.wantErr {
|
||
if err == nil {
|
||
t.Fatal("expected error, got nil")
|
||
}
|
||
return
|
||
}
|
||
if err != nil {
|
||
t.Fatalf("unexpected error: %v", err)
|
||
}
|
||
if got.Scheme != tc.want.Scheme {
|
||
t.Errorf("Scheme = %q, want %q", got.Scheme, tc.want.Scheme)
|
||
}
|
||
if got.Host != tc.want.Host {
|
||
t.Errorf("Host = %q, want %q", got.Host, tc.want.Host)
|
||
}
|
||
if got.Port != tc.want.Port {
|
||
t.Errorf("Port = %q, want %q", got.Port, tc.want.Port)
|
||
}
|
||
if got.Path != tc.want.Path {
|
||
t.Errorf("Path = %q, want %q", got.Path, tc.want.Path)
|
||
}
|
||
if got.Query != tc.want.Query {
|
||
t.Errorf("Query = %q, want %q", got.Query, tc.want.Query)
|
||
}
|
||
if got.Fragment != tc.want.Fragment {
|
||
t.Errorf("Fragment = %q, want %q", got.Fragment, tc.want.Fragment)
|
||
}
|
||
|
||
// Check userinfo
|
||
if tc.want.User != nil {
|
||
if got.User == nil {
|
||
t.Fatal("expected User to be non-nil")
|
||
}
|
||
gu := got.User.Username()
|
||
gp, _ := got.User.Password()
|
||
wu := tc.want.User.Username()
|
||
wp, _ := tc.want.User.Password()
|
||
if gu != wu || gp != wp {
|
||
t.Errorf("Userinfo = (%q,%q), want (%q,%q)", gu, gp, wu, wp)
|
||
}
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
func TestParseHTTPURLInvalidScheme(t *testing.T) {
|
||
invalidSchemes := []string{"ftp", "file", "data", "ws", "wss"}
|
||
for _, scheme := range invalidSchemes {
|
||
t.Run(scheme, func(t *testing.T) {
|
||
u := mustParse(scheme + "://example.com")
|
||
_, err := ParseHTTPURL(u)
|
||
if err == nil {
|
||
t.Errorf("expected error for scheme %q, got nil", scheme)
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – header extraction helpers
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestGetContentType(t *testing.T) {
|
||
resp := makeResponse(map[string]string{"Content-Type": "text/html; charset=utf-8"}, 200)
|
||
if got := GetContentType(resp); got != "text/html; charset=utf-8" {
|
||
t.Errorf("got %q, want %q", got, "text/html; charset=utf-8")
|
||
}
|
||
|
||
resp2 := makeResponse(nil, 200)
|
||
if got := GetContentType(resp2); got != "" {
|
||
t.Errorf("expected empty, got %q", got)
|
||
}
|
||
}
|
||
|
||
func TestGetContentLength(t *testing.T) {
|
||
resp := makeResponse(nil, 200)
|
||
resp.ContentLength = 42
|
||
if got := GetContentLength(resp); got != 42 {
|
||
t.Errorf("got %d, want %d", got, 42)
|
||
}
|
||
|
||
resp.ContentLength = -1
|
||
if got := GetContentLength(resp); got != -1 {
|
||
t.Errorf("got %d, want %d", got, -1)
|
||
}
|
||
}
|
||
|
||
func TestGetContentEncoding(t *testing.T) {
|
||
resp := makeResponse(map[string]string{"Content-Encoding": "gzip"}, 200)
|
||
if got := GetContentEncoding(resp); got != "gzip" {
|
||
t.Errorf("got %q, want %q", got, "gzip")
|
||
}
|
||
|
||
resp2 := makeResponse(nil, 200)
|
||
if got := GetContentEncoding(resp2); got != "" {
|
||
t.Errorf("expected empty, got %q", got)
|
||
}
|
||
}
|
||
|
||
func TestGetETag(t *testing.T) {
|
||
resp := makeResponse(map[string]string{"ETag": `"abc123"`}, 200)
|
||
if got := GetETag(resp); got != `"abc123"` {
|
||
t.Errorf("got %q, want %q", got, `"abc123"`)
|
||
}
|
||
|
||
resp2 := makeResponse(nil, 200)
|
||
if got := GetETag(resp2); got != "" {
|
||
t.Errorf("expected empty, got %q", got)
|
||
}
|
||
}
|
||
|
||
func TestGetLastModified(t *testing.T) {
|
||
resp := makeResponse(map[string]string{"Last-Modified": "Mon, 02 Jan 2006 15:04:05 GMT"}, 200)
|
||
if got := GetLastModified(resp); got != "Mon, 02 Jan 2006 15:04:05 GMT" {
|
||
t.Errorf("got %q, want %q", got, "Mon, 02 Jan 2006 15:04:05 GMT")
|
||
}
|
||
|
||
resp2 := makeResponse(nil, 200)
|
||
if got := GetLastModified(resp2); got != "" {
|
||
t.Errorf("expected empty, got %q", got)
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – SupportsRangeCheck
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestSupportsRangeCheck(t *testing.T) {
|
||
tests := []struct {
|
||
name string
|
||
header string
|
||
want bool
|
||
}{
|
||
{"bytes", "bytes", true},
|
||
{"Bytes (case insensitive)", "Bytes", true},
|
||
{"BYTES", "BYTES", true},
|
||
{"none", "", false},
|
||
{"other value", "none", false},
|
||
}
|
||
|
||
for _, tc := range tests {
|
||
t.Run(tc.name, func(t *testing.T) {
|
||
resp := makeResponse(nil, 200)
|
||
if tc.header != "" {
|
||
resp.Header.Set("Accept-Ranges", tc.header)
|
||
}
|
||
if got := SupportsRangeCheck(resp); got != tc.want {
|
||
t.Errorf("got %v, want %v", got, tc.want)
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – IsRedirect / GetRedirectLocation
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestIsRedirect(t *testing.T) {
|
||
codes := []int{300, 301, 302, 303, 304, 305, 306, 307, 308}
|
||
for _, code := range codes {
|
||
t.Run(http.StatusText(code), func(t *testing.T) {
|
||
if !IsRedirect(code) {
|
||
t.Errorf("IsRedirect(%d) = false, want true", code)
|
||
}
|
||
})
|
||
}
|
||
|
||
nonRedirect := []int{200, 201, 400, 404, 500}
|
||
for _, code := range nonRedirect {
|
||
t.Run(http.StatusText(code), func(t *testing.T) {
|
||
if IsRedirect(code) {
|
||
t.Errorf("IsRedirect(%d) = true, want false", code)
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
func TestGetRedirectLocation(t *testing.T) {
|
||
resp := makeResponse(map[string]string{"Location": "http://example.com/new"}, 302)
|
||
if got := GetRedirectLocation(resp); got != "http://example.com/new" {
|
||
t.Errorf("got %q, want %q", got, "http://example.com/new")
|
||
}
|
||
|
||
resp2 := makeResponse(nil, 200)
|
||
if got := GetRedirectLocation(resp2); got != "" {
|
||
t.Errorf("expected empty, got %q", got)
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – BuildURLWithQuery
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestBuildURLWithQuery(t *testing.T) {
|
||
t.Run("adds query params to bare URL", func(t *testing.T) {
|
||
base := mustParse("http://example.com/resource")
|
||
result, err := BuildURLWithQuery(base, map[string]string{"key": "val", "foo": "bar"})
|
||
if err != nil {
|
||
t.Fatalf("unexpected error: %v", err)
|
||
}
|
||
q := result.Query()
|
||
if q.Get("key") != "val" {
|
||
t.Errorf("key = %q, want %q", q.Get("key"), "val")
|
||
}
|
||
if q.Get("foo") != "bar" {
|
||
t.Errorf("foo = %q, want %q", q.Get("foo"), "bar")
|
||
}
|
||
})
|
||
|
||
t.Run("preserves existing query params", func(t *testing.T) {
|
||
base := mustParse("http://example.com/resource?existing=1")
|
||
result, err := BuildURLWithQuery(base, map[string]string{"new": "2"})
|
||
if err != nil {
|
||
t.Fatalf("unexpected error: %v", err)
|
||
}
|
||
q := result.Query()
|
||
if q.Get("existing") != "1" {
|
||
t.Errorf("existing = %q, want %q", q.Get("existing"), "1")
|
||
}
|
||
if q.Get("new") != "2" {
|
||
t.Errorf("new = %q, want %q", q.Get("new"), "2")
|
||
}
|
||
})
|
||
|
||
t.Run("overwrites existing param", func(t *testing.T) {
|
||
base := mustParse("http://example.com/resource?key=old")
|
||
result, err := BuildURLWithQuery(base, map[string]string{"key": "new"})
|
||
if err != nil {
|
||
t.Fatalf("unexpected error: %v", err)
|
||
}
|
||
if result.Query().Get("key") != "new" {
|
||
t.Errorf("got %q, want %q", result.Query().Get("key"), "new")
|
||
}
|
||
})
|
||
|
||
t.Run("nil base returns error", func(t *testing.T) {
|
||
_, err := BuildURLWithQuery(nil, map[string]string{"k": "v"})
|
||
if err == nil {
|
||
t.Fatal("expected error for nil base URL")
|
||
}
|
||
})
|
||
|
||
t.Run("empty params leaves URL unchanged", func(t *testing.T) {
|
||
base := mustParse("http://example.com/resource")
|
||
result, err := BuildURLWithQuery(base, nil)
|
||
if err != nil {
|
||
t.Fatalf("unexpected error: %v", err)
|
||
}
|
||
if result.String() != base.String() {
|
||
t.Errorf("got %q, want %q", result.String(), base.String())
|
||
}
|
||
})
|
||
|
||
t.Run("special characters are encoded", func(t *testing.T) {
|
||
base := mustParse("http://example.com/search")
|
||
result, err := BuildURLWithQuery(base, map[string]string{"q": "hello world"})
|
||
if err != nil {
|
||
t.Fatalf("unexpected error: %v", err)
|
||
}
|
||
if result.Query().Get("q") != "hello world" {
|
||
t.Errorf("q = %q, want %q", result.Query().Get("q"), "hello world")
|
||
}
|
||
})
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – GetUserinfo
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestGetUserinfo(t *testing.T) {
|
||
t.Run("no user info", func(t *testing.T) {
|
||
u := mustParse("http://example.com")
|
||
user, pass, ok := GetUserinfo(u)
|
||
if ok {
|
||
t.Errorf("expected ok=false, got true (user=%q, pass=%q)", user, pass)
|
||
}
|
||
})
|
||
|
||
t.Run("user without password", func(t *testing.T) {
|
||
u := mustParse("http://user@example.com")
|
||
user, pass, ok := GetUserinfo(u)
|
||
if !ok {
|
||
t.Fatal("expected ok=true")
|
||
}
|
||
if user != "user" {
|
||
t.Errorf("user = %q, want %q", user, "user")
|
||
}
|
||
if pass != "" {
|
||
t.Errorf("pass = %q, want empty", pass)
|
||
}
|
||
})
|
||
|
||
t.Run("user with password", func(t *testing.T) {
|
||
u := mustParse("http://user:secret@example.com")
|
||
user, pass, ok := GetUserinfo(u)
|
||
if !ok {
|
||
t.Fatal("expected ok=true")
|
||
}
|
||
if user != "user" {
|
||
t.Errorf("user = %q, want %q", user, "user")
|
||
}
|
||
if pass != "secret" {
|
||
t.Errorf("pass = %q, want %q", pass, "secret")
|
||
}
|
||
})
|
||
|
||
t.Run("empty user string in URL", func(t *testing.T) {
|
||
u := mustParse("http://:pass@example.com")
|
||
user, pass, ok := GetUserinfo(u)
|
||
if !ok {
|
||
t.Fatal("expected ok=true")
|
||
}
|
||
if user != "" {
|
||
t.Errorf("user = %q, want empty", user)
|
||
}
|
||
if pass != "pass" {
|
||
t.Errorf("pass = %q, want %q", pass, "pass")
|
||
}
|
||
})
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – IsSecure
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestIsSecure(t *testing.T) {
|
||
if !IsSecure(mustParse("https://example.com")) {
|
||
t.Error("IsSecure(https) = false, want true")
|
||
}
|
||
if IsSecure(mustParse("http://example.com")) {
|
||
t.Error("IsSecure(http) = true, want false")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// handler.go – GetDefaultPortForScheme
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestGetDefaultPortForScheme(t *testing.T) {
|
||
if got := GetDefaultPortForScheme("http"); got != "80" {
|
||
t.Errorf("http -> %q, want %q", got, "80")
|
||
}
|
||
if got := GetDefaultPortForScheme("https"); got != "443" {
|
||
t.Errorf("https -> %q, want %q", got, "443")
|
||
}
|
||
if got := GetDefaultPortForScheme("ftp"); got != "" {
|
||
t.Errorf("ftp -> %q, want empty", got)
|
||
}
|
||
if got := GetDefaultPortForScheme(""); got != "" {
|
||
t.Errorf("empty -> %q, want empty", got)
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// upload.go – Simple upload config / construction
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestUploadSimple(t *testing.T) {
|
||
cfg := DefaultConfig()
|
||
cfg.Transport.Timeout = 5 * time.Second
|
||
cfg.Headers["X-Custom"] = "test-value"
|
||
|
||
client, err := NewClient(cfg)
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
|
||
if client.config.Transport.Timeout != 5*time.Second {
|
||
t.Errorf("Timeout = %v, want %v", client.config.Transport.Timeout, 5*time.Second)
|
||
}
|
||
if client.config.Headers["X-Custom"] != "test-value" {
|
||
t.Errorf("Header X-Custom = %q, want %q", client.config.Headers["X-Custom"], "test-value")
|
||
}
|
||
if client.config.Headers == nil {
|
||
t.Error("config.Headers is nil")
|
||
}
|
||
}
|
||
|
||
func TestUploadSimpleNilConfig(t *testing.T) {
|
||
client, err := NewClient(nil)
|
||
if err != nil {
|
||
t.Fatalf("NewClient(nil): %v", err)
|
||
}
|
||
if client.config == nil {
|
||
t.Fatal("config is nil")
|
||
}
|
||
if !client.config.Transport.FollowRedirects {
|
||
t.Error("expected FollowRedirects to be true")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// upload.go – applyBasicAuth
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestApplyBasicAuth(t *testing.T) {
|
||
req, err := http.NewRequest("PUT", "http://example.com/upload", nil)
|
||
if err != nil {
|
||
t.Fatalf("NewRequest: %v", err)
|
||
}
|
||
|
||
applyBasicAuth(req, "alice", "s3cret")
|
||
|
||
user, pass, ok := req.BasicAuth()
|
||
if !ok {
|
||
t.Fatal("BasicAuth() returned ok=false, want true")
|
||
}
|
||
if user != "alice" {
|
||
t.Errorf("user = %q, want %q", user, "alice")
|
||
}
|
||
if pass != "s3cret" {
|
||
t.Errorf("pass = %q, want %q", pass, "s3cret")
|
||
}
|
||
}
|
||
|
||
func TestApplyBasicAuthEmptyCredentials(t *testing.T) {
|
||
req, err := http.NewRequest("PUT", "http://example.com/upload", nil)
|
||
if err != nil {
|
||
t.Fatalf("NewRequest: %v", err)
|
||
}
|
||
|
||
applyBasicAuth(req, "", "")
|
||
|
||
// SetBasicAuth always sets the Authorization header; verify it's present.
|
||
auth := req.Header.Get("Authorization")
|
||
if auth == "" {
|
||
t.Error("expected Authorization header to be set, got empty")
|
||
}
|
||
if len(auth) < 6 || auth[:6] != "Basic " {
|
||
t.Errorf("expected Basic auth prefix, got %q", auth)
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – DefaultConfig
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestDefaultConfig(t *testing.T) {
|
||
cfg := DefaultConfig()
|
||
|
||
if !cfg.Transport.FollowRedirects {
|
||
t.Error("FollowRedirects should be true")
|
||
}
|
||
if cfg.Transport.MaxRedirects != 10 {
|
||
t.Errorf("MaxRedirects = %d, want %d", cfg.Transport.MaxRedirects, 10)
|
||
}
|
||
if cfg.Transport.Timeout != 30*time.Minute {
|
||
t.Errorf("Timeout = %v, want %v", cfg.Transport.Timeout, 30*time.Minute)
|
||
}
|
||
if cfg.UserAgent == "" {
|
||
t.Error("UserAgent should not be empty")
|
||
}
|
||
if cfg.Transport.BufferSize != 32*1024 {
|
||
t.Errorf("BufferSize = %d, want %d", cfg.Transport.BufferSize, 32*1024)
|
||
}
|
||
if cfg.Headers == nil {
|
||
t.Error("Headers map should be initialised")
|
||
}
|
||
if cfg.Auth.AuthType != "auto" {
|
||
t.Errorf("AuthType = %q, want %q", cfg.Auth.AuthType, "auto")
|
||
}
|
||
if cfg.Recursive.MaxDepth != 3 {
|
||
t.Errorf("MaxDepth = %d, want %d", cfg.Recursive.MaxDepth, 3)
|
||
}
|
||
if cfg.Output.OutputDir != "." {
|
||
t.Errorf("OutputDir = %q, want %q", cfg.Output.OutputDir, ".")
|
||
}
|
||
if !cfg.Recursive.MirrorAssets {
|
||
t.Error("MirrorAssets should be true")
|
||
}
|
||
if !cfg.Recursive.ConvertLinks {
|
||
t.Error("ConvertLinks should be true")
|
||
}
|
||
if !cfg.Recursive.RespectRobots {
|
||
t.Error("RespectRobots should be true")
|
||
}
|
||
if cfg.Parallel == nil {
|
||
t.Fatal("Parallel config should not be nil")
|
||
}
|
||
if cfg.Transport.Retry == nil {
|
||
t.Error("Retry config should not be nil")
|
||
}
|
||
if cfg.Transport.Dialer == nil {
|
||
t.Error("Dialer config should not be nil")
|
||
}
|
||
if cfg.Transport.TLS == nil {
|
||
t.Error("TLS config should not be nil")
|
||
}
|
||
if cfg.Recursive.ExcludePatterns == nil {
|
||
t.Error("ExcludePatterns should be initialised (non-nil)")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – NewClient
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestNewClient(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
if client == nil {
|
||
t.Fatal("client is nil")
|
||
}
|
||
if client.httpClient == nil {
|
||
t.Error("httpClient is nil")
|
||
}
|
||
if client.transport == nil {
|
||
t.Error("transport is nil")
|
||
}
|
||
if client.dialer == nil {
|
||
t.Error("dialer is nil")
|
||
}
|
||
if client.config == nil {
|
||
t.Error("config is nil")
|
||
}
|
||
}
|
||
|
||
func TestNewClientNilConfig(t *testing.T) {
|
||
client, err := NewClient(nil)
|
||
if err != nil {
|
||
t.Fatalf("NewClient(nil): %v", err)
|
||
}
|
||
if client == nil {
|
||
t.Fatal("client is nil")
|
||
}
|
||
if client.config == nil {
|
||
t.Fatal("config is nil after nil input")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – Scheme
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestClientScheme(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
if got := client.Scheme(); got != "http" {
|
||
t.Errorf("Scheme() = %q, want %q", got, "http")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – Capabilities
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestClientCapabilities(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
caps := client.Capabilities()
|
||
if len(caps) == 0 {
|
||
t.Fatal("Capabilities() returned empty slice")
|
||
}
|
||
expected := []string{"download", "resume", "compression", "redirects", "tls", "proxy", "parallel", "auth", "cookies", "recursive", "extract"}
|
||
if len(caps) != len(expected) {
|
||
t.Fatalf("len(caps) = %d, want %d; caps = %v", len(caps), len(expected), caps)
|
||
}
|
||
for i, cap := range caps {
|
||
if cap != expected[i] {
|
||
t.Errorf("caps[%d] = %q, want %q", i, cap, expected[i])
|
||
}
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – Supports helpers
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestClientSupportsHelpers(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
if !client.SupportsResume() {
|
||
t.Error("SupportsResume() = false")
|
||
}
|
||
if !client.SupportsCompression() {
|
||
t.Error("SupportsCompression() = false")
|
||
}
|
||
if !client.SupportsParallel() {
|
||
t.Error("SupportsParallel() = false")
|
||
}
|
||
if !client.SupportsRecursive() {
|
||
t.Error("SupportsRecursive() = false")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – CanHandle
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestCanHandle(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
|
||
tests := []struct {
|
||
raw string
|
||
want bool
|
||
}{
|
||
{"http://example.com", true},
|
||
{"https://example.com", true},
|
||
{"HTTP://example.com", true},
|
||
{"HTTPS://example.com", true},
|
||
{"ftp://example.com", false},
|
||
{"file:///tmp/foo", false},
|
||
{"", false},
|
||
}
|
||
for _, tc := range tests {
|
||
t.Run(tc.raw, func(t *testing.T) {
|
||
var u *url.URL
|
||
if tc.raw != "" {
|
||
u = mustParse(tc.raw)
|
||
}
|
||
if got := client.CanHandle(u); got != tc.want {
|
||
t.Errorf("CanHandle(%q) = %v, want %v", tc.raw, got, tc.want)
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – SetHeader / GetTransport
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestSetHeader(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
client.SetHeader("Authorization", "Bearer token123")
|
||
if client.config.Headers["Authorization"] != "Bearer token123" {
|
||
t.Errorf("header = %q, want %q", client.config.Headers["Authorization"], "Bearer token123")
|
||
}
|
||
}
|
||
|
||
func TestGetTransport(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
tr := client.GetTransport()
|
||
if tr == nil {
|
||
t.Fatal("GetTransport() returned nil")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – createChunks
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestCreateChunks(t *testing.T) {
|
||
// Use a custom ParallelConfig with MinSize=1 so small files are still
|
||
// split into multiple chunks, and MinChunkSize=1 so even tiny chunks
|
||
// are accepted.
|
||
pc := &core.ParallelConfig{
|
||
Connections: 0,
|
||
MinSize: 1,
|
||
MinChunkSize: 1,
|
||
MaxChunkSize: 0, // unlimited
|
||
}
|
||
|
||
t.Run("single chunk when count <= 1", func(t *testing.T) {
|
||
chunks := createChunks(1000, 1, pc)
|
||
if len(chunks) != 1 {
|
||
t.Fatalf("expected 1 chunk, got %d", len(chunks))
|
||
}
|
||
if chunks[0].ID != 0 {
|
||
t.Errorf("chunk ID = %d, want 0", chunks[0].ID)
|
||
}
|
||
if chunks[0].Start != 0 {
|
||
t.Errorf("chunk Start = %d, want 0", chunks[0].Start)
|
||
}
|
||
if chunks[0].End != 999 {
|
||
t.Errorf("chunk End = %d, want 999", chunks[0].End)
|
||
}
|
||
})
|
||
|
||
t.Run("multiple chunks for large file", func(t *testing.T) {
|
||
totalSize := int64(1000)
|
||
count := 4
|
||
chunks := createChunks(totalSize, count, pc)
|
||
if len(chunks) != count {
|
||
t.Fatalf("expected %d chunks, got %d", count, len(chunks))
|
||
}
|
||
|
||
var totalDownloaded int64
|
||
for i, ch := range chunks {
|
||
if ch.ID != i {
|
||
t.Errorf("chunk[%d].ID = %d, want %d", i, ch.ID, i)
|
||
}
|
||
if ch.Start > ch.End {
|
||
t.Errorf("chunk[%d]: Start (%d) > End (%d)", i, ch.Start, ch.End)
|
||
}
|
||
totalDownloaded += ch.End - ch.Start + 1
|
||
}
|
||
if totalDownloaded != totalSize {
|
||
t.Errorf("total downloaded = %d, want %d", totalDownloaded, totalSize)
|
||
}
|
||
|
||
for i := 1; i < len(chunks); i++ {
|
||
if chunks[i].Start != chunks[i-1].End+1 {
|
||
t.Errorf("gap between chunk %d and %d: chunk[%d].End=%d, chunk[%d].Start=%d",
|
||
i-1, i, i-1, chunks[i-1].End, i, chunks[i].Start)
|
||
}
|
||
}
|
||
})
|
||
|
||
t.Run("exact multiples divide evenly", func(t *testing.T) {
|
||
chunks := createChunks(800, 4, pc)
|
||
if len(chunks) != 4 {
|
||
t.Fatalf("expected 4 chunks, got %d", len(chunks))
|
||
}
|
||
expectedEnds := []int64{199, 399, 599, 799}
|
||
for i, ch := range chunks {
|
||
if ch.End != expectedEnds[i] {
|
||
t.Errorf("chunk[%d].End = %d, want %d", i, ch.End, expectedEnds[i])
|
||
}
|
||
}
|
||
})
|
||
|
||
t.Run("last chunk absorbs remainder", func(t *testing.T) {
|
||
totalSize := int64(1000 + 37) // not evenly divisible
|
||
count := 4
|
||
chunks := createChunks(totalSize, count, pc)
|
||
if len(chunks) == 0 {
|
||
t.Fatal("expected at least 1 chunk")
|
||
}
|
||
last := chunks[len(chunks)-1]
|
||
if last.End != totalSize-1 {
|
||
t.Errorf("last chunk End = %d, want %d", last.End, totalSize-1)
|
||
}
|
||
})
|
||
|
||
t.Run("default config returns 1 chunk for <= 100MB", func(t *testing.T) {
|
||
defaultPc := core.DefaultParallelConfig()
|
||
chunks := createChunks(100*1024*1024, 4, defaultPc)
|
||
if len(chunks) != 1 {
|
||
t.Fatalf("expected 1 chunk (totalSize=%d <= MinSize=%d), got %d",
|
||
100*1024*1024, defaultPc.MinSize, len(chunks))
|
||
}
|
||
})
|
||
|
||
t.Run("count adjusted by MaxChunkSize", func(t *testing.T) {
|
||
pc2 := &core.ParallelConfig{
|
||
Connections: 100,
|
||
MinSize: 1,
|
||
MinChunkSize: 1,
|
||
MaxChunkSize: 100, // 100 bytes max per chunk
|
||
}
|
||
totalSize := int64(1000)
|
||
chunks := createChunks(totalSize, 100, pc2)
|
||
|
||
if len(chunks) < 10 {
|
||
t.Errorf("expected at least 10 chunks with MaxChunkSize=100, got %d", len(chunks))
|
||
}
|
||
for i, ch := range chunks {
|
||
chunkSize := ch.End - ch.Start + 1
|
||
if chunkSize > 100 {
|
||
t.Errorf("chunk[%d] size = %d, exceeds MaxChunkSize=%d", i, chunkSize, 100)
|
||
}
|
||
}
|
||
})
|
||
|
||
t.Run("zero total size", func(t *testing.T) {
|
||
chunks := createChunks(0, 4, pc)
|
||
if len(chunks) != 1 {
|
||
t.Fatalf("expected 1 chunk for zero size, got %d", len(chunks))
|
||
}
|
||
if chunks[0].End != -1 {
|
||
t.Errorf("End = %d, want -1 (empty file)", chunks[0].End)
|
||
}
|
||
})
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// client.go – NewHandler
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestNewHandler(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
h := NewHandler(client)
|
||
if h == nil {
|
||
t.Fatal("NewHandler returned nil")
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// UploadRequest construction (unit test without real files)
|
||
// ---------------------------------------------------------------------------
|
||
|
||
func TestUploadRequestDefaults(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
|
||
ctx, cancel := context.WithCancel(context.Background())
|
||
cancel()
|
||
|
||
_, err = client.Upload(ctx, &UploadRequest{
|
||
URL: "http://example.com/upload",
|
||
FilePath: "",
|
||
})
|
||
if err == nil {
|
||
t.Error("expected error for empty file path, got nil")
|
||
}
|
||
if err.Error() != "put upload requires --upload-file" {
|
||
t.Errorf("unexpected error message: %v", err)
|
||
}
|
||
}
|
||
|
||
func TestUploadRequestDefaultsNoData(t *testing.T) {
|
||
client, err := NewClient(DefaultConfig())
|
||
if err != nil {
|
||
t.Fatalf("NewClient: %v", err)
|
||
}
|
||
|
||
ctx, cancel := context.WithCancel(context.Background())
|
||
cancel()
|
||
|
||
// POST with no files, no form data, no file path -> no upload data
|
||
_, err = client.Upload(ctx, &UploadRequest{
|
||
URL: "http://example.com/upload",
|
||
Method: "POST",
|
||
})
|
||
if err == nil {
|
||
t.Error("expected error for no upload data, got nil")
|
||
}
|
||
if err.Error() != "no upload data provided" {
|
||
t.Errorf("unexpected error message: %v", err)
|
||
}
|
||
}
|
||
|
||
func TestParseContentDisposition(t *testing.T) {
|
||
tests := []struct {
|
||
name string
|
||
header string
|
||
expected string
|
||
}{
|
||
{"empty header", "", ""},
|
||
{"simple filename", `attachment; filename="file.zip"`, "file.zip"},
|
||
{"extended filename", `attachment; filename*=UTF-8''file%20.zip`, "file .zip"},
|
||
{"inline disposition", `inline; filename="document.pdf"`, "document.pdf"},
|
||
{"no filename", `attachment`, ""},
|
||
{"with path separators", `attachment; filename="../../etc/passwd"`, "../../etc/passwd"},
|
||
{"rfc 5987 utf-8", `attachment; filename*=UTF-8''t%C3%A9l%C3%A9chargement.zip`, "téléchargement.zip"},
|
||
}
|
||
|
||
for _, tc := range tests {
|
||
t.Run(tc.name, func(t *testing.T) {
|
||
result := ParseContentDisposition(tc.header)
|
||
if result != tc.expected {
|
||
t.Errorf("ParseContentDisposition(%q) = %q, want %q", tc.header, result, tc.expected)
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
func TestSanitizeFilename(t *testing.T) {
|
||
tests := []struct {
|
||
name string
|
||
input string
|
||
expected string
|
||
}{
|
||
{"simple name", "file.zip", "file.zip"},
|
||
{"path traversal", "../../etc/passwd", "passwd"},
|
||
{"empty string", "", ""},
|
||
{"dot only", ".", ""},
|
||
{"dot dot only", "..", ""},
|
||
{"with subdir", "dir/file.txt", "file.txt"},
|
||
{"url encoded", "file%20name.zip", "file%20name.zip"},
|
||
{"root path", "/etc/passwd", "passwd"},
|
||
}
|
||
|
||
for _, tc := range tests {
|
||
t.Run(tc.name, func(t *testing.T) {
|
||
result := SanitizeFilename(tc.input)
|
||
if result != tc.expected {
|
||
t.Errorf("SanitizeFilename(%q) = %q, want %q", tc.input, result, tc.expected)
|
||
}
|
||
})
|
||
}
|
||
}
|