diff --git a/ops/ops.go b/ops/ops.go index 4df7c1f7..4af52d61 100644 --- a/ops/ops.go +++ b/ops/ops.go @@ -13,6 +13,8 @@ import ( "github.com/specterops/dawgs/util/size" ) +var ErrGraphQueryExecutionFailed = errors.New("query execution failed") + func FetchAllNodeProperties(tx graph.Transaction, nodes graph.NodeSet) error { return tx.Nodes().Filter( query.InIDs(query.NodeID(), nodes.IDs()...), @@ -195,7 +197,7 @@ func FetchByQuery(tx graph.Transaction, query string) (QueryResult, error) { ) if queryResult := tx.Query(query, map[string]any{}); queryResult.Error() != nil { - return result, queryResult.Error() + return result, fmt.Errorf("%w: %w", ErrGraphQueryExecutionFailed, queryResult.Error()) } else { defer queryResult.Close() @@ -235,7 +237,7 @@ func FetchByQuery(tx graph.Transaction, query string) (QueryResult, error) { ) if currentPathSize > tx.GraphQueryMemoryLimit() || pathSetSize+literalSize > tx.GraphQueryMemoryLimit() { - return result, fmt.Errorf("%s - Limit: %.2f MB", "query required more memory than allowed", tx.GraphQueryMemoryLimit().Mebibytes()) + return result, fmt.Errorf("%w - Limit: %.2f MB", ErrGraphQueryMemoryLimit, tx.GraphQueryMemoryLimit().Mebibytes()) } } } diff --git a/util/errors.go b/util/errors.go index aefe76e2..d9d97690 100644 --- a/util/errors.go +++ b/util/errors.go @@ -5,6 +5,7 @@ import ( "strings" "sync" + "github.com/jackc/pgx/v5/pgconn" "github.com/neo4j/neo4j-go-driver/v5/neo4j" ) @@ -50,3 +51,15 @@ func IsNeoTimeoutError(err error) bool { return strings.Contains(e.Error(), "Neo.ClientError.Transaction.TransactionTimedOut") } } + +func IsPostgresTimeoutError(err error) bool { + if err == nil { + return false + } + + var pgErr *pgconn.PgError + return errors.As(err, &pgErr) && + pgErr != nil && + pgErr.Code == "57014" && + strings.Contains(strings.ToLower(pgErr.Message), "statement timeout") +} diff --git a/util/errors_test.go b/util/errors_test.go index f46fc0f7..b3ba99e2 100644 --- a/util/errors_test.go +++ b/util/errors_test.go @@ -1,10 +1,12 @@ package util_test import ( + "context" "errors" "fmt" "testing" + "github.com/jackc/pgx/v5/pgconn" "github.com/neo4j/neo4j-go-driver/v5/neo4j" "github.com/specterops/dawgs/graph" "github.com/specterops/dawgs/util" @@ -38,6 +40,27 @@ func TestIsNeoTimeoutError(t *testing.T) { require.False(t, util.IsNeoTimeoutError(notDriverTimeOutErr)) } +func TestIsPostgresTimeoutError(t *testing.T) { + statementTimeoutErr := &pgconn.PgError{ + Code: "57014", + Message: "canceling statement due to statement timeout", + } + + require.False(t, util.IsPostgresTimeoutError(nil)) + require.True(t, util.IsPostgresTimeoutError(statementTimeoutErr)) + require.True(t, util.IsPostgresTimeoutError(fmt.Errorf("wrapped: %w", statementTimeoutErr))) + require.False(t, util.IsPostgresTimeoutError(context.DeadlineExceeded)) + require.False(t, util.IsPostgresTimeoutError(&pgconn.PgError{ + Code: "57014", + Message: "canceling statement due to user request", + })) + require.False(t, util.IsPostgresTimeoutError(&pgconn.PgError{ + Code: "55P03", + Message: "canceling statement due to lock timeout", + })) + require.False(t, util.IsPostgresTimeoutError(errors.New("statement timeout"))) +} + func TestNewErrorCollector(t *testing.T) { errCollector := util.NewErrorCollector()