Files

401 lines
10 KiB
Go
Raw Permalink Normal View History

//go:build linux || freebsd
// +build linux freebsd
package ftp
import (
"errors"
"net/url"
"strings"
"testing"
"time"
)
// ---------- Protocol tests ----------
func TestNewProtocol(t *testing.T) {
p := NewProtocol()
if p == nil {
t.Fatal("NewProtocol() returned nil")
}
}
func TestScheme(t *testing.T) {
p := &Protocol{}
if got := p.Scheme(); got != "ftp" {
t.Errorf("Scheme() = %q, want %q", got, "ftp")
}
}
func TestCanHandle(t *testing.T) {
tests := []struct {
name string
raw string
want bool
}{
{"ftp scheme", "ftp://ftp.example.com/file.txt", true},
{"ftps scheme", "ftps://ftp.example.com/file.txt", true},
{"http scheme", "http://example.com/file.txt", false},
{"https scheme", "https://example.com/file.txt", false},
{"sftp scheme", "sftp://example.com/file.txt", false},
{"empty scheme", "file.txt", false},
{"nil URL", "", false},
}
p := &Protocol{}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var u *url.URL
if tt.raw != "" {
var err error
u, err = url.Parse(tt.raw)
if err != nil {
t.Fatalf("failed to parse URL %q: %v", tt.raw, err)
}
}
if got := p.CanHandle(u); got != tt.want {
t.Errorf("CanHandle(%q) = %v, want %v", tt.raw, got, tt.want)
}
})
}
}
func TestCanHandleCaseInsensitive(t *testing.T) {
p := &Protocol{}
u, _ := url.Parse("FTP://example.com/file.txt")
if !p.CanHandle(u) {
t.Errorf("CanHandle should be case-insensitive for FTP")
}
u, _ = url.Parse("FTPS://example.com/file.txt")
if !p.CanHandle(u) {
t.Errorf("CanHandle should be case-insensitive for FTPS")
}
}
func TestCapabilities(t *testing.T) {
p := &Protocol{}
caps := p.Capabilities()
if caps == nil {
t.Fatal("Capabilities() returned nil")
}
expected := map[string]bool{
"download": true,
"resume": true,
"recursive": true,
"tls": true,
}
if len(caps) != len(expected) {
t.Errorf("Capabilities() length = %d, want %d", len(caps), len(expected))
}
for _, cap := range caps {
if !expected[cap] {
t.Errorf("unexpected capability: %q", cap)
}
}
}
func TestSupportsResume(t *testing.T) {
p := &Protocol{}
if !p.SupportsResume() {
t.Errorf("SupportsResume() should return true")
}
}
func TestSupportsCompression(t *testing.T) {
p := &Protocol{}
if p.SupportsCompression() {
t.Errorf("SupportsCompression() should return false")
}
}
func TestSupportsParallel(t *testing.T) {
p := &Protocol{}
if p.SupportsParallel() {
t.Errorf("SupportsParallel() should return false")
}
}
func TestSupportsRecursive(t *testing.T) {
p := &Protocol{}
if !p.SupportsRecursive() {
t.Errorf("SupportsRecursive() should return true")
}
}
func TestReadAll(t *testing.T) {
t.Run("normal content", func(t *testing.T) {
input := "Hello, FTP world!"
r := strings.NewReader(input)
data, err := readAll(r)
if err != nil {
t.Fatalf("readAll() returned error: %v", err)
}
if string(data) != input {
t.Errorf("readAll() = %q, want %q", string(data), input)
}
})
t.Run("empty reader", func(t *testing.T) {
r := strings.NewReader("")
data, err := readAll(r)
if err != nil {
t.Fatalf("readAll() returned error: %v", err)
}
if len(data) != 0 {
t.Errorf("readAll() returned %d bytes, want 0", len(data))
}
})
t.Run("large content", func(t *testing.T) {
// Create content larger than the 32KB buffer to ensure multiple reads
content := strings.Repeat("A", 100*1024)
r := strings.NewReader(content)
data, err := readAll(r)
if err != nil {
t.Fatalf("readAll() returned error: %v", err)
}
if len(data) != len(content) {
t.Errorf("readAll() returned %d bytes, want %d", len(data), len(content))
}
})
t.Run("reader error", func(t *testing.T) {
expectedErr := errors.New("read error")
r := &errorReader{err: expectedErr}
_, err := readAll(r)
if err == nil {
t.Fatal("readAll() should return error")
}
if !strings.Contains(err.Error(), expectedErr.Error()) {
t.Errorf("readAll() error = %v, want %v", err, expectedErr)
}
})
}
// errorReader implements io.Reader that always returns an error
type errorReader struct {
err error
}
func (r *errorReader) Read(p []byte) (n int, err error) {
return 0, r.err
}
// ---------- Client tests ----------
func TestNewClient(t *testing.T) {
tests := []struct {
name string
rawURL string
wantErr bool
}{
{"basic FTP URL", "ftp://ftp.example.com/pub/file.txt", false},
{"FTP URL with path", "ftp://ftp.example.com/pub/", false},
{"FTP URL root", "ftp://ftp.example.com", false},
{"FTP URL with trailing slash", "ftp://ftp.example.com/", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
client, err := NewClient(tt.rawURL, 30*time.Second)
if tt.wantErr {
if err == nil {
t.Errorf("NewClient(%q) expected error", tt.rawURL)
}
return
}
if err != nil {
t.Fatalf("NewClient(%q) returned error: %v", tt.rawURL, err)
}
if client == nil {
t.Fatal("NewClient() returned nil")
}
})
}
}
func TestNewClientWithAuth(t *testing.T) {
client, err := NewClient("ftp://user:password@ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.user != "user" {
t.Errorf("client.user = %q, want %q", client.user, "user")
}
if client.password != "password" {
t.Errorf("client.password = %q, want %q", client.password, "password")
}
}
func TestNewClientWithUserOnly(t *testing.T) {
client, err := NewClient("ftp://user@ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.user != "user" {
t.Errorf("client.user = %q, want %q", client.user, "user")
}
if client.password != "anonymous@" {
t.Errorf("client.password should default to %q when only username is provided, got %q", "anonymous@", client.password)
}
}
func TestNewClientFTPS(t *testing.T) {
client, err := NewClient("ftps://ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if !client.useTLS {
t.Errorf("client.useTLS should be true for ftps:// scheme")
}
}
func TestNewClientFTPSWithAuth(t *testing.T) {
client, err := NewClient("ftps://alice:secret@ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if !client.useTLS {
t.Errorf("client.useTLS should be true for ftps:// scheme")
}
if client.user != "alice" {
t.Errorf("client.user = %q, want %q", client.user, "alice")
}
if client.password != "secret" {
t.Errorf("client.password = %q, want %q", client.password, "secret")
}
}
func TestNewClientInvalidURL(t *testing.T) {
// Only inputs where url.Parse itself fails
_, err := NewClient("://invalid", 30*time.Second)
if err == nil {
t.Errorf("NewClient(%q) expected error", "://invalid")
}
}
func TestNewClientEmptyURL(t *testing.T) {
client, err := NewClient("", 30*time.Second)
if err != nil {
t.Fatalf("NewClient(\"\") returned error: %v", err)
}
// url.Parse("") returns an empty URL struct, so host will be empty
if client.host != "" {
t.Errorf("client.host = %q, want empty", client.host)
}
if client.port != 21 {
t.Errorf("client.port = %d, want %d", client.port, 21)
}
}
func TestNewClientNoScheme(t *testing.T) {
client, err := NewClient("not-a-url", 30*time.Second)
if err != nil {
t.Fatalf("NewClient(\"not-a-url\") returned error: %v", err)
}
// url.Parse("not-a-url") parses it as a path with empty host
if client.host != "" {
t.Errorf("client.host = %q, want empty", client.host)
}
if client.port != 21 {
t.Errorf("client.port = %d, want %d", client.port, 21)
}
}
func TestNewClientDefaultPort(t *testing.T) {
client, err := NewClient("ftp://ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.port != 21 {
t.Errorf("client.port = %d, want %d", client.port, 21)
}
}
func TestNewClientCustomPort(t *testing.T) {
client, err := NewClient("ftp://ftp.example.com:2121/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.port != 2121 {
t.Errorf("client.port = %d, want %d", client.port, 2121)
}
}
func TestNewClientInvalidPort(t *testing.T) {
_, err := NewClient("ftp://ftp.example.com:invalid/file.txt", 30*time.Second)
if err == nil {
t.Errorf("NewClient() expected error for invalid port")
}
}
func TestNewClientDefaultCredentials(t *testing.T) {
client, err := NewClient("ftp://ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.user != "anonymous" {
t.Errorf("default user = %q, want %q", client.user, "anonymous")
}
if client.password != "anonymous@" {
t.Errorf("default password = %q, want %q", client.password, "anonymous@")
}
}
func TestNewClientHostParsing(t *testing.T) {
client, err := NewClient("ftp://ftp.example.com/path/to/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.host != "ftp.example.com" {
t.Errorf("client.host = %q, want %q", client.host, "ftp.example.com")
}
}
func TestNewClientPassiveModeDefault(t *testing.T) {
client, err := NewClient("ftp://ftp.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if !client.passive {
t.Errorf("client.passive should default to true")
}
}
func TestNewClientTLSConfig(t *testing.T) {
client, err := NewClient("ftps://secure.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
if client.tlsConfig == nil {
t.Fatal("client.tlsConfig should not be nil")
}
if client.tlsConfig.ServerName != "secure.example.com" {
t.Errorf("tlsConfig.ServerName = %q, want %q", client.tlsConfig.ServerName, "secure.example.com")
}
if client.tlsConfig.InsecureSkipVerify {
t.Errorf("tlsConfig.InsecureSkipVerify should be false by default")
}
}
func TestSetPinnedCert(t *testing.T) {
client, err := NewClient("ftps://secure.example.com/file.txt", 30*time.Second)
if err != nil {
t.Fatalf("NewClient() returned error: %v", err)
}
// Pinning is off by default.
if client.pinnedCertHash != "" {
t.Errorf("pinnedCertHash should default to empty, got %q", client.pinnedCertHash)
}
// SetPinnedCert normalises to lowercase and trims whitespace.
client.SetPinnedCert(" ABCDEF0123456789 ")
if client.pinnedCertHash != "abcdef0123456789" {
t.Errorf("pinnedCertHash = %q, want %q", client.pinnedCertHash, "abcdef0123456789")
}
}