1
0

serve_test.go 2.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111
  1. package network
  2. import (
  3. "errors"
  4. "go/ast"
  5. "go/parser"
  6. "go/token"
  7. "net"
  8. "net/http"
  9. "os"
  10. "path/filepath"
  11. "runtime"
  12. "strings"
  13. "testing"
  14. "github.com/mhsanaei/3x-ui/v3/internal/logger"
  15. )
  16. type failingListener struct{ err error }
  17. func (l failingListener) Accept() (net.Conn, error) { return nil, l.err }
  18. func (failingListener) Close() error { return nil }
  19. func (failingListener) Addr() net.Addr { return testAddr("failing") }
  20. type testAddr string
  21. func (a testAddr) Network() string { return string(a) }
  22. func (a testAddr) String() string { return string(a) }
  23. func TestServeHTTPLogsUnexpectedListenerFailure(t *testing.T) {
  24. errInjected := errors.New("injected listener failure")
  25. ServeHTTP(&http.Server{}, failingListener{err: errInjected}, "Test server")
  26. for _, line := range logger.GetLogs(100, "error") {
  27. if strings.Contains(line, errInjected.Error()) {
  28. return
  29. }
  30. }
  31. t.Fatal("unexpected listener failure was not recorded in the panel log")
  32. }
  33. func TestServeHTTPSuppressesNormalServerClose(t *testing.T) {
  34. const marker = "normal-close-must-stay-silent"
  35. ServeHTTP(&http.Server{}, failingListener{err: http.ErrServerClosed}, marker)
  36. for _, line := range logger.GetLogs(100, "error") {
  37. if strings.Contains(line, marker) {
  38. t.Fatalf("normal http.ErrServerClosed was recorded as an error: %s", line)
  39. }
  40. }
  41. }
  42. func TestProductionHTTPServersUseServeHTTPWrapper(t *testing.T) {
  43. _, currentFile, _, ok := runtime.Caller(0)
  44. if !ok {
  45. t.Fatal("locate test source")
  46. }
  47. repoRoot := filepath.Clean(filepath.Join(filepath.Dir(currentFile), "../../.."))
  48. fset := token.NewFileSet()
  49. err := filepath.WalkDir(repoRoot, func(path string, entry os.DirEntry, walkErr error) error {
  50. if walkErr != nil {
  51. return walkErr
  52. }
  53. if entry.IsDir() {
  54. if entry.Name() == ".git" || entry.Name() == "vendor" || entry.Name() == "node_modules" {
  55. return filepath.SkipDir
  56. }
  57. return nil
  58. }
  59. if !strings.HasSuffix(path, ".go") || strings.HasSuffix(path, "_test.go") || path == currentFile || path == filepath.Join(filepath.Dir(currentFile), "serve.go") {
  60. return nil
  61. }
  62. parsed, err := parser.ParseFile(fset, path, nil, parser.ImportsOnly)
  63. if err != nil {
  64. return err
  65. }
  66. usesHTTP := false
  67. for _, imp := range parsed.Imports {
  68. if imp.Path.Value == `"net/http"` {
  69. usesHTTP = true
  70. break
  71. }
  72. }
  73. if !usesHTTP {
  74. return nil
  75. }
  76. parsed, err = parser.ParseFile(fset, path, nil, 0)
  77. if err != nil {
  78. return err
  79. }
  80. ast.Inspect(parsed, func(node ast.Node) bool {
  81. call, ok := node.(*ast.CallExpr)
  82. if !ok {
  83. return true
  84. }
  85. selector, ok := call.Fun.(*ast.SelectorExpr)
  86. if ok && selector.Sel.Name == "Serve" {
  87. position := fset.Position(call.Pos())
  88. t.Errorf("direct Serve call at %s; production HTTP servers must use network.ServeHTTP", position)
  89. }
  90. return true
  91. })
  92. return nil
  93. })
  94. if err != nil {
  95. t.Fatalf("scan production Go files: %v", err)
  96. }
  97. }