@@ -0,0 +1,41 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDBFromEnv(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
driver string
|
||||
dsn string
|
||||
wantDriver string
|
||||
wantErr string
|
||||
}{
|
||||
{name: "缺省走 sqlite", wantDriver: "sqlite"},
|
||||
{name: "postgres 带 DSN", driver: "postgres", dsn: "host=x", wantDriver: "postgres"},
|
||||
{name: "mysql 缺 DSN 报错", driver: "mysql", wantErr: "DB_DSN is required"},
|
||||
{name: "postgres 缺 DSN 报错", driver: "postgres", wantErr: "DB_DSN is required"},
|
||||
{name: "非法驱动报错", driver: "oracle", wantErr: "unsupported DB_DRIVER"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Setenv("DB_DRIVER", tt.driver)
|
||||
t.Setenv("DB_DSN", tt.dsn)
|
||||
driver, dsn, err := dbFromEnv()
|
||||
if tt.wantErr != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("err = %v, want contains %q", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("dbFromEnv: %v", err)
|
||||
}
|
||||
if driver != tt.wantDriver || dsn != tt.dsn {
|
||||
t.Fatalf("got (%q,%q), want (%q,%q)", driver, dsn, tt.wantDriver, tt.dsn)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user