Files
goget/internal/protocol/http/http_test.go
T

1055 lines
28 KiB
Go
Raw Normal View History

//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)
}
})
}
}