fix: harden scan edge cases

This commit is contained in:
ZacharyZcR
2026-06-01 03:32:13 +08:00
parent 8b558b4f12
commit 8ec96bfe6d
21 changed files with 534 additions and 67 deletions
+28 -2
View File
@@ -163,6 +163,21 @@ func (p *Neo4jPlugin) testUnauthorizedAccess(ctx context.Context, info *common.H
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode == 200 {
body, err := io.ReadAll(resp.Body)
if err != nil {
return &ScanResult{
Success: false,
Service: "neo4j",
Error: err,
}
}
if !strings.Contains(strings.ToLower(string(body)), "neo4j") {
return &ScanResult{
Success: false,
Service: "neo4j",
Error: fmt.Errorf("%s", i18n.Tr("service_not_identified", "Neo4j")),
}
}
return &ScanResult{
Type: plugins.ResultTypeVuln,
Success: true,
@@ -206,11 +221,22 @@ func (p *Neo4jPlugin) identifyService(ctx context.Context, info *common.HostInfo
if serverHeader != "" && strings.Contains(strings.ToLower(serverHeader), "neo4j") {
banner = "Neo4j"
} else if resp.StatusCode == 200 || resp.StatusCode == 401 {
body, _ := io.ReadAll(resp.Body)
body, err := io.ReadAll(resp.Body)
if err != nil {
return &ScanResult{
Success: false,
Service: "neo4j",
Error: err,
}
}
if strings.Contains(strings.ToLower(string(body)), "neo4j") {
banner = "Neo4j"
} else {
banner = "Neo4j"
return &ScanResult{
Success: false,
Service: "neo4j",
Error: fmt.Errorf("%s", i18n.Tr("service_not_identified", "Neo4j")),
}
}
} else {
return &ScanResult{
+59
View File
@@ -0,0 +1,59 @@
package services
import (
"context"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strconv"
"testing"
"github.com/shadow1ng/fscan/common"
)
func testSession() *common.ScanSession {
cfg := common.NewConfig()
return common.NewScanSession(cfg, common.NewState(), &common.FlagVars{})
}
func hostInfoFromServer(t *testing.T, server *httptest.Server) *common.HostInfo {
t.Helper()
u, err := url.Parse(server.URL)
if err != nil {
t.Fatalf("Parse server URL error = %v", err)
}
host, portText, err := net.SplitHostPort(u.Host)
if err != nil {
t.Fatalf("SplitHostPort error = %v", err)
}
port, err := strconv.Atoi(portText)
if err != nil {
t.Fatalf("Atoi port error = %v", err)
}
return &common.HostInfo{Host: host, Port: port}
}
func TestNeo4jIdentifyRejectsGenericHTTP(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write([]byte("plain http service"))
}))
defer server.Close()
result := NewNeo4jPlugin().identifyService(context.Background(), hostInfoFromServer(t, server), testSession())
if result.Success {
t.Fatalf("identifyService reported generic HTTP as Neo4j: %#v", result)
}
}
func TestNeo4jUnauthorizedRequiresNeo4jBody(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write([]byte("ok"))
}))
defer server.Close()
result := NewNeo4jPlugin().testUnauthorizedAccess(context.Background(), hostInfoFromServer(t, server), testSession())
if result != nil && result.Success {
t.Fatalf("testUnauthorizedAccess reported generic 200 as Neo4j: %#v", result)
}
}
+8 -1
View File
@@ -285,7 +285,14 @@ func (p *RabbitMQPlugin) testManagementInterface(ctx context.Context, info *comm
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode == 200 || resp.StatusCode == 401 {
body, _ := io.ReadAll(resp.Body)
body, err := io.ReadAll(resp.Body)
if err != nil {
return &ScanResult{
Success: false,
Service: "rabbitmq",
Error: err,
}
}
if strings.Contains(strings.ToLower(string(body)), "rabbitmq") {
banner := "RabbitMQ Management"
common.LogSuccess(i18n.Tr("rabbitmq_detected", target, banner))
+20
View File
@@ -0,0 +1,20 @@
package services
import (
"context"
"net/http"
"net/http/httptest"
"testing"
)
func TestRabbitMQManagementRejectsGenericHTTP(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write([]byte("plain http service"))
}))
defer server.Close()
result := NewRabbitMQPlugin().testManagementInterface(context.Background(), hostInfoFromServer(t, server), testSession())
if result.Success {
t.Fatalf("testManagementInterface reported generic HTTP as RabbitMQ: %#v", result)
}
}