fix: comprehensive bug fixes, hardening, and test coverage

Bug fixes:
- Fix nil-pointer panic in sendMessage SOAP fault handler (used err.Error() on nil)
- Fix missing return after SOAP fault response (caused fall-through to 200 OK)
- Fix timeout=0 shadowing: removed package-level constant, use configured value
- Fix receiveMessage not checking SOAP faults from CIIMS

Codec hardening:
- Replace fragile regex parsing with encoding/xml.Decoder for namespace-agnostic XML
- XML-escape all user inputs (user, pass, event) — not just message body
- Remove magic-number-based string slicing (split, GetErrMsg)

Transport improvements:
- Extract Client type with connection pooling (reuse http.Client across requests)
- Add CIIMSResponse domain type with IsFault(), ErrorMessage(), Messages()
- Add HTTPStatusError and ResponseTooLargeError typed errors
- Replace deprecated ioutil.ReadAll with io.ReadAll
- Add response body size limit (64 MiB)

Security hardening:
- Add MaxBytesReader (1MB request body limit) on both endpoints
- Add URL allowlist via CIIMS_ALLOWED_SERVERS env var
- Add normalizeBaseURL with scheme/host/query validation
- Add count bounds validation (1-1000) on /receive
- Add server-level timeouts (ReadHeaderTimeout, ReadTimeout, IdleTimeout)

Configuration:
- Introduce Config struct to replace package-level globals
- Add loadConfig() with full validation and error propagation
- Add getIntEnv() with positive-value enforcement

Test coverage (75 tests, 35 new):
- Phase 1: 8 internal tests (nil receiver, error types, SOAP+HTTP500, malformed XML)
- Phase 2: 19 pure function tests (normalizeBaseURL, parseAllowedServers, getIntEnv)
- Phase 3: 16 handler tests (send/receive success, SOAP fault, network error, HTTP 500,
  body too large, invalid JSON, missing fields, URL allowlist)
- Phase 4: 8 config/router tests (loadConfig, newServer, newRouter)

Toolchain:
- Upgrade Go 1.13 → 1.22, gin 1.6.3 → 1.10.0, testify 1.5.1 → 1.10.0

Cleanup:
- Remove dead code (unused post() function, commented-out defaults)
- Replace println with structured logging
- Add .gitignore
- Rewrite README with API docs, env vars, security considerations
This commit is contained in:
zhiqiang feng
2026-07-08 16:23:00 +08:00
parent dd3f6769f6
commit ce4c2b088f
12 changed files with 1704 additions and 85 deletions
+145 -44
View File
@@ -5,31 +5,57 @@ import (
"fmt"
"log"
"net/http"
"net/url"
"os"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"gzzn.com/mini/ciimsproxy/internal"
)
const (
servicePrefix string = "/services/ExchangeService"
defaultTimeout = 240
maxMsgLen = 1048576 // 1MB
maxCount = 1000
servicePrefix = "/services/ExchangeService"
defaultTimeout = 240
maxMsgLen = 1048576 // 1MB
maxCount = 1000
readHeaderTimeout = 5 * time.Second
readTimeout = 10 * time.Second
writeTimeoutGrace = 10 * time.Second
idleTimeout = 60 * time.Second
allowedServersEnvVarName = "CIIMS_ALLOWED_SERVERS"
)
type Config struct {
ServerURL string
Listen string
Timeout int
Client *internal.Client
ServerURL string
AllowedServers map[string]string
Listen string
Timeout int
Client *internal.Client
}
type App struct {
config Config
}
type sendRequest struct {
URL string `json:"url"`
User string `json:"user" binding:"required"`
Pass string `json:"pass" binding:"required"`
Event string `json:"event" binding:"required"`
Priority int `json:"priority"`
Val bool `json:"val"`
Msg string `json:"msg" binding:"required"`
}
type receiveRequest struct {
URL string `json:"url"`
User string `json:"user" binding:"required"`
Pass string `json:"pass" binding:"required"`
Count int `json:"count" binding:"required"`
}
func main() {
config, err := loadConfig()
if err != nil {
@@ -38,10 +64,10 @@ func main() {
}
log.Printf("use ciims: %s", config.ServerURL)
r := newRouter(config)
srv := newServer(config)
log.Printf("Starting ciims proxy for %s", config.ServerURL)
log.Printf("timeout: %d", config.Timeout)
if err := r.Run(config.Listen); err != nil {
if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Printf("[ERROR] server failed: %v", err)
os.Exit(1)
}
@@ -49,10 +75,25 @@ func main() {
func loadConfig() (Config, error) {
config := Config{
ServerURL: os.Getenv("CIIMS_SERVER"),
Listen: os.Getenv("PROXY_LISTEN"),
Timeout: defaultTimeout,
Listen: os.Getenv("PROXY_LISTEN"),
Timeout: defaultTimeout,
}
serverURL, err := normalizeBaseURL(os.Getenv("CIIMS_SERVER"))
if err != nil {
return Config{}, fmt.Errorf("invalid CIIMS_SERVER: %w", err)
}
config.ServerURL = serverURL
config.AllowedServers = map[string]string{serverURL: serverURL}
allowedServers, err := parseAllowedServers(os.Getenv(allowedServersEnvVarName))
if err != nil {
return Config{}, err
}
for _, allowed := range allowedServers {
config.AllowedServers[allowed] = allowed
}
t, err := getIntEnv("CIIMS_TIMEOUT")
if err != nil {
return Config{}, err
@@ -70,7 +111,21 @@ func loadConfig() (Config, error) {
return config, nil
}
func newServer(config Config) *http.Server {
return &http.Server{
Addr: config.Listen,
Handler: newRouter(config),
ReadHeaderTimeout: readHeaderTimeout,
ReadTimeout: readTimeout,
WriteTimeout: time.Duration(config.Timeout)*time.Second + writeTimeoutGrace,
IdleTimeout: idleTimeout,
}
}
func newRouter(config Config) *gin.Engine {
if config.AllowedServers == nil && config.ServerURL != "" {
config.AllowedServers = map[string]string{config.ServerURL: config.ServerURL}
}
app := &App{config: config}
r := gin.Default()
r.GET("/ping", func(c *gin.Context) {
@@ -84,32 +139,25 @@ func newRouter(config Config) *gin.Engine {
}
func (a *App) sendMessage(c *gin.Context) {
// Limit request body size
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxMsgLen)
var message struct {
URL string `json:"url"`
User string `json:"user" binding:"required"`
Pass string `json:"pass" binding:"required"`
Event string `json:"event" binding:"required"`
Priority int `json:"priority"`
Val bool `json:"val"`
Msg string `json:"msg" binding:"required"`
}
var message sendRequest
if err := c.ShouldBind(&message); err != nil {
a.handleBindError(c, err)
return
}
ciimsURL, ok := a.resolveTargetURL(c, message.URL)
if !ok {
return
}
msg := internal.CreateSend(message.User, message.Pass,
message.Priority, message.Event, message.Val, message.Msg)
url := a.getURL(message.URL)
resp, err := a.config.Client.Send(url, msg)
resp, err := a.config.Client.Send(c.Request.Context(), ciimsURL, msg)
if err != nil {
a.handleSendError(c, "send", resp, err)
return
}
errMsg := resp.ErrorMessage()
if len(errMsg) > 0 {
if errMsg := resp.ErrorMessage(); errMsg != "" {
log.Printf("[ERROR] SOAP fault: %s", errMsg)
c.JSON(http.StatusInternalServerError, gin.H{"error": errMsg})
return
@@ -118,15 +166,9 @@ func (a *App) sendMessage(c *gin.Context) {
}
func (a *App) receiveMessage(c *gin.Context) {
// Limit request body size
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxMsgLen)
var message struct {
URL string `json:"url"`
User string `json:"user" binding:"required"`
Pass string `json:"pass" binding:"required"`
Count int `json:"count" binding:"required"`
}
var message receiveRequest
if err := c.ShouldBind(&message); err != nil {
a.handleBindError(c, err)
return
@@ -137,9 +179,12 @@ func (a *App) receiveMessage(c *gin.Context) {
return
}
ciimsURL, ok := a.resolveTargetURL(c, message.URL)
if !ok {
return
}
msg := internal.CreateReceive(message.User, message.Pass, message.Count)
url := a.getURL(message.URL)
resp, err := a.config.Client.Send(url, msg)
resp, err := a.config.Client.Send(c.Request.Context(), ciimsURL, msg)
if err != nil {
a.handleSendError(c, "receive", resp, err)
return
@@ -150,17 +195,70 @@ func (a *App) receiveMessage(c *gin.Context) {
return
}
c.JSON(http.StatusOK, gin.H{"msgs": resp.Messages()})
}
func (a *App) getURL(url string) string {
var result string
if len(url) > 0 {
result = url + servicePrefix
} else {
result = a.config.ServerURL + servicePrefix
func (a *App) resolveTargetURL(c *gin.Context, requestedURL string) (string, bool) {
targetBase, err := a.getBaseURL(requestedURL)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return "", false
}
return result
return targetBase + servicePrefix, true
}
func (a *App) getBaseURL(requestedURL string) (string, error) {
if strings.TrimSpace(requestedURL) == "" {
return a.config.ServerURL, nil
}
normalized, err := normalizeBaseURL(requestedURL)
if err != nil {
return "", fmt.Errorf("invalid url: %w", err)
}
if _, ok := a.config.AllowedServers[normalized]; !ok {
return "", fmt.Errorf("url is not allowed")
}
return normalized, nil
}
func normalizeBaseURL(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", fmt.Errorf("must not be empty")
}
u, err := url.Parse(raw)
if err != nil {
return "", err
}
u.Scheme = strings.ToLower(u.Scheme)
u.Host = strings.ToLower(u.Host)
if u.Scheme != "http" && u.Scheme != "https" {
return "", fmt.Errorf("scheme must be http or https")
}
if u.Host == "" {
return "", fmt.Errorf("host is required")
}
if u.RawQuery != "" || u.Fragment != "" {
return "", fmt.Errorf("query and fragment are not allowed")
}
u.Path = strings.TrimRight(u.Path, "/")
u.RawPath = ""
return u.String(), nil
}
func parseAllowedServers(raw string) ([]string, error) {
if strings.TrimSpace(raw) == "" {
return nil, nil
}
parts := strings.Split(raw, ",")
allowed := make([]string, 0, len(parts))
for _, part := range parts {
normalized, err := normalizeBaseURL(part)
if err != nil {
return nil, fmt.Errorf("invalid %s entry %q: %w", allowedServersEnvVarName, part, err)
}
allowed = append(allowed, normalized)
}
return allowed, nil
}
func (a *App) handleBindError(c *gin.Context, err error) {
@@ -202,5 +300,8 @@ func getIntEnv(key string) (int, error) {
if err != nil {
return 0, fmt.Errorf("invalid %s value %q: must be an integer", key, val)
}
if ret <= 0 {
return 0, fmt.Errorf("invalid %s value %q: must be positive", key, val)
}
return ret, nil
}
+396 -3
View File
@@ -7,6 +7,7 @@ import (
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
@@ -38,11 +39,287 @@ func performJSON(r http.Handler, method, path, body string) *httptest.ResponseRe
return w
}
func TestPing(t *testing.T) {
r := testRouter(t, func(w http.ResponseWriter, r *http.Request) {})
// =============================================================================
// Phase 2: Pure Function Tests — normalizeBaseURL
// =============================================================================
func TestNormalizeBaseURL_ValidHTTP(t *testing.T) {
result, err := normalizeBaseURL("http://example.com/path/")
require.NoError(t, err)
assert.Equal(t, "http://example.com/path", result)
}
func TestNormalizeBaseURL_ValidHTTPS(t *testing.T) {
result, err := normalizeBaseURL("https://EXAMPLE.COM:8443")
require.NoError(t, err)
assert.Equal(t, "https://example.com:8443", result)
}
func TestNormalizeBaseURL_Empty(t *testing.T) {
_, err := normalizeBaseURL("")
require.Error(t, err)
assert.Contains(t, err.Error(), "must not be empty")
}
func TestNormalizeBaseURL_WhitespaceOnly(t *testing.T) {
_, err := normalizeBaseURL(" ")
require.Error(t, err)
}
func TestNormalizeBaseURL_InvalidScheme(t *testing.T) {
_, err := normalizeBaseURL("ftp://example.com")
require.Error(t, err)
assert.Contains(t, err.Error(), "scheme must be http or https")
}
func TestNormalizeBaseURL_MissingHost(t *testing.T) {
_, err := normalizeBaseURL("http:///path")
require.Error(t, err)
assert.Contains(t, err.Error(), "host is required")
}
func TestNormalizeBaseURL_QueryNotAllowed(t *testing.T) {
_, err := normalizeBaseURL("http://example.com?a=1")
require.Error(t, err)
assert.Contains(t, err.Error(), "query and fragment")
}
func TestNormalizeBaseURL_FragmentNotAllowed(t *testing.T) {
_, err := normalizeBaseURL("http://example.com#section")
require.Error(t, err)
assert.Contains(t, err.Error(), "query and fragment")
}
func TestNormalizeBaseURL_InvalidURL(t *testing.T) {
_, err := normalizeBaseURL("://bad")
require.Error(t, err)
}
// =============================================================================
// Phase 2: Pure Function Tests — parseAllowedServers
// =============================================================================
func TestParseAllowedServers_Empty(t *testing.T) {
result, err := parseAllowedServers("")
require.NoError(t, err)
assert.Nil(t, result)
}
func TestParseAllowedServers_WhitespaceOnly(t *testing.T) {
_, err := parseAllowedServers(" , ")
require.Error(t, err)
assert.Contains(t, err.Error(), "CIIMS_ALLOWED_SERVERS")
}
func TestParseAllowedServers_ValidSingle(t *testing.T) {
result, err := parseAllowedServers("http://other.example.com")
require.NoError(t, err)
assert.Equal(t, []string{"http://other.example.com"}, result)
}
func TestParseAllowedServers_ValidMultiple(t *testing.T) {
result, err := parseAllowedServers("http://a.com,https://b.com:8443")
require.NoError(t, err)
assert.ElementsMatch(t, []string{"http://a.com", "https://b.com:8443"}, result)
}
func TestParseAllowedServers_InvalidEntry(t *testing.T) {
_, err := parseAllowedServers("http://good.com,ftp://bad.com")
require.Error(t, err)
assert.Contains(t, err.Error(), "CIIMS_ALLOWED_SERVERS")
assert.Contains(t, err.Error(), "ftp://bad.com")
}
// =============================================================================
// Phase 2: Pure Function Tests — getIntEnv
// =============================================================================
func TestGetIntEnv_Unset(t *testing.T) {
t.Setenv("CIIMS_TEST_UNSET_KEY", "")
result, err := getIntEnv("CIIMS_TEST_UNSET_KEY")
require.NoError(t, err)
assert.Equal(t, 0, result)
}
func TestGetIntEnv_Valid(t *testing.T) {
t.Setenv("CIIMS_TEST_VALID_KEY", "30")
result, err := getIntEnv("CIIMS_TEST_VALID_KEY")
require.NoError(t, err)
assert.Equal(t, 30, result)
}
func TestGetIntEnv_NotAnInteger(t *testing.T) {
t.Setenv("CIIMS_TEST_BAD_KEY", "abc")
_, err := getIntEnv("CIIMS_TEST_BAD_KEY")
require.Error(t, err)
assert.Contains(t, err.Error(), "must be an integer")
}
func TestGetIntEnv_Negative(t *testing.T) {
t.Setenv("CIIMS_TEST_NEG_KEY", "-5")
_, err := getIntEnv("CIIMS_TEST_NEG_KEY")
require.Error(t, err)
assert.Contains(t, err.Error(), "must be positive")
}
func TestGetIntEnv_Zero(t *testing.T) {
t.Setenv("CIIMS_TEST_ZERO_KEY", "0")
_, err := getIntEnv("CIIMS_TEST_ZERO_KEY")
require.Error(t, err)
assert.Contains(t, err.Error(), "must be positive")
}
// =============================================================================
// Phase 3: Handler-Level Tests (remaining gaps)
// =============================================================================
func TestSendNetworkError(t *testing.T) {
gin.SetMode(gin.TestMode)
ciims := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
fmt.Fprint(w, testSendOK)
}))
ciims.Close() // close before creating router so the URL is dead
client, err := internal.NewClient(1)
require.NoError(t, err)
r := newRouter(Config{
ServerURL: ciims.URL,
Listen: ":0",
Timeout: 1,
Client: client,
AllowedServers: map[string]string{ciims.URL: ciims.URL},
})
w := performJSON(r, http.MethodPost, "/send", `{"user":"FIMS","pass":"x","event":"E1","msg":"<MSG/>"}`)
assert.Equal(t, http.StatusInternalServerError, w.Code)
assert.Contains(t, w.Body.String(), "connection refused")
}
func TestSendInvalidJSON(t *testing.T) {
gin.SetMode(gin.TestMode)
client, err := internal.NewClient(10)
require.NoError(t, err)
r := newRouter(Config{
ServerURL: "http://127.0.0.1:1",
Listen: ":0",
Timeout: 10,
Client: client,
AllowedServers: map[string]string{},
})
w := performJSON(r, http.MethodPost, "/send", `not json`)
assert.Equal(t, http.StatusBadRequest, w.Code)
}
func TestReceiveCountValidBoundary(t *testing.T) {
gin.SetMode(gin.TestMode)
ciims := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
fmt.Fprint(w, testReceiveResp)
}))
defer ciims.Close()
client, err := internal.NewClient(10)
require.NoError(t, err)
r := newRouter(Config{
ServerURL: ciims.URL,
Listen: ":0",
Timeout: 10,
Client: client,
AllowedServers: map[string]string{ciims.URL: ciims.URL},
})
w := performJSON(r, http.MethodPost, "/receive", `{"user":"FIMS","pass":"x","count":1000}`)
assert.Equal(t, http.StatusOK, w.Code)
}
// =============================================================================
// Phase 4: Config & Router Tests (remaining gaps)
// =============================================================================
func TestLoadConfig_Minimal(t *testing.T) {
t.Setenv("CIIMS_SERVER", "http://example.com")
t.Setenv("PROXY_LISTEN", "")
t.Setenv("CIIMS_TIMEOUT", "")
t.Setenv("CIIMS_ALLOWED_SERVERS", "")
config, err := loadConfig()
require.NoError(t, err)
assert.Equal(t, ":9090", config.Listen)
assert.Equal(t, 240, config.Timeout)
assert.Equal(t, "http://example.com", config.ServerURL)
assert.Contains(t, config.AllowedServers, "http://example.com")
require.NotNil(t, config.Client)
}
func TestLoadConfig_Full(t *testing.T) {
t.Setenv("CIIMS_SERVER", "http://ciims.example.com")
t.Setenv("PROXY_LISTEN", ":8080")
t.Setenv("CIIMS_TIMEOUT", "120")
t.Setenv("CIIMS_ALLOWED_SERVERS", "http://backup1.example.com,https://backup2.example.com:8443")
config, err := loadConfig()
require.NoError(t, err)
assert.Equal(t, ":8080", config.Listen)
assert.Equal(t, 120, config.Timeout)
assert.Equal(t, "http://ciims.example.com", config.ServerURL)
assert.Contains(t, config.AllowedServers, "http://ciims.example.com")
assert.Contains(t, config.AllowedServers, "http://backup1.example.com")
assert.Contains(t, config.AllowedServers, "https://backup2.example.com:8443")
require.NotNil(t, config.Client)
}
func TestLoadConfig_InvalidTimeout(t *testing.T) {
t.Setenv("CIIMS_SERVER", "http://example.com")
t.Setenv("CIIMS_TIMEOUT", "abc")
t.Setenv("CIIMS_ALLOWED_SERVERS", "")
_, err := loadConfig()
require.Error(t, err)
assert.Contains(t, err.Error(), "CIIMS_TIMEOUT")
}
func TestNewServerConfig(t *testing.T) {
client, err := internal.NewClient(30)
require.NoError(t, err)
config := Config{
ServerURL: "http://example.com",
AllowedServers: map[string]string{"http://example.com": "http://example.com"},
Listen: ":9999",
Timeout: 30,
Client: client,
}
srv := newServer(config)
assert.Equal(t, ":9999", srv.Addr)
assert.Equal(t, readHeaderTimeout, srv.ReadHeaderTimeout)
assert.Equal(t, readTimeout, srv.ReadTimeout)
assert.Equal(t, idleTimeout, srv.IdleTimeout)
assert.NotNil(t, srv.Handler)
}
func TestNewRouter_RoutesExist(t *testing.T) {
gin.SetMode(gin.TestMode)
client, err := internal.NewClient(10)
require.NoError(t, err)
config := Config{
ServerURL: "http://example.com",
AllowedServers: map[string]string{"http://example.com": "http://example.com"},
Listen: ":0",
Timeout: 10,
Client: client,
}
r := newRouter(config)
// Ping route
w := performJSON(r, http.MethodGet, "/ping", "")
assert.Equal(t, http.StatusOK, w.Code)
assert.JSONEq(t, `{"message":"pong"}`, w.Body.String())
// Send route (will fail with 400 due to missing fields, but not 404)
w = performJSON(r, http.MethodPost, "/send", `{}`)
assert.NotEqual(t, http.StatusNotFound, w.Code)
// Receive route (will fail with 400 due to missing fields, but not 404)
w = performJSON(r, http.MethodPost, "/receive", `{}`)
assert.NotEqual(t, http.StatusNotFound, w.Code)
}
func TestSendSuccess(t *testing.T) {
@@ -117,3 +394,119 @@ func TestHTTP500WithSOAPFaultReturnsFaultText(t *testing.T) {
assert.Contains(t, w.Body.String(), "Can not find the event")
assert.NotContains(t, w.Body.String(), "ciims returned 500")
}
func TestLoadConfigRequiresCIIMSServer(t *testing.T) {
t.Setenv("CIIMS_SERVER", "")
t.Setenv("CIIMS_TIMEOUT", "")
t.Setenv("CIIMS_ALLOWED_SERVERS", "")
_, err := loadConfig()
assert.Error(t, err)
assert.Contains(t, err.Error(), "CIIMS_SERVER")
}
func TestLoadConfigRejectsInvalidCIIMSServer(t *testing.T) {
t.Setenv("CIIMS_SERVER", "ftp://ciims.example.com")
t.Setenv("CIIMS_TIMEOUT", "")
t.Setenv("CIIMS_ALLOWED_SERVERS", "")
_, err := loadConfig()
assert.Error(t, err)
assert.Contains(t, err.Error(), "scheme")
}
func TestLoadConfigRejectsNonPositiveTimeout(t *testing.T) {
t.Setenv("CIIMS_SERVER", "http://ciims.example.com")
t.Setenv("CIIMS_TIMEOUT", "0")
t.Setenv("CIIMS_ALLOWED_SERVERS", "")
_, err := loadConfig()
assert.Error(t, err)
assert.Contains(t, err.Error(), "positive")
}
func TestLoadConfigAllowedServers(t *testing.T) {
t.Setenv("CIIMS_SERVER", "HTTP://ciims.example.com/base/")
t.Setenv("CIIMS_ALLOWED_SERVERS", "https://backup.example.com/ciims/")
t.Setenv("CIIMS_TIMEOUT", "3")
config, err := loadConfig()
require.NoError(t, err)
assert.Equal(t, "http://ciims.example.com/base", config.ServerURL)
assert.Contains(t, config.AllowedServers, "http://ciims.example.com/base")
assert.Contains(t, config.AllowedServers, "https://backup.example.com/ciims")
assert.Equal(t, 3, config.Timeout)
}
func TestLoadConfigRejectsInvalidAllowedServer(t *testing.T) {
t.Setenv("CIIMS_SERVER", "http://ciims.example.com")
t.Setenv("CIIMS_ALLOWED_SERVERS", "http://bad.example.com?x=1")
t.Setenv("CIIMS_TIMEOUT", "")
_, err := loadConfig()
assert.Error(t, err)
assert.Contains(t, err.Error(), "CIIMS_ALLOWED_SERVERS")
}
func TestNewServerTimeouts(t *testing.T) {
client, err := internal.NewClient(10)
require.NoError(t, err)
config := Config{ServerURL: "http://ciims.example.com", Listen: ":0", Timeout: 10, Client: client}
srv := newServer(config)
assert.Equal(t, readHeaderTimeout, srv.ReadHeaderTimeout)
assert.Equal(t, readTimeout, srv.ReadTimeout)
assert.Equal(t, 20*time.Second, srv.WriteTimeout)
assert.Equal(t, idleTimeout, srv.IdleTimeout)
}
func TestDefaultURLUsesCIIMSServer(t *testing.T) {
var gotPath string
r := testRouter(t, func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
fmt.Fprint(w, testSendOK)
})
w := performJSON(r, http.MethodPost, "/send", `{"user":"FIMS","pass":"x","event":"E1","msg":"<MSG/>"}`)
assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, servicePrefix, gotPath)
}
func TestAllowedRequestURLSucceeds(t *testing.T) {
gin.SetMode(gin.TestMode)
ciims := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
fmt.Fprint(w, testSendOK)
}))
defer ciims.Close()
client, err := internal.NewClient(10)
require.NoError(t, err)
r := newRouter(Config{
ServerURL: "http://default.example.com",
Listen: ":0",
Timeout: 10,
Client: client,
AllowedServers: map[string]string{ciims.URL: ciims.URL, "http://default.example.com": "http://default.example.com"},
})
body := fmt.Sprintf(`{"url":%q,"user":"FIMS","pass":"x","event":"E1","msg":"<MSG/>"}`, ciims.URL+"/")
w := performJSON(r, http.MethodPost, "/send", body)
assert.Equal(t, http.StatusOK, w.Code)
}
func TestDisallowedRequestURLReturns400AndDoesNotCallBackend(t *testing.T) {
calls := 0
r := testRouter(t, func(w http.ResponseWriter, r *http.Request) {
calls++
fmt.Fprint(w, testSendOK)
})
w := performJSON(r, http.MethodPost, "/send", `{"url":"http://evil.example.com","user":"FIMS","pass":"x","event":"E1","msg":"<MSG/>"}`)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Equal(t, 0, calls)
}
func TestRequestURLWithQueryReturns400(t *testing.T) {
r := testRouter(t, func(w http.ResponseWriter, r *http.Request) {})
w := performJSON(r, http.MethodPost, "/send", `{"url":"http://evil.example.com?x=1","user":"FIMS","pass":"x","event":"E1","msg":"<MSG/>"}`)
assert.Equal(t, http.StatusBadRequest, w.Code)
}
func TestHTTP500WithoutSOAPFaultReturnsStatusError(t *testing.T) {
r := testRouter(t, func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "backend failed", http.StatusInternalServerError)
})
w := performJSON(r, http.MethodPost, "/send", `{"user":"FIMS","pass":"x","event":"E1","msg":"<MSG/>"}`)
assert.Equal(t, http.StatusInternalServerError, w.Code)
assert.Contains(t, w.Body.String(), "ciims returned 500")
}