fix(go.sum): update ResolveSpec dependency to v1.0.87
This commit is contained in:
@@ -3,30 +3,29 @@ module git.warky.dev/wdevs/amcs
|
|||||||
go 1.26.1
|
go 1.26.1
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/bitechdev/ResolveSpec v1.0.87
|
github.com/bitechdev/ResolveSpec v1.1.15
|
||||||
github.com/google/jsonschema-go v0.4.2
|
github.com/google/jsonschema-go v0.4.3
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/jackc/pgx/v5 v5.9.1
|
github.com/jackc/pgx/v5 v5.9.2
|
||||||
github.com/modelcontextprotocol/go-sdk v1.4.1
|
github.com/modelcontextprotocol/go-sdk v1.4.1
|
||||||
github.com/pgvector/pgvector-go v0.3.0
|
github.com/pgvector/pgvector-go v0.3.0
|
||||||
github.com/spf13/cobra v1.10.2
|
github.com/spf13/cobra v1.10.2
|
||||||
github.com/uptrace/bun v1.2.16
|
github.com/uptrace/bun v1.2.18
|
||||||
github.com/uptrace/bun/dialect/pgdialect v1.2.16
|
github.com/uptrace/bun/dialect/pgdialect v1.2.16
|
||||||
github.com/uptrace/bun/driver/pgdriver v1.1.12
|
github.com/uptrace/bun/driver/pgdriver v1.1.12
|
||||||
github.com/uptrace/bunrouter v1.0.23
|
github.com/uptrace/bunrouter v1.0.23
|
||||||
golang.org/x/sync v0.19.0
|
golang.org/x/sync v0.20.0
|
||||||
gopkg.in/yaml.v3 v3.0.1
|
gopkg.in/yaml.v3 v3.0.1
|
||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/beorn7/perks v1.0.1 // indirect
|
github.com/beorn7/perks v1.0.1 // indirect
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf // indirect
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c // indirect
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
github.com/fsnotify/fsnotify v1.10.1 // indirect
|
||||||
github.com/fsnotify/fsnotify v1.9.0 // indirect
|
github.com/getsentry/sentry-go v0.46.2 // indirect
|
||||||
github.com/getsentry/sentry-go v0.40.0 // indirect
|
github.com/go-viper/mapstructure/v2 v2.5.0 // indirect
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 // indirect
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 // indirect
|
||||||
github.com/golang-sql/sqlexp v0.1.0 // indirect
|
github.com/golang-sql/sqlexp v0.1.0 // indirect
|
||||||
github.com/gorilla/mux v1.8.1 // indirect
|
github.com/gorilla/mux v1.8.1 // indirect
|
||||||
@@ -36,17 +35,17 @@ require (
|
|||||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||||
github.com/jinzhu/inflection v1.0.0 // indirect
|
github.com/jinzhu/inflection v1.0.0 // indirect
|
||||||
github.com/jinzhu/now v1.1.5 // indirect
|
github.com/jinzhu/now v1.1.5 // indirect
|
||||||
github.com/mattn/go-sqlite3 v1.14.33 // indirect
|
github.com/mattn/go-sqlite3 v1.14.44 // indirect
|
||||||
github.com/microsoft/go-mssqldb v1.9.5 // indirect
|
github.com/microsoft/go-mssqldb v1.10.0 // indirect
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
github.com/pelletier/go-toml/v2 v2.3.1 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/prometheus/client_golang v1.23.2 // indirect
|
github.com/prometheus/client_golang v1.23.2 // indirect
|
||||||
github.com/prometheus/client_model v0.6.2 // indirect
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
github.com/prometheus/common v0.67.4 // indirect
|
github.com/prometheus/common v0.67.5 // indirect
|
||||||
github.com/prometheus/procfs v0.19.2 // indirect
|
github.com/prometheus/procfs v0.20.1 // indirect
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
||||||
github.com/redis/go-redis/v9 v9.17.2 // indirect
|
github.com/redis/go-redis/v9 v9.19.0 // indirect
|
||||||
github.com/rogpeppe/go-internal v1.14.1 // indirect
|
github.com/rogpeppe/go-internal v1.14.1 // indirect
|
||||||
github.com/sagikazarmark/locafero v0.12.0 // indirect
|
github.com/sagikazarmark/locafero v0.12.0 // indirect
|
||||||
github.com/segmentio/asm v1.1.3 // indirect
|
github.com/segmentio/asm v1.1.3 // indirect
|
||||||
@@ -58,7 +57,7 @@ require (
|
|||||||
github.com/spf13/viper v1.21.0 // indirect
|
github.com/spf13/viper v1.21.0 // indirect
|
||||||
github.com/stretchr/testify v1.11.1 // indirect
|
github.com/stretchr/testify v1.11.1 // indirect
|
||||||
github.com/subosito/gotenv v1.6.0 // indirect
|
github.com/subosito/gotenv v1.6.0 // indirect
|
||||||
github.com/tidwall/gjson v1.18.0 // indirect
|
github.com/tidwall/gjson v1.19.0 // indirect
|
||||||
github.com/tidwall/match v1.2.0 // indirect
|
github.com/tidwall/match v1.2.0 // indirect
|
||||||
github.com/tidwall/pretty v1.2.1 // indirect
|
github.com/tidwall/pretty v1.2.1 // indirect
|
||||||
github.com/tidwall/sjson v1.2.5 // indirect
|
github.com/tidwall/sjson v1.2.5 // indirect
|
||||||
@@ -69,15 +68,16 @@ require (
|
|||||||
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
github.com/vmihailenco/tagparser/v2 v2.0.0 // indirect
|
||||||
github.com/x448/float16 v0.8.4 // indirect
|
github.com/x448/float16 v0.8.4 // indirect
|
||||||
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
|
||||||
|
go.uber.org/atomic v1.11.0 // indirect
|
||||||
go.uber.org/multierr v1.11.0 // indirect
|
go.uber.org/multierr v1.11.0 // indirect
|
||||||
go.uber.org/zap v1.27.1 // indirect
|
go.uber.org/zap v1.28.0 // indirect
|
||||||
go.yaml.in/yaml/v2 v2.4.3 // indirect
|
go.yaml.in/yaml/v2 v2.4.4 // indirect
|
||||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||||
golang.org/x/crypto v0.46.0 // indirect
|
golang.org/x/crypto v0.51.0 // indirect
|
||||||
golang.org/x/mod v0.31.0 // indirect
|
golang.org/x/mod v0.36.0 // indirect
|
||||||
golang.org/x/oauth2 v0.34.0 // indirect
|
golang.org/x/oauth2 v0.36.0 // indirect
|
||||||
golang.org/x/sys v0.40.0 // indirect
|
golang.org/x/sys v0.44.0 // indirect
|
||||||
golang.org/x/text v0.32.0 // indirect
|
golang.org/x/text v0.37.0 // indirect
|
||||||
google.golang.org/protobuf v1.36.11 // indirect
|
google.golang.org/protobuf v1.36.11 // indirect
|
||||||
gorm.io/driver/postgres v1.6.0 // indirect
|
gorm.io/driver/postgres v1.6.0 // indirect
|
||||||
gorm.io/driver/sqlite v1.6.0 // indirect
|
gorm.io/driver/sqlite v1.6.0 // indirect
|
||||||
|
|||||||
@@ -5,41 +5,39 @@ entgo.io/ent v0.14.3/go.mod h1:aDPE/OziPEu8+OWbzy4UlvWmD2/kbRuWfK2A40hcxJM=
|
|||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.0/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.7.1/go.mod h1:bjGvMhVMb+EEm3VRNQawDMUyMMjo+S5ewNjflkep/0Q=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.11.1/go.mod h1:a6xsAQUZg+VsS3TJ05SRp524Hs4pZ/AeFSr5ENf0Yjo=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0 h1:Gt0j3wceWMwPmiazCa8MzMA0MfhmPIz0Qp0FJ6qcM0U=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.18.0/go.mod h1:Ot/6aikWnKWi4l9QB7qVSwa8iMphQNqkWALMoNT3rzM=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.3.1/go.mod h1:uE9zaUfEQT/nbQjVi2IblCG9iaLtZsuYZ8ne+PuQ02M=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.3.1/go.mod h1:uE9zaUfEQT/nbQjVi2IblCG9iaLtZsuYZ8ne+PuQ02M=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.6.0/go.mod h1:9kIvujWAA58nmPmWB1m23fyWic1kYZMxD9CxaWn4Qpg=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1 h1:B+blDbyVIG3WaikNxPnhPiJ1MThR03b3vKGtER95TP4=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.10.1/go.mod h1:JdM5psgjfBf5fo2uWOZhflPWyDBZ/O/CNAH9CtsuZE4=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1/go.mod h1:IYus9qsFobWIc2YVwe/WPjcnyCkPKtnHAqUYeebc8z0=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.3.0/go.mod h1:okt5dMMTOFjX/aovMlrjvvXoPMBVSPzk9185BT0+eZM=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.3.0/go.mod h1:okt5dMMTOFjX/aovMlrjvvXoPMBVSPzk9185BT0+eZM=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.5.2/go.mod h1:yInRyqWXAuaPrgI7p70+lDDgh3mlBohis29jGMISnmc=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.8.0/go.mod h1:4OG6tQ9EOP/MT0NMjDlRzWoVFxfu9rN9B2X+tlSVktg=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1 h1:FPKJS1T+clwv+OLGt13a8UjqeRuh0O4SJ3lUriThc+4=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0 h1:fhqpLE3UEXi9lPaBRpQ6XuRW0nU7hgg4zlmZZa+a9q4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/internal v1.11.1/go.mod h1:j2chePtV91HrC22tGoRX3sGY42uF13WzmmV80/OdVAA=
|
github.com/Azure/azure-sdk-for-go/sdk/internal v1.12.0/go.mod h1:7dCRMLwisfRH3dBupKeNCioWYUZ4SS09Z14H+7i8ZoY=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.0.1/go.mod h1:GpPjLhVR9dnUoJMyHWSPy71xY9/lcmpzIPZXmF0FCVY=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.0.1/go.mod h1:GpPjLhVR9dnUoJMyHWSPy71xY9/lcmpzIPZXmF0FCVY=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1 h1:Wgf5rZba3YZqeTNJPtvqZoBu1sBN/L4sry+u2U3Y75w=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0 h1:E4MgwLBGeVB5f2MdcIVD3ELVAWpr+WD6MUe1i+tM/PA=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.3.1/go.mod h1:xxCBG/f/4Vbmh2XQJBsOmNdxWUY5j/s27jujKPbQf14=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys v1.4.0/go.mod h1:Y2b/1clN4zsAoUd/pgNAQHjLDnTis/6ROkUfyob6psM=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.0.0/go.mod h1:bTSOgj05NGRuHHhQwAdPnYr9TOdNmKlZTgGLL6nyAdI=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.0.0/go.mod h1:bTSOgj05NGRuHHhQwAdPnYr9TOdNmKlZTgGLL6nyAdI=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1 h1:bFWuoEKg+gImo7pvkiQEFAc8ocibADgXeiLAxWhWmkI=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0 h1:nCYfgcSyHZXJI8J0IWE5MsCGlb2xp9fJiXyxWgmOFg4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.1.1/go.mod h1:Vih/3yc6yac2JzU4hzpaDupBJP0Flaia9rXXrU8xyww=
|
github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/internal v1.2.0/go.mod h1:ucUjca2JtSZboY8IoUqyQyuuXvwbMBVwFOm0vdQPNhA=
|
||||||
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
|
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
|
||||||
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.1.1/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.2.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2 h1:oygO0locgZJe7PpYPXT5A29ZkwJaPqcva7BVeemZOZs=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0 h1:XRzhVemXdgvJqCH0sFfrBUTnUJSBrBf7++ypk+twtRs=
|
||||||
github.com/AzureAD/microsoft-authentication-library-for-go v1.4.2/go.mod h1:wP83P5OoQ5p6ip3ScPr0BAq0BvuPAvacpEuSzyouqAI=
|
github.com/AzureAD/microsoft-authentication-library-for-go v1.6.0/go.mod h1:HKpQxkWaGLJ+D/5H8QRpyQXA1eKjxkFlOMwck5+33Jk=
|
||||||
github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
|
github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
|
||||||
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
|
||||||
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
||||||
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
||||||
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||||
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||||
github.com/bitechdev/ResolveSpec v1.0.86 h1:a4yFMMDizrmvDOV61cj/+kD+mEtKL/5EIHY2GcP3uJU=
|
github.com/bitechdev/ResolveSpec v1.1.15 h1:Bhy1ZWGUg9AJOhOLk5Di9ilOVKmZ40SkGIGDraAlti8=
|
||||||
github.com/bitechdev/ResolveSpec v1.0.86/go.mod h1:YZOY2YCD0Kmb+pjAMhOqPh4q82Hij57F/CLlCMkzT78=
|
github.com/bitechdev/ResolveSpec v1.1.15/go.mod h1:GF51sMRCWbAyri2WNae3IZAFM/2s6DG6i3eTTrobbVs=
|
||||||
github.com/bitechdev/ResolveSpec v1.0.87 h1:zLiHynLK8LLpXIfCZOjL5Iy1COBS6YZcWE1BHKfYqbA=
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c h1:6Gpm9YYUEQx2T9zMsYolQhr6sjwwGtFitSA0pQsa7a8=
|
||||||
github.com/bitechdev/ResolveSpec v1.0.87/go.mod h1:YZOY2YCD0Kmb+pjAMhOqPh4q82Hij57F/CLlCMkzT78=
|
github.com/bradfitz/gomemcache v0.0.0-20260422231931-4d751bb6e37c/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf h1:TqhNAT4zKbTdLa62d2HDBFdvgSbIGB3eJE8HqhgiL9I=
|
|
||||||
github.com/bradfitz/gomemcache v0.0.0-20250403215159-8d39553ac7cf/go.mod h1:r5xuitiExdLAJ09PR7vBVENGvp4ZuTBeWTGtxuX3K+c=
|
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
||||||
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
||||||
@@ -61,10 +59,9 @@ github.com/cpuguy83/dockercfg v0.3.2/go.mod h1:sugsbF4//dDlL/i+S+rtpIWp+5h0BHJHf
|
|||||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||||
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
|
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
|
||||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
|
||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
|
||||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
|
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
|
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
|
||||||
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
|
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
|
||||||
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
|
github.com/dnaeon/go-vcr v1.1.0/go.mod h1:M7tiix8f0r6mKKJ3Yq/kqU1OYf3MnfmBWVbPx/yU9ko=
|
||||||
@@ -83,10 +80,10 @@ github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2
|
|||||||
github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
|
github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
|
||||||
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
||||||
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||||
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
|
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||||
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
||||||
github.com/getsentry/sentry-go v0.40.0 h1:VTJMN9zbTvqDqPwheRVLcp0qcUcM+8eFivvGocAaSbo=
|
github.com/getsentry/sentry-go v0.46.2 h1:1jhYwrKGa3sIpo/y5iDNXS5wDoT7I1KNzMHrnK6ojns=
|
||||||
github.com/getsentry/sentry-go v0.40.0/go.mod h1:eRXCoh3uvmjQLY6qu63BjUZnaBu5L5WhMV1RwYO8W5s=
|
github.com/getsentry/sentry-go v0.46.2/go.mod h1:evVbw2qotNUdYG8KxXbAdjOQWWvWIwKxpjdZZIvcIPw=
|
||||||
github.com/go-errors/errors v1.4.2 h1:J6MZopCL4uSllY1OfXM374weqZFFItUbrImctkmUxIA=
|
github.com/go-errors/errors v1.4.2 h1:J6MZopCL4uSllY1OfXM374weqZFFItUbrImctkmUxIA=
|
||||||
github.com/go-errors/errors v1.4.2/go.mod h1:sIVyrIiJhuEF+Pj9Ebtd6P/rEYROXFi3BopGUQ5a5Og=
|
github.com/go-errors/errors v1.4.2/go.mod h1:sIVyrIiJhuEF+Pj9Ebtd6P/rEYROXFi3BopGUQ5a5Og=
|
||||||
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
||||||
@@ -99,13 +96,13 @@ github.com/go-pg/pg/v10 v10.11.0 h1:CMKJqLgTrfpE/aOVeLdybezR2om071Vh38OLZjsyMI0=
|
|||||||
github.com/go-pg/pg/v10 v10.11.0/go.mod h1:4BpHRoxE61y4Onpof3x1a2SQvi9c+q1dJnrNdMjsroA=
|
github.com/go-pg/pg/v10 v10.11.0/go.mod h1:4BpHRoxE61y4Onpof3x1a2SQvi9c+q1dJnrNdMjsroA=
|
||||||
github.com/go-pg/zerochecker v0.2.0 h1:pp7f72c3DobMWOb2ErtZsnrPaSvHd2W4o9//8HtF4mU=
|
github.com/go-pg/zerochecker v0.2.0 h1:pp7f72c3DobMWOb2ErtZsnrPaSvHd2W4o9//8HtF4mU=
|
||||||
github.com/go-pg/zerochecker v0.2.0/go.mod h1:NJZ4wKL0NmTtz0GKCoJ8kym6Xn/EQzXRl2OnAe7MmDo=
|
github.com/go-pg/zerochecker v0.2.0/go.mod h1:NJZ4wKL0NmTtz0GKCoJ8kym6Xn/EQzXRl2OnAe7MmDo=
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
|
github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
||||||
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
|
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0=
|
||||||
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A=
|
||||||
@@ -113,8 +110,8 @@ github.com/golang-sql/sqlexp v0.1.0/go.mod h1:J4ad9Vo8ZCWQ2GMrC4UCQy1JpCbwU9m3EO
|
|||||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||||
github.com/google/jsonschema-go v0.4.2 h1:tmrUohrwoLZZS/P3x7ex0WAVknEkBZM46iALbcqoRA8=
|
github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0=
|
||||||
github.com/google/jsonschema-go v0.4.2/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
|
||||||
github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||||
@@ -130,8 +127,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
|
|||||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||||
github.com/jackc/pgx/v5 v5.9.1 h1:uwrxJXBnx76nyISkhr33kQLlUqjv7et7b9FjCen/tdc=
|
github.com/jackc/pgx/v5 v5.9.2 h1:3ZhOzMWnR4yJ+RW1XImIPsD1aNSz4T4fyP7zlQb56hw=
|
||||||
github.com/jackc/pgx/v5 v5.9.1/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
github.com/jackc/pgx/v5 v5.9.2/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||||
github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs=
|
github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs=
|
||||||
@@ -146,8 +143,10 @@ github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
|
|||||||
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
|
||||||
github.com/jmoiron/sqlx v1.3.5 h1:vFFPA71p1o5gAeqtEAwLU4dnX2napprKtHr7PYIcN3g=
|
github.com/jmoiron/sqlx v1.3.5 h1:vFFPA71p1o5gAeqtEAwLU4dnX2napprKtHr7PYIcN3g=
|
||||||
github.com/jmoiron/sqlx v1.3.5/go.mod h1:nRVWtLre0KfCLJvgxzCsLVMogSvQ1zNJtpYr2Ccp0mQ=
|
github.com/jmoiron/sqlx v1.3.5/go.mod h1:nRVWtLre0KfCLJvgxzCsLVMogSvQ1zNJtpYr2Ccp0mQ=
|
||||||
github.com/klauspost/compress v1.18.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk=
|
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
|
||||||
github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
|
github.com/klauspost/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||||
|
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
|
||||||
|
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
|
||||||
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
|
||||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||||
@@ -163,13 +162,13 @@ github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 h1:6E+4a0GO5zZEnZ
|
|||||||
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
|
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
|
||||||
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
|
||||||
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
|
||||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
|
||||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
|
||||||
github.com/mattn/go-sqlite3 v1.14.33 h1:A5blZ5ulQo2AtayQ9/limgHEkFreKj1Dv226a1K73s0=
|
github.com/mattn/go-sqlite3 v1.14.44 h1:3VSe+xafpbzsLbdr2AWlAZk9yRHiBhTBakioXaCKTF8=
|
||||||
github.com/mattn/go-sqlite3 v1.14.33/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
|
github.com/mattn/go-sqlite3 v1.14.44/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ=
|
||||||
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
|
github.com/microsoft/go-mssqldb v1.8.2/go.mod h1:vp38dT33FGfVotRiTmDo3bFyaHq+p3LektQrjTULowo=
|
||||||
github.com/microsoft/go-mssqldb v1.9.5 h1:orwya0X/5bsL1o+KasupTkk2eNTNFkTQG0BEe/HxCn0=
|
github.com/microsoft/go-mssqldb v1.10.0 h1:pHEt+Qz6YFPWqREq10mqSE524QQo+/QremwTCQht7TY=
|
||||||
github.com/microsoft/go-mssqldb v1.9.5/go.mod h1:VCP2a0KEZZtGLRHd1PsLavLFYy/3xX2yJUPycv3Sr2Q=
|
github.com/microsoft/go-mssqldb v1.10.0/go.mod h1:mnG7lGa9iYJbzJqGCXyuQCegStKMr3kogDLD6+bmggg=
|
||||||
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
|
github.com/moby/docker-image-spec v1.3.1 h1:jMKff3w6PgbfSa69GfNg+zN/XLhfXJGnEx3Nl2EsFP0=
|
||||||
github.com/moby/docker-image-spec v1.3.1/go.mod h1:eKmb5VW8vQEh/BAr2yvVNvuiJuY6UIocYsFu/DxxRpo=
|
github.com/moby/docker-image-spec v1.3.1/go.mod h1:eKmb5VW8vQEh/BAr2yvVNvuiJuY6UIocYsFu/DxxRpo=
|
||||||
github.com/moby/go-archive v0.1.0 h1:Kk/5rdW/g+H8NHdJW2gsXyZ7UnzvJNOy6VKJqueWdcQ=
|
github.com/moby/go-archive v0.1.0 h1:Kk/5rdW/g+H8NHdJW2gsXyZ7UnzvJNOy6VKJqueWdcQ=
|
||||||
@@ -198,8 +197,8 @@ github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8
|
|||||||
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
|
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
|
||||||
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
|
||||||
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
|
github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
|
github.com/pelletier/go-toml/v2 v2.3.1 h1:MYEvvGnQjeNkRF1qUuGolNtNExTDwct51yp7olPtrEc=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
github.com/pelletier/go-toml/v2 v2.3.1/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
||||||
github.com/pgvector/pgvector-go v0.3.0 h1:Ij+Yt78R//uYqs3Zk35evZFvr+G0blW0OUN+Q2D1RWc=
|
github.com/pgvector/pgvector-go v0.3.0 h1:Ij+Yt78R//uYqs3Zk35evZFvr+G0blW0OUN+Q2D1RWc=
|
||||||
github.com/pgvector/pgvector-go v0.3.0/go.mod h1:duFy+PXWfW7QQd5ibqutBO4GxLsUZ9RVXhFZGIBsWSA=
|
github.com/pgvector/pgvector-go v0.3.0/go.mod h1:duFy+PXWfW7QQd5ibqutBO4GxLsUZ9RVXhFZGIBsWSA=
|
||||||
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
|
github.com/pingcap/errors v0.11.4 h1:lFuQV/oaUMGcD2tqt+01ROSmJs75VG1ToEOkZIZ4nE4=
|
||||||
@@ -218,14 +217,14 @@ github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h
|
|||||||
github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
|
github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg=
|
||||||
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||||
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
||||||
github.com/prometheus/common v0.67.4 h1:yR3NqWO1/UyO1w2PhUvXlGQs/PtFmoveVO0KZ4+Lvsc=
|
github.com/prometheus/common v0.67.5 h1:pIgK94WWlQt1WLwAC5j2ynLaBRDiinoAb86HZHTUGI4=
|
||||||
github.com/prometheus/common v0.67.4/go.mod h1:gP0fq6YjjNCLssJCQp0yk4M8W6ikLURwkdd/YKtTbyI=
|
github.com/prometheus/common v0.67.5/go.mod h1:SjE/0MzDEEAyrdr5Gqc6G+sXI67maCxzaT3A2+HqjUw=
|
||||||
github.com/prometheus/procfs v0.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4Vws=
|
github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEycfc=
|
||||||
github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw=
|
github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo=
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
|
github.com/puzpuzpuz/xsync/v3 v3.5.1 h1:GJYJZwO6IdxN/IKbneznS6yPkVC+c3zyY/j19c++5Fg=
|
||||||
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
github.com/puzpuzpuz/xsync/v3 v3.5.1/go.mod h1:VjzYrABPabuM4KyBh1Ftq6u8nhwY5tBPKP9jpmh0nnA=
|
||||||
github.com/redis/go-redis/v9 v9.17.2 h1:P2EGsA4qVIM3Pp+aPocCJ7DguDHhqrXNhVcEp4ViluI=
|
github.com/redis/go-redis/v9 v9.19.0 h1:XPVaaPSnG6RhYf7p+rmSa9zZfeVAnWsH5h3lxthOm/k=
|
||||||
github.com/redis/go-redis/v9 v9.17.2/go.mod h1:u410H11HMLoB+TP67dz8rL9s6QW2j76l0//kSOd3370=
|
github.com/redis/go-redis/v9 v9.19.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||||
@@ -275,8 +274,8 @@ github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSW
|
|||||||
github.com/testcontainers/testcontainers-go v0.40.0 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU=
|
github.com/testcontainers/testcontainers-go v0.40.0 h1:pSdJYLOVgLE8YdUY2FHQ1Fxu+aMnb6JfVz1mxk7OeMU=
|
||||||
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
github.com/testcontainers/testcontainers-go v0.40.0/go.mod h1:FSXV5KQtX2HAMlm7U3APNyLkkap35zNLxukw9oBi/MY=
|
||||||
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
||||||
github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
|
github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU=
|
||||||
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc=
|
||||||
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
||||||
github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM=
|
github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM=
|
||||||
github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
|
||||||
@@ -291,8 +290,8 @@ github.com/tklauser/numcpus v0.6.1 h1:ng9scYS7az0Bk4OZLvrNXNSAO2Pxr1XXRAPyjhIx+F
|
|||||||
github.com/tklauser/numcpus v0.6.1/go.mod h1:1XfjsgE2zo8GVw7POkMbHENHzVg3GzmoZ9fESEdAacY=
|
github.com/tklauser/numcpus v0.6.1/go.mod h1:1XfjsgE2zo8GVw7POkMbHENHzVg3GzmoZ9fESEdAacY=
|
||||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYmW5DyG0UqvY96Bu5QYsTLvCHdrgo=
|
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYmW5DyG0UqvY96Bu5QYsTLvCHdrgo=
|
||||||
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
|
github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs=
|
||||||
github.com/uptrace/bun v1.2.16 h1:QlObi6ZIK5Ao7kAALnh91HWYNZUBbVwye52fmlQM9kc=
|
github.com/uptrace/bun v1.2.18 h1:3HnRcMfS6OBPMG1eSOzlbFJ/X/AyMEJb7rMxE6VQvDU=
|
||||||
github.com/uptrace/bun v1.2.16/go.mod h1:jMoNg2n56ckaawi/O/J92BHaECmrz6IRjuMWqlMaMTM=
|
github.com/uptrace/bun v1.2.18/go.mod h1:wNltaKJk4JtOt4SG5I5zmA7v0/Mzjh1+/S906Rayd3Y=
|
||||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.16 h1:rKv0cKPNBviXadB/+2Y/UedA/c1JnwGzUWZkdN5FdSQ=
|
github.com/uptrace/bun/dialect/mssqldialect v1.2.16 h1:rKv0cKPNBviXadB/+2Y/UedA/c1JnwGzUWZkdN5FdSQ=
|
||||||
github.com/uptrace/bun/dialect/mssqldialect v1.2.16/go.mod h1:J5U7tGKWDsx2Q7MwDZF2417jCdpD6yD/ZMFJcCR80bk=
|
github.com/uptrace/bun/dialect/mssqldialect v1.2.16/go.mod h1:J5U7tGKWDsx2Q7MwDZF2417jCdpD6yD/ZMFJcCR80bk=
|
||||||
github.com/uptrace/bun/dialect/pgdialect v1.2.16 h1:KFNZ0LxAyczKNfK/IJWMyaleO6eI9/Z5tUv3DE1NVL4=
|
github.com/uptrace/bun/dialect/pgdialect v1.2.16 h1:KFNZ0LxAyczKNfK/IJWMyaleO6eI9/Z5tUv3DE1NVL4=
|
||||||
@@ -320,24 +319,28 @@ github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT0
|
|||||||
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
||||||
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
github.com/yusufpapurcu/wmi v1.2.4 h1:zFUKzehAFReQwLys1b/iSMl+JQGSCSjtVqQn9bBrPo0=
|
||||||
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
|
||||||
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
|
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
||||||
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
|
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
||||||
|
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
|
||||||
|
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
|
||||||
go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8=
|
go.opentelemetry.io/otel v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I=
|
||||||
go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM=
|
go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0=
|
||||||
go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA=
|
go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM=
|
||||||
go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI=
|
go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY=
|
||||||
go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE=
|
go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A=
|
||||||
go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs=
|
go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0=
|
||||||
|
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
|
||||||
|
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
|
||||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
go.uber.org/multierr v1.11.0 h1:blXXJkSxSSfBVBlC76pxqeO+LN3aDfLQo+309xJstO0=
|
||||||
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
go.uber.org/multierr v1.11.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||||
go.uber.org/zap v1.27.1 h1:08RqriUEv8+ArZRYSTXy1LeBScaMpVSTBhCeaZYfMYc=
|
go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
||||||
go.uber.org/zap v1.27.1/go.mod h1:GB2qFLM7cTU87MWRP2mPIjqfIDnGu+VIO4V/SdhGo2E=
|
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
||||||
go.yaml.in/yaml/v2 v2.4.3 h1:6gvOSjQoTB3vt1l+CU+tSyi/HOjfOjRLJ4YwYZGwRO0=
|
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||||
go.yaml.in/yaml/v2 v2.4.3/go.mod h1:zSxWcmIDjOzPXpjlTTbAsKokqkDNAVtZO0WOMiT90s8=
|
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||||
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
||||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||||
@@ -352,18 +355,16 @@ golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOM
|
|||||||
golang.org/x/crypto v0.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
|
golang.org/x/crypto v0.22.0/go.mod h1:vr6Su+7cTlO45qkww3VDJlzDn0ctJvRgYbC2NvXHt+M=
|
||||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||||
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
golang.org/x/crypto v0.24.0/go.mod h1:Z1PMYSOR5nyMcyAVAIQSKCDwalqy85Aqn1x3Ws4L5DM=
|
||||||
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
|
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI=
|
||||||
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
|
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8=
|
||||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93 h1:fQsdNF2N+/YewlRZiricy4P1iimyPKZ/xwniHj8Q2a0=
|
|
||||||
golang.org/x/exp v0.0.0-20251219203646-944ab1f22d93/go.mod h1:EPRbTFwzwjXj9NpYyyrvenVh9Y+GFeEvMNh7Xuz7xgU=
|
|
||||||
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
||||||
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.9.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||||
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||||
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
|
||||||
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4=
|
||||||
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
|
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ=
|
||||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
@@ -381,10 +382,10 @@ golang.org/x/net v0.22.0/go.mod h1:JKghWKKOSdJwpW2GEx0Ja7fmaKnMsbu+MWVZTokSYmg=
|
|||||||
golang.org/x/net v0.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8=
|
golang.org/x/net v0.24.0/go.mod h1:2Q7sJY5mzlzWjKtYUEXSlBWCdyaioyXzRB2RtU8KVE8=
|
||||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||||
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
golang.org/x/net v0.26.0/go.mod h1:5YKkiSynbBIh3p6iOc/vibscux0x38BZDkn8sCUPxHE=
|
||||||
golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU=
|
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w=
|
||||||
golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY=
|
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ=
|
||||||
golang.org/x/oauth2 v0.34.0 h1:hqK/t4AKgbqWkdkcAeI8XLmbK+4m4G5YeQRrmiotGlw=
|
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
||||||
golang.org/x/oauth2 v0.34.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
|
golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q=
|
||||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
@@ -392,8 +393,8 @@ golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
|
|||||||
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
golang.org/x/sync v0.9.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
|
||||||
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||||
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
@@ -413,8 +414,8 @@ golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
|||||||
golang.org/x/sys v0.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.19.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.21.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ=
|
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ=
|
||||||
golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
|
||||||
@@ -443,16 +444,16 @@ golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
|||||||
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||||
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
|
golang.org/x/text v0.16.0/go.mod h1:GhwF1Be+LQoKShO3cGOHzqOgRrGaYc9AvblQOmPVHnI=
|
||||||
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
|
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
|
||||||
golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
|
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc=
|
||||||
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
|
golang.org/x/text v0.37.0/go.mod h1:a5sjxXGs9hsn/AJVwuElvCAo9v8QYLzvavO5z2PiM38=
|
||||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||||
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
||||||
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
||||||
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
|
||||||
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
golang.org/x/tools v0.13.0/go.mod h1:HvlwmtVNQAhOuCjW7xxvovg8wbNq7LwfXh/k7wXUl58=
|
||||||
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
|
||||||
golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
|
golang.org/x/tools v0.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c=
|
||||||
golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
|
golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI=
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||||
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
@@ -477,11 +478,11 @@ gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg=
|
|||||||
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
|
gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
|
||||||
mellium.im/sasl v0.3.1 h1:wE0LW6g7U83vhvxjC1IY8DnXM+EU095yeo8XClvCdfo=
|
mellium.im/sasl v0.3.1 h1:wE0LW6g7U83vhvxjC1IY8DnXM+EU095yeo8XClvCdfo=
|
||||||
mellium.im/sasl v0.3.1/go.mod h1:xm59PUYpZHhgQ9ZqoJ5QaCqzWMi8IeS49dhp6plPCzw=
|
mellium.im/sasl v0.3.1/go.mod h1:xm59PUYpZHhgQ9ZqoJ5QaCqzWMi8IeS49dhp6plPCzw=
|
||||||
modernc.org/libc v1.67.4 h1:zZGmCMUVPORtKv95c2ReQN5VDjvkoRm9GWPTEPuvlWg=
|
modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU=
|
||||||
modernc.org/libc v1.67.4/go.mod h1:QvvnnJ5P7aitu0ReNpVIEyesuhmDLQ8kaEoyMjIFZJA=
|
modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs=
|
||||||
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
|
||||||
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
|
||||||
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=
|
||||||
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=
|
||||||
modernc.org/sqlite v1.42.2 h1:7hkZUNJvJFN2PgfUdjni9Kbvd4ef4mNLOu0B9FGxM74=
|
modernc.org/sqlite v1.50.1 h1:l+cQvn0sd0zJJtfygGHuQJ5AjlrwXmWPw4KP3ZMwr9w=
|
||||||
modernc.org/sqlite v1.42.2/go.mod h1:+VkC6v3pLOAE0A0uVucQEcbVW0I5nHCeDaBf+DpsQT8=
|
modernc.org/sqlite v1.50.1/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM=
|
||||||
|
|||||||
Vendored
-2
@@ -1,2 +0,0 @@
|
|||||||
This placeholder keeps internal/app/ui/dist present in clean source checkouts.
|
|
||||||
The real UI bundle is generated by the frontend build into this directory.
|
|
||||||
Executable
+8
@@ -0,0 +1,8 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
export GONOSUMDB="*"
|
||||||
|
export GOSUMDB=off
|
||||||
|
export GOPRIVATE="github.com/bitechdev/"
|
||||||
|
go get -u -v github.com/bitechdev/ResolveSpec@latest
|
||||||
|
go mod tidy
|
||||||
|
go mod vendor
|
||||||
+20
@@ -0,0 +1,20 @@
|
|||||||
|
Copyright (C) 2013 Blake Mizerany
|
||||||
|
|
||||||
|
Permission is hereby granted, free of charge, to any person obtaining
|
||||||
|
a copy of this software and associated documentation files (the
|
||||||
|
"Software"), to deal in the Software without restriction, including
|
||||||
|
without limitation the rights to use, copy, modify, merge, publish,
|
||||||
|
distribute, sublicense, and/or sell copies of the Software, and to
|
||||||
|
permit persons to whom the Software is furnished to do so, subject to
|
||||||
|
the following conditions:
|
||||||
|
|
||||||
|
The above copyright notice and this permission notice shall be
|
||||||
|
included in all copies or substantial portions of the Software.
|
||||||
|
|
||||||
|
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
|
||||||
|
EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF
|
||||||
|
MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
|
||||||
|
NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE
|
||||||
|
LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION
|
||||||
|
OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
|
||||||
|
WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
+2388
File diff suppressed because it is too large
Load Diff
+316
@@ -0,0 +1,316 @@
|
|||||||
|
// Package quantile computes approximate quantiles over an unbounded data
|
||||||
|
// stream within low memory and CPU bounds.
|
||||||
|
//
|
||||||
|
// A small amount of accuracy is traded to achieve the above properties.
|
||||||
|
//
|
||||||
|
// Multiple streams can be merged before calling Query to generate a single set
|
||||||
|
// of results. This is meaningful when the streams represent the same type of
|
||||||
|
// data. See Merge and Samples.
|
||||||
|
//
|
||||||
|
// For more detailed information about the algorithm used, see:
|
||||||
|
//
|
||||||
|
// Effective Computation of Biased Quantiles over Data Streams
|
||||||
|
//
|
||||||
|
// http://www.cs.rutgers.edu/~muthu/bquant.pdf
|
||||||
|
package quantile
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"sort"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Sample holds an observed value and meta information for compression. JSON
|
||||||
|
// tags have been added for convenience.
|
||||||
|
type Sample struct {
|
||||||
|
Value float64 `json:",string"`
|
||||||
|
Width float64 `json:",string"`
|
||||||
|
Delta float64 `json:",string"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Samples represents a slice of samples. It implements sort.Interface.
|
||||||
|
type Samples []Sample
|
||||||
|
|
||||||
|
func (a Samples) Len() int { return len(a) }
|
||||||
|
func (a Samples) Less(i, j int) bool { return a[i].Value < a[j].Value }
|
||||||
|
func (a Samples) Swap(i, j int) { a[i], a[j] = a[j], a[i] }
|
||||||
|
|
||||||
|
type invariant func(s *stream, r float64) float64
|
||||||
|
|
||||||
|
// NewLowBiased returns an initialized Stream for low-biased quantiles
|
||||||
|
// (e.g. 0.01, 0.1, 0.5) where the needed quantiles are not known a priori, but
|
||||||
|
// error guarantees can still be given even for the lower ranks of the data
|
||||||
|
// distribution.
|
||||||
|
//
|
||||||
|
// The provided epsilon is a relative error, i.e. the true quantile of a value
|
||||||
|
// returned by a query is guaranteed to be within (1±Epsilon)*Quantile.
|
||||||
|
//
|
||||||
|
// See http://www.cs.rutgers.edu/~muthu/bquant.pdf for time, space, and error
|
||||||
|
// properties.
|
||||||
|
func NewLowBiased(epsilon float64) *Stream {
|
||||||
|
ƒ := func(s *stream, r float64) float64 {
|
||||||
|
return 2 * epsilon * r
|
||||||
|
}
|
||||||
|
return newStream(ƒ)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewHighBiased returns an initialized Stream for high-biased quantiles
|
||||||
|
// (e.g. 0.01, 0.1, 0.5) where the needed quantiles are not known a priori, but
|
||||||
|
// error guarantees can still be given even for the higher ranks of the data
|
||||||
|
// distribution.
|
||||||
|
//
|
||||||
|
// The provided epsilon is a relative error, i.e. the true quantile of a value
|
||||||
|
// returned by a query is guaranteed to be within 1-(1±Epsilon)*(1-Quantile).
|
||||||
|
//
|
||||||
|
// See http://www.cs.rutgers.edu/~muthu/bquant.pdf for time, space, and error
|
||||||
|
// properties.
|
||||||
|
func NewHighBiased(epsilon float64) *Stream {
|
||||||
|
ƒ := func(s *stream, r float64) float64 {
|
||||||
|
return 2 * epsilon * (s.n - r)
|
||||||
|
}
|
||||||
|
return newStream(ƒ)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTargeted returns an initialized Stream concerned with a particular set of
|
||||||
|
// quantile values that are supplied a priori. Knowing these a priori reduces
|
||||||
|
// space and computation time. The targets map maps the desired quantiles to
|
||||||
|
// their absolute errors, i.e. the true quantile of a value returned by a query
|
||||||
|
// is guaranteed to be within (Quantile±Epsilon).
|
||||||
|
//
|
||||||
|
// See http://www.cs.rutgers.edu/~muthu/bquant.pdf for time, space, and error properties.
|
||||||
|
func NewTargeted(targetMap map[float64]float64) *Stream {
|
||||||
|
// Convert map to slice to avoid slow iterations on a map.
|
||||||
|
// ƒ is called on the hot path, so converting the map to a slice
|
||||||
|
// beforehand results in significant CPU savings.
|
||||||
|
targets := targetMapToSlice(targetMap)
|
||||||
|
|
||||||
|
ƒ := func(s *stream, r float64) float64 {
|
||||||
|
var m = math.MaxFloat64
|
||||||
|
var f float64
|
||||||
|
for _, t := range targets {
|
||||||
|
if t.quantile*s.n <= r {
|
||||||
|
f = (2 * t.epsilon * r) / t.quantile
|
||||||
|
} else {
|
||||||
|
f = (2 * t.epsilon * (s.n - r)) / (1 - t.quantile)
|
||||||
|
}
|
||||||
|
if f < m {
|
||||||
|
m = f
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
return newStream(ƒ)
|
||||||
|
}
|
||||||
|
|
||||||
|
type target struct {
|
||||||
|
quantile float64
|
||||||
|
epsilon float64
|
||||||
|
}
|
||||||
|
|
||||||
|
func targetMapToSlice(targetMap map[float64]float64) []target {
|
||||||
|
targets := make([]target, 0, len(targetMap))
|
||||||
|
|
||||||
|
for quantile, epsilon := range targetMap {
|
||||||
|
t := target{
|
||||||
|
quantile: quantile,
|
||||||
|
epsilon: epsilon,
|
||||||
|
}
|
||||||
|
targets = append(targets, t)
|
||||||
|
}
|
||||||
|
|
||||||
|
return targets
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stream computes quantiles for a stream of float64s. It is not thread-safe by
|
||||||
|
// design. Take care when using across multiple goroutines.
|
||||||
|
type Stream struct {
|
||||||
|
*stream
|
||||||
|
b Samples
|
||||||
|
sorted bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newStream(ƒ invariant) *Stream {
|
||||||
|
x := &stream{ƒ: ƒ}
|
||||||
|
return &Stream{x, make(Samples, 0, 500), true}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Insert inserts v into the stream.
|
||||||
|
func (s *Stream) Insert(v float64) {
|
||||||
|
s.insert(Sample{Value: v, Width: 1})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stream) insert(sample Sample) {
|
||||||
|
s.b = append(s.b, sample)
|
||||||
|
s.sorted = false
|
||||||
|
if len(s.b) == cap(s.b) {
|
||||||
|
s.flush()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Query returns the computed qth percentiles value. If s was created with
|
||||||
|
// NewTargeted, and q is not in the set of quantiles provided a priori, Query
|
||||||
|
// will return an unspecified result.
|
||||||
|
func (s *Stream) Query(q float64) float64 {
|
||||||
|
if !s.flushed() {
|
||||||
|
// Fast path when there hasn't been enough data for a flush;
|
||||||
|
// this also yields better accuracy for small sets of data.
|
||||||
|
l := len(s.b)
|
||||||
|
if l == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
i := int(math.Ceil(float64(l) * q))
|
||||||
|
if i > 0 {
|
||||||
|
i -= 1
|
||||||
|
}
|
||||||
|
s.maybeSort()
|
||||||
|
return s.b[i].Value
|
||||||
|
}
|
||||||
|
s.flush()
|
||||||
|
return s.stream.query(q)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Merge merges samples into the underlying streams samples. This is handy when
|
||||||
|
// merging multiple streams from separate threads, database shards, etc.
|
||||||
|
//
|
||||||
|
// ATTENTION: This method is broken and does not yield correct results. The
|
||||||
|
// underlying algorithm is not capable of merging streams correctly.
|
||||||
|
func (s *Stream) Merge(samples Samples) {
|
||||||
|
sort.Sort(samples)
|
||||||
|
s.stream.merge(samples)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset reinitializes and clears the list reusing the samples buffer memory.
|
||||||
|
func (s *Stream) Reset() {
|
||||||
|
s.stream.reset()
|
||||||
|
s.b = s.b[:0]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Samples returns stream samples held by s.
|
||||||
|
func (s *Stream) Samples() Samples {
|
||||||
|
if !s.flushed() {
|
||||||
|
return s.b
|
||||||
|
}
|
||||||
|
s.flush()
|
||||||
|
return s.stream.samples()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Count returns the total number of samples observed in the stream
|
||||||
|
// since initialization.
|
||||||
|
func (s *Stream) Count() int {
|
||||||
|
return len(s.b) + s.stream.count()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stream) flush() {
|
||||||
|
s.maybeSort()
|
||||||
|
s.stream.merge(s.b)
|
||||||
|
s.b = s.b[:0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stream) maybeSort() {
|
||||||
|
if !s.sorted {
|
||||||
|
s.sorted = true
|
||||||
|
sort.Sort(s.b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Stream) flushed() bool {
|
||||||
|
return len(s.stream.l) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
type stream struct {
|
||||||
|
n float64
|
||||||
|
l []Sample
|
||||||
|
ƒ invariant
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) reset() {
|
||||||
|
s.l = s.l[:0]
|
||||||
|
s.n = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) insert(v float64) {
|
||||||
|
s.merge(Samples{{v, 1, 0}})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) merge(samples Samples) {
|
||||||
|
// TODO(beorn7): This tries to merge not only individual samples, but
|
||||||
|
// whole summaries. The paper doesn't mention merging summaries at
|
||||||
|
// all. Unittests show that the merging is inaccurate. Find out how to
|
||||||
|
// do merges properly.
|
||||||
|
var r float64
|
||||||
|
i := 0
|
||||||
|
for _, sample := range samples {
|
||||||
|
for ; i < len(s.l); i++ {
|
||||||
|
c := s.l[i]
|
||||||
|
if c.Value > sample.Value {
|
||||||
|
// Insert at position i.
|
||||||
|
s.l = append(s.l, Sample{})
|
||||||
|
copy(s.l[i+1:], s.l[i:])
|
||||||
|
s.l[i] = Sample{
|
||||||
|
sample.Value,
|
||||||
|
sample.Width,
|
||||||
|
math.Max(sample.Delta, math.Floor(s.ƒ(s, r))-1),
|
||||||
|
// TODO(beorn7): How to calculate delta correctly?
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
goto inserted
|
||||||
|
}
|
||||||
|
r += c.Width
|
||||||
|
}
|
||||||
|
s.l = append(s.l, Sample{sample.Value, sample.Width, 0})
|
||||||
|
i++
|
||||||
|
inserted:
|
||||||
|
s.n += sample.Width
|
||||||
|
r += sample.Width
|
||||||
|
}
|
||||||
|
s.compress()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) count() int {
|
||||||
|
return int(s.n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) query(q float64) float64 {
|
||||||
|
t := math.Ceil(q * s.n)
|
||||||
|
t += math.Ceil(s.ƒ(s, t) / 2)
|
||||||
|
p := s.l[0]
|
||||||
|
var r float64
|
||||||
|
for _, c := range s.l[1:] {
|
||||||
|
r += p.Width
|
||||||
|
if r+c.Width+c.Delta > t {
|
||||||
|
return p.Value
|
||||||
|
}
|
||||||
|
p = c
|
||||||
|
}
|
||||||
|
return p.Value
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) compress() {
|
||||||
|
if len(s.l) < 2 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
x := s.l[len(s.l)-1]
|
||||||
|
xi := len(s.l) - 1
|
||||||
|
r := s.n - 1 - x.Width
|
||||||
|
|
||||||
|
for i := len(s.l) - 2; i >= 0; i-- {
|
||||||
|
c := s.l[i]
|
||||||
|
if c.Width+x.Width+x.Delta <= s.ƒ(s, r) {
|
||||||
|
x.Width += c.Width
|
||||||
|
s.l[xi] = x
|
||||||
|
// Remove element at i.
|
||||||
|
copy(s.l[i:], s.l[i+1:])
|
||||||
|
s.l = s.l[:len(s.l)-1]
|
||||||
|
xi -= 1
|
||||||
|
} else {
|
||||||
|
x = c
|
||||||
|
xi = i
|
||||||
|
}
|
||||||
|
r -= c.Width
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stream) samples() Samples {
|
||||||
|
samples := make(Samples, len(s.l))
|
||||||
|
copy(samples, s.l)
|
||||||
|
return samples
|
||||||
|
}
|
||||||
+88
@@ -0,0 +1,88 @@
|
|||||||
|
Project Notice
|
||||||
|
|
||||||
|
This project was independently developed.
|
||||||
|
|
||||||
|
The contents of this repository were prepared and published outside any time
|
||||||
|
allocated to Bitech Systems CC and do not contain, incorporate, disclose,
|
||||||
|
or rely upon any proprietary or confidential information, trade secrets,
|
||||||
|
protected designs, or other intellectual property of Bitech Systems CC.
|
||||||
|
|
||||||
|
No portion of this repository reproduces any Bitech Systems CC-specific
|
||||||
|
implementation, design asset, confidential workflow, or non-public technical material.
|
||||||
|
|
||||||
|
This notice is provided for clarification only and does not modify the terms of
|
||||||
|
the Apache License, Version 2.0.
|
||||||
|
|
||||||
|
Apache License
|
||||||
|
Version 2.0, January 2004
|
||||||
|
http://www.apache.org/licenses/
|
||||||
|
|
||||||
|
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||||
|
|
||||||
|
1. Definitions.
|
||||||
|
|
||||||
|
"License" shall mean the terms and conditions for use, reproduction, and distribution as defined by Sections 1 through 9 of this document.
|
||||||
|
|
||||||
|
"Licensor" shall mean the copyright owner or entity authorized by the copyright owner that is granting the License.
|
||||||
|
|
||||||
|
"Legal Entity" shall mean the union of the acting entity and all other entities that control, are controlled by, or are under common control with that entity. For the purposes of this definition, "control" means (i) the power, direct or indirect, to cause the direction or management of such entity, whether by contract or otherwise, or (ii) ownership of fifty percent (50%) or more of the outstanding shares, or (iii) beneficial ownership of such entity.
|
||||||
|
|
||||||
|
"You" (or "Your") shall mean an individual or Legal Entity exercising permissions granted by this License.
|
||||||
|
|
||||||
|
"Source" form shall mean the preferred form for making modifications, including but not limited to software source code, documentation source, and configuration files.
|
||||||
|
|
||||||
|
"Object" form shall mean any form resulting from mechanical transformation or translation of a Source form, including but not limited to compiled object code, generated documentation, and conversions to other media types.
|
||||||
|
|
||||||
|
"Work" shall mean the work of authorship, whether in Source or Object form, made available under the License, as indicated by a copyright notice that is included in or attached to the work (an example is provided in the Appendix below).
|
||||||
|
|
||||||
|
"Derivative Works" shall mean any work, whether in Source or Object form, that is based on (or derived from) the Work and for which the editorial revisions, annotations, elaborations, or other modifications represent, as a whole, an original work of authorship. For the purposes of this License, Derivative Works shall not include works that remain separable from, or merely link (or bind by name) to the interfaces of, the Work and Derivative Works thereof.
|
||||||
|
|
||||||
|
"Contribution" shall mean any work of authorship, including the original version of the Work and any modifications or additions to that Work or Derivative Works thereof, that is intentionally submitted to Licensor for inclusion in the Work by the copyright owner or by an individual or Legal Entity authorized to submit on behalf of the copyright owner. For the purposes of this definition, "submitted" means any form of electronic, verbal, or written communication sent to the Licensor or its representatives, including but not limited to communication on electronic mailing lists, source code control systems, and issue tracking systems that are managed by, or on behalf of, the Licensor for the purpose of discussing and improving the Work, but excluding communication that is conspicuously marked or otherwise designated in writing by the copyright owner as "Not a Contribution."
|
||||||
|
|
||||||
|
"Contributor" shall mean Licensor and any individual or Legal Entity on behalf of whom a Contribution has been received by Licensor and subsequently incorporated within the Work.
|
||||||
|
|
||||||
|
2. Grant of Copyright License. Subject to the terms and conditions of this License, each Contributor hereby grants to You a perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable copyright license to reproduce, prepare Derivative Works of, publicly display, publicly perform, sublicense, and distribute the Work and such Derivative Works in Source or Object form.
|
||||||
|
|
||||||
|
3. Grant of Patent License. Subject to the terms and conditions of this License, each Contributor hereby grants to You a perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable (except as stated in this section) patent license to make, have made, use, offer to sell, sell, import, and otherwise transfer the Work, where such license applies only to those patent claims licensable by such Contributor that are necessarily infringed by their Contribution(s) alone or by combination of their Contribution(s) with the Work to which such Contribution(s) was submitted. If You institute patent litigation against any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the Work or a Contribution incorporated within the Work constitutes direct or contributory patent infringement, then any patent licenses granted to You under this License for that Work shall terminate as of the date such litigation is filed.
|
||||||
|
|
||||||
|
4. Redistribution. You may reproduce and distribute copies of the Work or Derivative Works thereof in any medium, with or without modifications, and in Source or Object form, provided that You meet the following conditions:
|
||||||
|
|
||||||
|
(a) You must give any other recipients of the Work or Derivative Works a copy of this License; and
|
||||||
|
|
||||||
|
(b) You must cause any modified files to carry prominent notices stating that You changed the files; and
|
||||||
|
|
||||||
|
(c) You must retain, in the Source form of any Derivative Works that You distribute, all copyright, patent, trademark, and attribution notices from the Source form of the Work, excluding those notices that do not pertain to any part of the Derivative Works; and
|
||||||
|
|
||||||
|
(d) If the Work includes a "NOTICE" text file as part of its distribution, then any Derivative Works that You distribute must include a readable copy of the attribution notices contained within such NOTICE file, excluding those notices that do not pertain to any part of the Derivative Works, in at least one of the following places: within a NOTICE text file distributed as part of the Derivative Works; within the Source form or documentation, if provided along with the Derivative Works; or, within a display generated by the Derivative Works, if and wherever such third-party notices normally appear. The contents of the NOTICE file are for informational purposes only and do not modify the License. You may add Your own attribution notices within Derivative Works that You distribute, alongside or as an addendum to the NOTICE text from the Work, provided that such additional attribution notices cannot be construed as modifying the License.
|
||||||
|
|
||||||
|
You may add Your own copyright statement to Your modifications and may provide additional or different license terms and conditions for use, reproduction, or distribution of Your modifications, or for any such Derivative Works as a whole, provided Your use, reproduction, and distribution of the Work otherwise complies with the conditions stated in this License.
|
||||||
|
|
||||||
|
5. Submission of Contributions. Unless You explicitly state otherwise, any Contribution intentionally submitted for inclusion in the Work by You to the Licensor shall be under the terms and conditions of this License, without any additional terms or conditions. Notwithstanding the above, nothing herein shall supersede or modify the terms of any separate license agreement you may have executed with Licensor regarding such Contributions.
|
||||||
|
|
||||||
|
6. Trademarks. This License does not grant permission to use the trade names, trademarks, service marks, or product names of the Licensor, except as required for reasonable and customary use in describing the origin of the Work and reproducing the content of the NOTICE file.
|
||||||
|
|
||||||
|
7. Disclaimer of Warranty. Unless required by applicable law or agreed to in writing, Licensor provides the Work (and each Contributor provides its Contributions) on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied, including, without limitation, any warranties or conditions of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A PARTICULAR PURPOSE. You are solely responsible for determining the appropriateness of using or redistributing the Work and assume any risks associated with Your exercise of permissions under this License.
|
||||||
|
|
||||||
|
8. Limitation of Liability. In no event and under no legal theory, whether in tort (including negligence), contract, or otherwise, unless required by applicable law (such as deliberate and grossly negligent acts) or agreed to in writing, shall any Contributor be liable to You for damages, including any direct, indirect, special, incidental, or consequential damages of any character arising as a result of this License or out of the use or inability to use the Work (including but not limited to damages for loss of goodwill, work stoppage, computer failure or malfunction, or any and all other commercial damages or losses), even if such Contributor has been advised of the possibility of such damages.
|
||||||
|
|
||||||
|
9. Accepting Warranty or Additional Liability. While redistributing the Work or Derivative Works thereof, You may choose to offer, and charge a fee for, acceptance of support, warranty, indemnity, or other liability obligations and/or rights consistent with this License. However, in accepting such obligations, You may act only on Your own behalf and on Your sole responsibility, not on behalf of any other Contributor, and only if You agree to indemnify, defend, and hold each Contributor harmless for any liability incurred by, or claims asserted against, such Contributor by reason of your accepting any such warranty or additional liability.
|
||||||
|
|
||||||
|
END OF TERMS AND CONDITIONS
|
||||||
|
|
||||||
|
APPENDIX: How to apply the Apache License to your work.
|
||||||
|
|
||||||
|
To apply the Apache License to your work, attach the following boilerplate notice, with the fields enclosed by brackets "[]" replaced with your own identifying information. (Don't include the brackets!) The text should be enclosed in the appropriate comment syntax for the file format. We also recommend that a file or class name and description of purpose be included on the same "printed page" as the copyright notice for easier identification within third-party archives.
|
||||||
|
|
||||||
|
Copyright 2025 wdevs
|
||||||
|
|
||||||
|
Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
you may not use this file except in compliance with the License.
|
||||||
|
You may obtain a copy of the License at
|
||||||
|
|
||||||
|
http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
|
||||||
|
Unless required by applicable law or agreed to in writing, software
|
||||||
|
distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
See the License for the specific language governing permissions and
|
||||||
|
limitations under the License.
|
||||||
+340
@@ -0,0 +1,340 @@
|
|||||||
|
# Cache Package
|
||||||
|
|
||||||
|
A flexible, provider-based caching library for Go that supports multiple backend storage systems including in-memory, Redis, and Memcache.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- **Multiple Providers**: Support for in-memory, Redis, and Memcache backends
|
||||||
|
- **Pluggable Architecture**: Easy to add custom cache providers
|
||||||
|
- **Type-Safe API**: Automatic JSON serialization/deserialization
|
||||||
|
- **TTL Support**: Configurable time-to-live for cache entries
|
||||||
|
- **Context-Aware**: All operations support Go contexts
|
||||||
|
- **Statistics**: Built-in cache statistics and monitoring
|
||||||
|
- **Pattern Deletion**: Delete keys by pattern (Redis)
|
||||||
|
- **Lazy Loading**: GetOrSet pattern for easy cache-aside implementation
|
||||||
|
|
||||||
|
## Installation
|
||||||
|
|
||||||
|
```bash
|
||||||
|
go get github.com/bitechdev/ResolveSpec/pkg/cache
|
||||||
|
```
|
||||||
|
|
||||||
|
For Redis support:
|
||||||
|
```bash
|
||||||
|
go get github.com/redis/go-redis/v9
|
||||||
|
```
|
||||||
|
|
||||||
|
For Memcache support:
|
||||||
|
```bash
|
||||||
|
go get github.com/bradfitz/gomemcache/memcache
|
||||||
|
```
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
### In-Memory Cache
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Initialize with in-memory provider
|
||||||
|
cache.UseMemory(&cache.Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
MaxSize: 10000,
|
||||||
|
})
|
||||||
|
defer cache.Close()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
c := cache.GetDefaultCache()
|
||||||
|
|
||||||
|
// Store a value
|
||||||
|
type User struct {
|
||||||
|
ID int
|
||||||
|
Name string
|
||||||
|
}
|
||||||
|
user := User{ID: 1, Name: "John"}
|
||||||
|
c.Set(ctx, "user:1", user, 10*time.Minute)
|
||||||
|
|
||||||
|
// Retrieve a value
|
||||||
|
var retrieved User
|
||||||
|
c.Get(ctx, "user:1", &retrieved)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Redis Cache
|
||||||
|
|
||||||
|
```go
|
||||||
|
cache.UseRedis(&cache.RedisConfig{
|
||||||
|
Host: "localhost",
|
||||||
|
Port: 6379,
|
||||||
|
Password: "",
|
||||||
|
DB: 0,
|
||||||
|
Options: &cache.Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
defer cache.Close()
|
||||||
|
```
|
||||||
|
|
||||||
|
### Memcache
|
||||||
|
|
||||||
|
```go
|
||||||
|
cache.UseMemcache(&cache.MemcacheConfig{
|
||||||
|
Servers: []string{"localhost:11211"},
|
||||||
|
Options: &cache.Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
defer cache.Close()
|
||||||
|
```
|
||||||
|
|
||||||
|
## API Reference
|
||||||
|
|
||||||
|
### Core Methods
|
||||||
|
|
||||||
|
#### Set
|
||||||
|
```go
|
||||||
|
Set(ctx context.Context, key string, value interface{}, ttl time.Duration) error
|
||||||
|
```
|
||||||
|
Stores a value in the cache with automatic JSON serialization.
|
||||||
|
|
||||||
|
#### Get
|
||||||
|
```go
|
||||||
|
Get(ctx context.Context, key string, dest interface{}) error
|
||||||
|
```
|
||||||
|
Retrieves and deserializes a value from the cache.
|
||||||
|
|
||||||
|
#### SetBytes / GetBytes
|
||||||
|
```go
|
||||||
|
SetBytes(ctx context.Context, key string, value []byte, ttl time.Duration) error
|
||||||
|
GetBytes(ctx context.Context, key string) ([]byte, error)
|
||||||
|
```
|
||||||
|
Store and retrieve raw bytes without serialization.
|
||||||
|
|
||||||
|
#### Delete
|
||||||
|
```go
|
||||||
|
Delete(ctx context.Context, key string) error
|
||||||
|
```
|
||||||
|
Removes a key from the cache.
|
||||||
|
|
||||||
|
#### DeleteByPattern
|
||||||
|
```go
|
||||||
|
DeleteByPattern(ctx context.Context, pattern string) error
|
||||||
|
```
|
||||||
|
Removes all keys matching a pattern (Redis only).
|
||||||
|
|
||||||
|
#### Clear
|
||||||
|
```go
|
||||||
|
Clear(ctx context.Context) error
|
||||||
|
```
|
||||||
|
Removes all items from the cache.
|
||||||
|
|
||||||
|
#### Exists
|
||||||
|
```go
|
||||||
|
Exists(ctx context.Context, key string) bool
|
||||||
|
```
|
||||||
|
Checks if a key exists in the cache.
|
||||||
|
|
||||||
|
#### GetOrSet
|
||||||
|
```go
|
||||||
|
GetOrSet(ctx context.Context, key string, dest interface{}, ttl time.Duration,
|
||||||
|
loader func() (interface{}, error)) error
|
||||||
|
```
|
||||||
|
Retrieves a value from cache, or loads and caches it if not found (lazy loading).
|
||||||
|
|
||||||
|
#### Stats
|
||||||
|
```go
|
||||||
|
Stats(ctx context.Context) (*CacheStats, error)
|
||||||
|
```
|
||||||
|
Returns cache statistics including hits, misses, and key counts.
|
||||||
|
|
||||||
|
## Provider Configuration
|
||||||
|
|
||||||
|
### In-Memory Options
|
||||||
|
|
||||||
|
```go
|
||||||
|
&cache.Options{
|
||||||
|
DefaultTTL: 5 * time.Minute, // Default expiration time
|
||||||
|
MaxSize: 10000, // Maximum number of items
|
||||||
|
EvictionPolicy: "LRU", // Eviction strategy (future)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Redis Configuration
|
||||||
|
|
||||||
|
```go
|
||||||
|
&cache.RedisConfig{
|
||||||
|
Host: "localhost",
|
||||||
|
Port: 6379,
|
||||||
|
Password: "", // Optional authentication
|
||||||
|
DB: 0, // Database number
|
||||||
|
PoolSize: 10, // Connection pool size
|
||||||
|
Options: &cache.Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Memcache Configuration
|
||||||
|
|
||||||
|
```go
|
||||||
|
&cache.MemcacheConfig{
|
||||||
|
Servers: []string{"localhost:11211"},
|
||||||
|
MaxIdleConns: 2,
|
||||||
|
Timeout: 1 * time.Second,
|
||||||
|
Options: &cache.Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Advanced Usage
|
||||||
|
|
||||||
|
### Custom Provider
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Create a custom provider instance
|
||||||
|
memProvider := cache.NewMemoryProvider(&cache.Options{
|
||||||
|
DefaultTTL: 10 * time.Minute,
|
||||||
|
MaxSize: 500,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Initialize with custom provider
|
||||||
|
cache.Initialize(memProvider)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Lazy Loading Pattern
|
||||||
|
|
||||||
|
```go
|
||||||
|
var data ExpensiveData
|
||||||
|
err := c.GetOrSet(ctx, "expensive:key", &data, 10*time.Minute, func() (interface{}, error) {
|
||||||
|
// This expensive operation only runs if key is not in cache
|
||||||
|
return computeExpensiveData(), nil
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### Query API Cache
|
||||||
|
|
||||||
|
The package includes specialized functions for caching query results:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Cache a query result
|
||||||
|
api := "GetUsers"
|
||||||
|
query := "SELECT * FROM users WHERE active = true"
|
||||||
|
tablenames := "users"
|
||||||
|
total := int64(150)
|
||||||
|
|
||||||
|
cache.PutQueryAPICache(ctx, api, query, tablenames, total)
|
||||||
|
|
||||||
|
// Retrieve cached query
|
||||||
|
hash := cache.HashQueryAPICache(api, query)
|
||||||
|
cachedQuery, err := cache.FetchQueryAPICache(ctx, hash)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Provider Comparison
|
||||||
|
|
||||||
|
| Feature | In-Memory | Redis | Memcache |
|
||||||
|
|---------|-----------|-------|----------|
|
||||||
|
| Persistence | No | Yes | No |
|
||||||
|
| Distributed | No | Yes | Yes |
|
||||||
|
| Pattern Delete | No | Yes | No |
|
||||||
|
| Statistics | Full | Full | Limited |
|
||||||
|
| Atomic Operations | Yes | Yes | Yes |
|
||||||
|
| Max Item Size | Memory | 512MB | 1MB |
|
||||||
|
|
||||||
|
## Best Practices
|
||||||
|
|
||||||
|
1. **Use contexts**: Always pass context for cancellation and timeout control
|
||||||
|
2. **Set appropriate TTLs**: Balance between freshness and performance
|
||||||
|
3. **Handle errors**: Cache misses and errors should be handled gracefully
|
||||||
|
4. **Monitor statistics**: Use Stats() to monitor cache performance
|
||||||
|
5. **Clean up**: Always call Close() when shutting down
|
||||||
|
6. **Pattern consistency**: Use consistent key naming patterns (e.g., "user:id:field")
|
||||||
|
|
||||||
|
## Example: Complete Application
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log"
|
||||||
|
"time"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
|
)
|
||||||
|
|
||||||
|
type UserService struct {
|
||||||
|
cache *cache.Cache
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewUserService() *UserService {
|
||||||
|
// Initialize with Redis in production, memory for testing
|
||||||
|
cache.UseRedis(&cache.RedisConfig{
|
||||||
|
Host: "localhost",
|
||||||
|
Port: 6379,
|
||||||
|
Options: &cache.Options{
|
||||||
|
DefaultTTL: 10 * time.Minute,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
return &UserService{
|
||||||
|
cache: cache.GetDefaultCache(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *UserService) GetUser(ctx context.Context, userID int) (*User, error) {
|
||||||
|
var user User
|
||||||
|
cacheKey := fmt.Sprintf("user:%d", userID)
|
||||||
|
|
||||||
|
// Try to get from cache first
|
||||||
|
err := s.cache.GetOrSet(ctx, cacheKey, &user, 15*time.Minute, func() (interface{}, error) {
|
||||||
|
// Load from database if not in cache
|
||||||
|
return s.loadUserFromDB(userID)
|
||||||
|
})
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *UserService) InvalidateUser(ctx context.Context, userID int) error {
|
||||||
|
cacheKey := fmt.Sprintf("user:%d", userID)
|
||||||
|
return s.cache.Delete(ctx, cacheKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
service := NewUserService()
|
||||||
|
defer cache.Close()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
user, err := service.GetUser(ctx, 123)
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Printf("User: %+v", user)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Performance Considerations
|
||||||
|
|
||||||
|
- **In-Memory**: Fastest but limited by RAM and not distributed
|
||||||
|
- **Redis**: Great for distributed systems, persistent, but network overhead
|
||||||
|
- **Memcache**: Good for distributed caching, simpler than Redis but less features
|
||||||
|
|
||||||
|
Choose based on your needs:
|
||||||
|
- Single instance? Use in-memory
|
||||||
|
- Need persistence or advanced features? Use Redis
|
||||||
|
- Simple distributed cache? Use Memcache
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
See repository license.
|
||||||
+76
@@ -0,0 +1,76 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
defaultCache *Cache
|
||||||
|
)
|
||||||
|
|
||||||
|
// Initialize initializes the cache with a provider.
|
||||||
|
// If not called, the package will use an in-memory provider by default.
|
||||||
|
func Initialize(provider Provider) {
|
||||||
|
defaultCache = NewCache(provider)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UseMemory configures the cache to use in-memory storage.
|
||||||
|
func UseMemory(opts *Options) error {
|
||||||
|
provider := NewMemoryProvider(opts)
|
||||||
|
defaultCache = NewCache(provider)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UseRedis configures the cache to use Redis storage.
|
||||||
|
func UseRedis(config *RedisConfig) error {
|
||||||
|
provider, err := NewRedisProvider(config)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to initialize Redis provider: %w", err)
|
||||||
|
}
|
||||||
|
defaultCache = NewCache(provider)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UseMemcache configures the cache to use Memcache storage.
|
||||||
|
func UseMemcache(config *MemcacheConfig) error {
|
||||||
|
provider, err := NewMemcacheProvider(config)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to initialize Memcache provider: %w", err)
|
||||||
|
}
|
||||||
|
defaultCache = NewCache(provider)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDefaultCache returns the default cache instance.
|
||||||
|
// Initializes with in-memory provider if not already initialized.
|
||||||
|
func GetDefaultCache() *Cache {
|
||||||
|
if defaultCache == nil {
|
||||||
|
_ = UseMemory(&Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
MaxSize: 10000,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return defaultCache
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetDefaultCache sets a custom cache instance as the default cache.
|
||||||
|
// This is useful for testing or when you want to use a pre-configured cache instance.
|
||||||
|
func SetDefaultCache(cache *Cache) {
|
||||||
|
defaultCache = cache
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetStats returns cache statistics.
|
||||||
|
func GetStats(ctx context.Context) (*CacheStats, error) {
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
return cache.Stats(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes the cache and releases resources.
|
||||||
|
func Close() error {
|
||||||
|
if defaultCache != nil {
|
||||||
|
return defaultCache.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+167
@@ -0,0 +1,167 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Cache is the main cache manager that wraps a Provider.
|
||||||
|
type Cache struct {
|
||||||
|
provider Provider
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCache creates a new cache manager with the specified provider.
|
||||||
|
func NewCache(provider Provider) *Cache {
|
||||||
|
return &Cache{
|
||||||
|
provider: provider,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get retrieves and deserializes a value from the cache.
|
||||||
|
func (c *Cache) Get(ctx context.Context, key string, dest interface{}) error {
|
||||||
|
data, exists := c.provider.Get(ctx, key)
|
||||||
|
if !exists {
|
||||||
|
return fmt.Errorf("key not found: %s", key)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(data, dest); err != nil {
|
||||||
|
return fmt.Errorf("failed to deserialize: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetBytes retrieves raw bytes from the cache.
|
||||||
|
func (c *Cache) GetBytes(ctx context.Context, key string) ([]byte, error) {
|
||||||
|
data, exists := c.provider.Get(ctx, key)
|
||||||
|
if !exists {
|
||||||
|
return nil, fmt.Errorf("key not found: %s", key)
|
||||||
|
}
|
||||||
|
return data, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set serializes and stores a value in the cache with the specified TTL.
|
||||||
|
func (c *Cache) Set(ctx context.Context, key string, value interface{}, ttl time.Duration) error {
|
||||||
|
data, err := json.Marshal(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to serialize: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return c.provider.Set(ctx, key, data, ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetBytes stores raw bytes in the cache with the specified TTL.
|
||||||
|
func (c *Cache) SetBytes(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||||
|
return c.provider.Set(ctx, key, value, ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetWithTags serializes and stores a value in the cache with the specified TTL and tags.
|
||||||
|
func (c *Cache) SetWithTags(ctx context.Context, key string, value interface{}, ttl time.Duration, tags []string) error {
|
||||||
|
data, err := json.Marshal(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to serialize: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return c.provider.SetWithTags(ctx, key, data, ttl, tags)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetBytesWithTags stores raw bytes in the cache with the specified TTL and tags.
|
||||||
|
func (c *Cache) SetBytesWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
|
||||||
|
return c.provider.SetWithTags(ctx, key, value, ttl, tags)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete removes a key from the cache.
|
||||||
|
func (c *Cache) Delete(ctx context.Context, key string) error {
|
||||||
|
return c.provider.Delete(ctx, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByTag removes all keys associated with the given tag.
|
||||||
|
func (c *Cache) DeleteByTag(ctx context.Context, tag string) error {
|
||||||
|
return c.provider.DeleteByTag(ctx, tag)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByPattern removes all keys matching the pattern.
|
||||||
|
func (c *Cache) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||||
|
return c.provider.DeleteByPattern(ctx, pattern)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear removes all items from the cache.
|
||||||
|
func (c *Cache) Clear(ctx context.Context) error {
|
||||||
|
return c.provider.Clear(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exists checks if a key exists in the cache.
|
||||||
|
func (c *Cache) Exists(ctx context.Context, key string) bool {
|
||||||
|
return c.provider.Exists(ctx, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stats returns statistics about the cache.
|
||||||
|
func (c *Cache) Stats(ctx context.Context) (*CacheStats, error) {
|
||||||
|
return c.provider.Stats(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes the cache and releases any resources.
|
||||||
|
func (c *Cache) Close() error {
|
||||||
|
return c.provider.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetOrSet retrieves a value from cache, or sets it if it doesn't exist.
|
||||||
|
// The loader function is called only if the key is not found in cache.
|
||||||
|
func (c *Cache) GetOrSet(ctx context.Context, key string, dest interface{}, ttl time.Duration, loader func() (interface{}, error)) error {
|
||||||
|
// Try to get from cache first
|
||||||
|
err := c.Get(ctx, key, dest)
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load the value
|
||||||
|
value, err := loader()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("loader failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store in cache
|
||||||
|
if err := c.Set(ctx, key, value, ttl); err != nil {
|
||||||
|
return fmt.Errorf("failed to cache value: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Populate dest with the loaded value
|
||||||
|
data, err := json.Marshal(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to serialize loaded value: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(data, dest); err != nil {
|
||||||
|
return fmt.Errorf("failed to deserialize loaded value: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remember is a convenience function that caches the result of a function call.
|
||||||
|
// It's similar to GetOrSet but returns the value directly.
|
||||||
|
func (c *Cache) Remember(ctx context.Context, key string, ttl time.Duration, loader func() (interface{}, error)) (interface{}, error) {
|
||||||
|
// Try to get from cache first as bytes
|
||||||
|
data, err := c.GetBytes(ctx, key)
|
||||||
|
if err == nil {
|
||||||
|
var result interface{}
|
||||||
|
if err := json.Unmarshal(data, &result); err == nil {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load the value
|
||||||
|
value, err := loader()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("loader failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store in cache
|
||||||
|
if err := c.Set(ctx, key, value, ttl); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to cache value: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return value, nil
|
||||||
|
}
|
||||||
+266
@@ -0,0 +1,266 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ExampleInMemoryCache demonstrates using the in-memory cache provider.
|
||||||
|
func ExampleInMemoryCache() {
|
||||||
|
// Initialize with in-memory provider
|
||||||
|
err := UseMemory(&Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
MaxSize: 1000,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Get the cache instance
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
|
||||||
|
// Store a value
|
||||||
|
type User struct {
|
||||||
|
ID int
|
||||||
|
Name string
|
||||||
|
}
|
||||||
|
|
||||||
|
user := User{ID: 1, Name: "John Doe"}
|
||||||
|
err = cache.Set(ctx, "user:1", user, 10*time.Minute)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Retrieve a value
|
||||||
|
var retrieved User
|
||||||
|
err = cache.Get(ctx, "user:1", &retrieved)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Retrieved user: %+v\n", retrieved)
|
||||||
|
|
||||||
|
// Check if key exists
|
||||||
|
exists := cache.Exists(ctx, "user:1")
|
||||||
|
fmt.Printf("Key exists: %v\n", exists)
|
||||||
|
|
||||||
|
// Delete a key
|
||||||
|
err = cache.Delete(ctx, "user:1")
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get statistics
|
||||||
|
stats, err := cache.Stats(ctx)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
fmt.Printf("Cache stats: %+v\n", stats)
|
||||||
|
_ = Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleRedisCache demonstrates using the Redis cache provider.
|
||||||
|
func ExampleRedisCache() {
|
||||||
|
// Initialize with Redis provider
|
||||||
|
err := UseRedis(&RedisConfig{
|
||||||
|
Host: "localhost",
|
||||||
|
Port: 6379,
|
||||||
|
Password: "", // Set if Redis requires authentication
|
||||||
|
DB: 0,
|
||||||
|
Options: &Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Get the cache instance
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
|
||||||
|
// Store raw bytes
|
||||||
|
data := []byte("Hello, Redis!")
|
||||||
|
err = cache.SetBytes(ctx, "greeting", data, 1*time.Hour)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Retrieve raw bytes
|
||||||
|
retrieved, err := cache.GetBytes(ctx, "greeting")
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Retrieved data: %s\n", string(retrieved))
|
||||||
|
|
||||||
|
// Clear all cache
|
||||||
|
err = cache.Clear(ctx)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleMemcacheCache demonstrates using the Memcache cache provider.
|
||||||
|
func ExampleMemcacheCache() {
|
||||||
|
// Initialize with Memcache provider
|
||||||
|
err := UseMemcache(&MemcacheConfig{
|
||||||
|
Servers: []string{"localhost:11211"},
|
||||||
|
Options: &Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Get the cache instance
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
|
||||||
|
// Store a value
|
||||||
|
type Product struct {
|
||||||
|
ID int
|
||||||
|
Name string
|
||||||
|
Price float64
|
||||||
|
}
|
||||||
|
|
||||||
|
product := Product{ID: 100, Name: "Widget", Price: 29.99}
|
||||||
|
err = cache.Set(ctx, "product:100", product, 30*time.Minute)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Retrieve a value
|
||||||
|
var retrieved Product
|
||||||
|
err = cache.Get(ctx, "product:100", &retrieved)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Retrieved product: %+v\n", retrieved)
|
||||||
|
_ = Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleGetOrSet demonstrates the GetOrSet pattern for lazy loading.
|
||||||
|
func ExampleGetOrSet() {
|
||||||
|
err := UseMemory(&Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
MaxSize: 1000,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
|
||||||
|
type ExpensiveData struct {
|
||||||
|
Result string
|
||||||
|
}
|
||||||
|
|
||||||
|
var data ExpensiveData
|
||||||
|
err = cache.GetOrSet(ctx, "expensive:computation", &data, 10*time.Minute, func() (interface{}, error) {
|
||||||
|
// This expensive operation only runs if the key is not in cache
|
||||||
|
fmt.Println("Computing expensive result...")
|
||||||
|
time.Sleep(1 * time.Second)
|
||||||
|
return ExpensiveData{Result: "computed value"}, nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Data: %+v\n", data)
|
||||||
|
|
||||||
|
// Second call will use cached value
|
||||||
|
err = cache.GetOrSet(ctx, "expensive:computation", &data, 10*time.Minute, func() (interface{}, error) {
|
||||||
|
fmt.Println("This won't be called!")
|
||||||
|
return ExpensiveData{Result: "new value"}, nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Cached data: %+v\n", data)
|
||||||
|
_ = Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleCustomProvider demonstrates using a custom provider.
|
||||||
|
func ExampleCustomProvider() {
|
||||||
|
// Create a custom provider
|
||||||
|
memProvider := NewMemoryProvider(&Options{
|
||||||
|
DefaultTTL: 10 * time.Minute,
|
||||||
|
MaxSize: 500,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Initialize with custom provider
|
||||||
|
Initialize(memProvider)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
|
||||||
|
// Use the cache
|
||||||
|
err := cache.SetBytes(ctx, "key", []byte("value"), 5*time.Minute)
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clean expired items (memory provider specific)
|
||||||
|
if mp, ok := cache.provider.(*MemoryProvider); ok {
|
||||||
|
count := mp.CleanExpired(ctx)
|
||||||
|
fmt.Printf("Cleaned %d expired items\n", count)
|
||||||
|
}
|
||||||
|
_ = Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleDeleteByPattern demonstrates pattern-based deletion (Redis only).
|
||||||
|
func ExampleDeleteByPattern() {
|
||||||
|
err := UseRedis(&RedisConfig{
|
||||||
|
Host: "localhost",
|
||||||
|
Port: 6379,
|
||||||
|
Options: &Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
cache := GetDefaultCache()
|
||||||
|
|
||||||
|
// Store multiple keys with a pattern
|
||||||
|
_ = cache.SetBytes(ctx, "user:1:profile", []byte("profile1"), 10*time.Minute)
|
||||||
|
_ = cache.SetBytes(ctx, "user:2:profile", []byte("profile2"), 10*time.Minute)
|
||||||
|
_ = cache.SetBytes(ctx, "user:1:settings", []byte("settings1"), 10*time.Minute)
|
||||||
|
|
||||||
|
// Delete all keys matching pattern (Redis glob pattern)
|
||||||
|
err = cache.DeleteByPattern(ctx, "user:*:profile")
|
||||||
|
if err != nil {
|
||||||
|
_ = Close()
|
||||||
|
log.Print(err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Println("Deleted all user profile keys")
|
||||||
|
_ = Close()
|
||||||
|
}
|
||||||
+65
@@ -0,0 +1,65 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Provider defines the interface that all cache providers must implement.
|
||||||
|
type Provider interface {
|
||||||
|
// Get retrieves a value from the cache by key.
|
||||||
|
// Returns nil, false if key doesn't exist or is expired.
|
||||||
|
Get(ctx context.Context, key string) ([]byte, bool)
|
||||||
|
|
||||||
|
// Set stores a value in the cache with the specified TTL.
|
||||||
|
// If ttl is 0, the item never expires.
|
||||||
|
Set(ctx context.Context, key string, value []byte, ttl time.Duration) error
|
||||||
|
|
||||||
|
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
||||||
|
// Tags can be used to invalidate groups of related keys.
|
||||||
|
// If ttl is 0, the item never expires.
|
||||||
|
SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error
|
||||||
|
|
||||||
|
// Delete removes a key from the cache.
|
||||||
|
Delete(ctx context.Context, key string) error
|
||||||
|
|
||||||
|
// DeleteByTag removes all keys associated with the given tag.
|
||||||
|
DeleteByTag(ctx context.Context, tag string) error
|
||||||
|
|
||||||
|
// DeleteByPattern removes all keys matching the pattern.
|
||||||
|
// Pattern syntax depends on the provider implementation.
|
||||||
|
DeleteByPattern(ctx context.Context, pattern string) error
|
||||||
|
|
||||||
|
// Clear removes all items from the cache.
|
||||||
|
Clear(ctx context.Context) error
|
||||||
|
|
||||||
|
// Exists checks if a key exists in the cache.
|
||||||
|
Exists(ctx context.Context, key string) bool
|
||||||
|
|
||||||
|
// Close closes the provider and releases any resources.
|
||||||
|
Close() error
|
||||||
|
|
||||||
|
// Stats returns statistics about the cache provider.
|
||||||
|
Stats(ctx context.Context) (*CacheStats, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CacheStats contains cache statistics.
|
||||||
|
type CacheStats struct {
|
||||||
|
Hits int64 `json:"hits"`
|
||||||
|
Misses int64 `json:"misses"`
|
||||||
|
Keys int64 `json:"keys"`
|
||||||
|
ProviderType string `json:"provider_type"`
|
||||||
|
ProviderStats map[string]any `json:"provider_stats,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Options contains configuration options for cache providers.
|
||||||
|
type Options struct {
|
||||||
|
// DefaultTTL is the default time-to-live for cache items.
|
||||||
|
DefaultTTL time.Duration
|
||||||
|
|
||||||
|
// MaxSize is the maximum number of items (for in-memory provider).
|
||||||
|
MaxSize int
|
||||||
|
|
||||||
|
// EvictionPolicy determines how items are evicted (LRU, LFU, etc).
|
||||||
|
EvictionPolicy string
|
||||||
|
}
|
||||||
+284
@@ -0,0 +1,284 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bradfitz/gomemcache/memcache"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MemcacheProvider is a Memcache implementation of the Provider interface.
|
||||||
|
type MemcacheProvider struct {
|
||||||
|
client *memcache.Client
|
||||||
|
options *Options
|
||||||
|
}
|
||||||
|
|
||||||
|
// MemcacheConfig contains Memcache-specific configuration.
|
||||||
|
type MemcacheConfig struct {
|
||||||
|
// Servers is a list of memcache server addresses (e.g., "localhost:11211")
|
||||||
|
Servers []string
|
||||||
|
|
||||||
|
// MaxIdleConns is the maximum number of idle connections (default: 2)
|
||||||
|
MaxIdleConns int
|
||||||
|
|
||||||
|
// Timeout for connection operations (default: 1 second)
|
||||||
|
Timeout time.Duration
|
||||||
|
|
||||||
|
// Options contains general cache options
|
||||||
|
Options *Options
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMemcacheProvider creates a new Memcache cache provider.
|
||||||
|
func NewMemcacheProvider(config *MemcacheConfig) (*MemcacheProvider, error) {
|
||||||
|
if config == nil {
|
||||||
|
config = &MemcacheConfig{
|
||||||
|
Servers: []string{"localhost:11211"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(config.Servers) == 0 {
|
||||||
|
config.Servers = []string{"localhost:11211"}
|
||||||
|
}
|
||||||
|
|
||||||
|
if config.MaxIdleConns == 0 {
|
||||||
|
config.MaxIdleConns = 2
|
||||||
|
}
|
||||||
|
|
||||||
|
if config.Timeout == 0 {
|
||||||
|
config.Timeout = 1 * time.Second
|
||||||
|
}
|
||||||
|
|
||||||
|
if config.Options == nil {
|
||||||
|
config.Options = &Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
client := memcache.New(config.Servers...)
|
||||||
|
client.MaxIdleConns = config.MaxIdleConns
|
||||||
|
client.Timeout = config.Timeout
|
||||||
|
|
||||||
|
// Test connection
|
||||||
|
if err := client.Ping(); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to connect to Memcache: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &MemcacheProvider{
|
||||||
|
client: client,
|
||||||
|
options: config.Options,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get retrieves a value from the cache by key.
|
||||||
|
func (m *MemcacheProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||||
|
item, err := m.client.Get(key)
|
||||||
|
if err == memcache.ErrCacheMiss {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return item.Value, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set stores a value in the cache with the specified TTL.
|
||||||
|
func (m *MemcacheProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||||
|
if ttl == 0 {
|
||||||
|
ttl = m.options.DefaultTTL
|
||||||
|
}
|
||||||
|
|
||||||
|
item := &memcache.Item{
|
||||||
|
Key: key,
|
||||||
|
Value: value,
|
||||||
|
Expiration: int32(ttl.Seconds()),
|
||||||
|
}
|
||||||
|
|
||||||
|
return m.client.Set(item)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
||||||
|
// Note: Tag support in Memcache is limited and less efficient than Redis.
|
||||||
|
func (m *MemcacheProvider) SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
|
||||||
|
if ttl == 0 {
|
||||||
|
ttl = m.options.DefaultTTL
|
||||||
|
}
|
||||||
|
|
||||||
|
expiration := int32(ttl.Seconds())
|
||||||
|
|
||||||
|
// Set the main value
|
||||||
|
item := &memcache.Item{
|
||||||
|
Key: key,
|
||||||
|
Value: value,
|
||||||
|
Expiration: expiration,
|
||||||
|
}
|
||||||
|
if err := m.client.Set(item); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store tags for this key
|
||||||
|
if len(tags) > 0 {
|
||||||
|
tagsData, err := json.Marshal(tags)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal tags: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tagsItem := &memcache.Item{
|
||||||
|
Key: fmt.Sprintf("cache:tags:%s", key),
|
||||||
|
Value: tagsData,
|
||||||
|
Expiration: expiration,
|
||||||
|
}
|
||||||
|
if err := m.client.Set(tagsItem); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add key to each tag's key list
|
||||||
|
for _, tag := range tags {
|
||||||
|
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||||
|
|
||||||
|
// Get existing keys for this tag
|
||||||
|
var keys []string
|
||||||
|
if item, err := m.client.Get(tagKey); err == nil {
|
||||||
|
_ = json.Unmarshal(item.Value, &keys)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add current key if not already present
|
||||||
|
found := false
|
||||||
|
for _, k := range keys {
|
||||||
|
if k == key {
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
keys = append(keys, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store updated key list
|
||||||
|
keysData, err := json.Marshal(keys)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
tagItem := &memcache.Item{
|
||||||
|
Key: tagKey,
|
||||||
|
Value: keysData,
|
||||||
|
Expiration: expiration + 3600, // Give tag lists longer TTL
|
||||||
|
}
|
||||||
|
_ = m.client.Set(tagItem)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete removes a key from the cache.
|
||||||
|
func (m *MemcacheProvider) Delete(ctx context.Context, key string) error {
|
||||||
|
// Get tags for this key
|
||||||
|
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||||
|
if item, err := m.client.Get(tagsKey); err == nil {
|
||||||
|
var tags []string
|
||||||
|
if err := json.Unmarshal(item.Value, &tags); err == nil {
|
||||||
|
// Remove key from each tag's key list
|
||||||
|
for _, tag := range tags {
|
||||||
|
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||||
|
if tagItem, err := m.client.Get(tagKey); err == nil {
|
||||||
|
var keys []string
|
||||||
|
if err := json.Unmarshal(tagItem.Value, &keys); err == nil {
|
||||||
|
// Remove current key from the list
|
||||||
|
newKeys := make([]string, 0, len(keys))
|
||||||
|
for _, k := range keys {
|
||||||
|
if k != key {
|
||||||
|
newKeys = append(newKeys, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Update the tag's key list
|
||||||
|
if keysData, err := json.Marshal(newKeys); err == nil {
|
||||||
|
tagItem.Value = keysData
|
||||||
|
_ = m.client.Set(tagItem)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Delete the tags key
|
||||||
|
_ = m.client.Delete(tagsKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete the actual key
|
||||||
|
err := m.client.Delete(key)
|
||||||
|
if err == memcache.ErrCacheMiss {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByTag removes all keys associated with the given tag.
|
||||||
|
func (m *MemcacheProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||||
|
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||||
|
|
||||||
|
// Get all keys associated with this tag
|
||||||
|
item, err := m.client.Get(tagKey)
|
||||||
|
if err == memcache.ErrCacheMiss {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
var keys []string
|
||||||
|
if err := json.Unmarshal(item.Value, &keys); err != nil {
|
||||||
|
return fmt.Errorf("failed to unmarshal tag keys: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete all keys
|
||||||
|
for _, key := range keys {
|
||||||
|
_ = m.client.Delete(key)
|
||||||
|
// Also delete the tags key for this cache key
|
||||||
|
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||||
|
_ = m.client.Delete(tagsKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete the tag key itself
|
||||||
|
_ = m.client.Delete(tagKey)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByPattern removes all keys matching the pattern.
|
||||||
|
// Note: Memcache does not support pattern-based deletion natively.
|
||||||
|
// This is a no-op for memcache and returns an error.
|
||||||
|
func (m *MemcacheProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||||
|
return fmt.Errorf("pattern-based deletion is not supported by Memcache")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear removes all items from the cache.
|
||||||
|
func (m *MemcacheProvider) Clear(ctx context.Context) error {
|
||||||
|
return m.client.FlushAll()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exists checks if a key exists in the cache.
|
||||||
|
func (m *MemcacheProvider) Exists(ctx context.Context, key string) bool {
|
||||||
|
_, err := m.client.Get(key)
|
||||||
|
return err == nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes the provider and releases any resources.
|
||||||
|
func (m *MemcacheProvider) Close() error {
|
||||||
|
// Memcache client doesn't have a close method
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stats returns statistics about the cache provider.
|
||||||
|
// Note: Memcache provider returns limited statistics.
|
||||||
|
func (m *MemcacheProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||||
|
stats := &CacheStats{
|
||||||
|
ProviderType: "memcache",
|
||||||
|
ProviderStats: map[string]any{
|
||||||
|
"note": "Memcache does not provide detailed statistics through the standard client",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return stats, nil
|
||||||
|
}
|
||||||
+342
@@ -0,0 +1,342 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// memoryItem represents a cached item in memory.
|
||||||
|
type memoryItem struct {
|
||||||
|
Value []byte
|
||||||
|
Expiration time.Time
|
||||||
|
LastAccess time.Time
|
||||||
|
HitCount int64
|
||||||
|
Tags []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// isExpired checks if the item has expired.
|
||||||
|
func (m *memoryItem) isExpired() bool {
|
||||||
|
if m.Expiration.IsZero() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return time.Now().After(m.Expiration)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MemoryProvider is an in-memory implementation of the Provider interface.
|
||||||
|
type MemoryProvider struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
items map[string]*memoryItem
|
||||||
|
tagToKeys map[string]map[string]struct{} // tag -> set of keys
|
||||||
|
options *Options
|
||||||
|
hits atomic.Int64
|
||||||
|
misses atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMemoryProvider creates a new in-memory cache provider.
|
||||||
|
func NewMemoryProvider(opts *Options) *MemoryProvider {
|
||||||
|
if opts == nil {
|
||||||
|
opts = &Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
MaxSize: 10000,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &MemoryProvider{
|
||||||
|
items: make(map[string]*memoryItem),
|
||||||
|
tagToKeys: make(map[string]map[string]struct{}),
|
||||||
|
options: opts,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get retrieves a value from the cache by key.
|
||||||
|
func (m *MemoryProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||||
|
// First try with read lock for fast path
|
||||||
|
m.mu.RLock()
|
||||||
|
item, exists := m.items[key]
|
||||||
|
if !exists {
|
||||||
|
m.mu.RUnlock()
|
||||||
|
m.misses.Add(1)
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
if item.isExpired() {
|
||||||
|
m.mu.RUnlock()
|
||||||
|
// Upgrade to write lock to delete expired item
|
||||||
|
m.mu.Lock()
|
||||||
|
delete(m.items, key)
|
||||||
|
m.mu.Unlock()
|
||||||
|
m.misses.Add(1)
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update stats and access time with write lock
|
||||||
|
value := item.Value
|
||||||
|
m.mu.RUnlock()
|
||||||
|
|
||||||
|
// Update access tracking with write lock
|
||||||
|
m.mu.Lock()
|
||||||
|
item.LastAccess = time.Now()
|
||||||
|
item.HitCount++
|
||||||
|
m.mu.Unlock()
|
||||||
|
|
||||||
|
m.hits.Add(1)
|
||||||
|
return value, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set stores a value in the cache with the specified TTL.
|
||||||
|
func (m *MemoryProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
if ttl == 0 {
|
||||||
|
ttl = m.options.DefaultTTL
|
||||||
|
}
|
||||||
|
|
||||||
|
var expiration time.Time
|
||||||
|
if ttl > 0 {
|
||||||
|
expiration = time.Now().Add(ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check max size and evict if necessary
|
||||||
|
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
||||||
|
if _, exists := m.items[key]; !exists {
|
||||||
|
m.evictOne()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
m.items[key] = &memoryItem{
|
||||||
|
Value: value,
|
||||||
|
Expiration: expiration,
|
||||||
|
LastAccess: time.Now(),
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
||||||
|
func (m *MemoryProvider) SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
if ttl == 0 {
|
||||||
|
ttl = m.options.DefaultTTL
|
||||||
|
}
|
||||||
|
|
||||||
|
var expiration time.Time
|
||||||
|
if ttl > 0 {
|
||||||
|
expiration = time.Now().Add(ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check max size and evict if necessary
|
||||||
|
if m.options.MaxSize > 0 && len(m.items) >= m.options.MaxSize {
|
||||||
|
if _, exists := m.items[key]; !exists {
|
||||||
|
m.evictOne()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove old tag associations if key exists
|
||||||
|
if oldItem, exists := m.items[key]; exists {
|
||||||
|
for _, tag := range oldItem.Tags {
|
||||||
|
if keySet, ok := m.tagToKeys[tag]; ok {
|
||||||
|
delete(keySet, key)
|
||||||
|
if len(keySet) == 0 {
|
||||||
|
delete(m.tagToKeys, tag)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store the item
|
||||||
|
m.items[key] = &memoryItem{
|
||||||
|
Value: value,
|
||||||
|
Expiration: expiration,
|
||||||
|
LastAccess: time.Now(),
|
||||||
|
Tags: tags,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add new tag associations
|
||||||
|
for _, tag := range tags {
|
||||||
|
if m.tagToKeys[tag] == nil {
|
||||||
|
m.tagToKeys[tag] = make(map[string]struct{})
|
||||||
|
}
|
||||||
|
m.tagToKeys[tag][key] = struct{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete removes a key from the cache.
|
||||||
|
func (m *MemoryProvider) Delete(ctx context.Context, key string) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
// Remove tag associations
|
||||||
|
if item, exists := m.items[key]; exists {
|
||||||
|
for _, tag := range item.Tags {
|
||||||
|
if keySet, ok := m.tagToKeys[tag]; ok {
|
||||||
|
delete(keySet, key)
|
||||||
|
if len(keySet) == 0 {
|
||||||
|
delete(m.tagToKeys, tag)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(m.items, key)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByTag removes all keys associated with the given tag.
|
||||||
|
func (m *MemoryProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
// Get all keys associated with this tag
|
||||||
|
keySet, exists := m.tagToKeys[tag]
|
||||||
|
if !exists {
|
||||||
|
return nil // No keys with this tag
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete all items with this tag
|
||||||
|
for key := range keySet {
|
||||||
|
if item, ok := m.items[key]; ok {
|
||||||
|
// Remove this tag from the item's tag list
|
||||||
|
newTags := make([]string, 0, len(item.Tags))
|
||||||
|
for _, t := range item.Tags {
|
||||||
|
if t != tag {
|
||||||
|
newTags = append(newTags, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If item has no more tags, delete it
|
||||||
|
// Otherwise update its tags
|
||||||
|
if len(newTags) == 0 {
|
||||||
|
delete(m.items, key)
|
||||||
|
} else {
|
||||||
|
item.Tags = newTags
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove the tag mapping
|
||||||
|
delete(m.tagToKeys, tag)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByPattern removes all keys matching the pattern.
|
||||||
|
func (m *MemoryProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid pattern: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for key := range m.items {
|
||||||
|
if re.MatchString(key) {
|
||||||
|
delete(m.items, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear removes all items from the cache.
|
||||||
|
func (m *MemoryProvider) Clear(ctx context.Context) error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
m.items = make(map[string]*memoryItem)
|
||||||
|
m.hits.Store(0)
|
||||||
|
m.misses.Store(0)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exists checks if a key exists in the cache.
|
||||||
|
func (m *MemoryProvider) Exists(ctx context.Context, key string) bool {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
|
item, exists := m.items[key]
|
||||||
|
if !exists {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return !item.isExpired()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes the provider and releases any resources.
|
||||||
|
func (m *MemoryProvider) Close() error {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
m.items = nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stats returns statistics about the cache provider.
|
||||||
|
func (m *MemoryProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
|
||||||
|
// Clean expired items first
|
||||||
|
validKeys := 0
|
||||||
|
for _, item := range m.items {
|
||||||
|
if !item.isExpired() {
|
||||||
|
validKeys++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &CacheStats{
|
||||||
|
Hits: m.hits.Load(),
|
||||||
|
Misses: m.misses.Load(),
|
||||||
|
Keys: int64(validKeys),
|
||||||
|
ProviderType: "memory",
|
||||||
|
ProviderStats: map[string]any{
|
||||||
|
"capacity": m.options.MaxSize,
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// evictOne removes one item from the cache using LRU strategy.
|
||||||
|
func (m *MemoryProvider) evictOne() {
|
||||||
|
var oldestKey string
|
||||||
|
var oldestTime time.Time
|
||||||
|
|
||||||
|
for key, item := range m.items {
|
||||||
|
if item.isExpired() {
|
||||||
|
delete(m.items, key)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if oldestKey == "" || item.LastAccess.Before(oldestTime) {
|
||||||
|
oldestKey = key
|
||||||
|
oldestTime = item.LastAccess
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if oldestKey != "" {
|
||||||
|
delete(m.items, oldestKey)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CleanExpired removes all expired items from the cache.
|
||||||
|
func (m *MemoryProvider) CleanExpired(ctx context.Context) int {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
|
||||||
|
count := 0
|
||||||
|
for key, item := range m.items {
|
||||||
|
if item.isExpired() {
|
||||||
|
delete(m.items, key)
|
||||||
|
count++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return count
|
||||||
|
}
|
||||||
+269
@@ -0,0 +1,269 @@
|
|||||||
|
package cache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RedisProvider is a Redis implementation of the Provider interface.
|
||||||
|
type RedisProvider struct {
|
||||||
|
client *redis.Client
|
||||||
|
options *Options
|
||||||
|
}
|
||||||
|
|
||||||
|
// RedisConfig contains Redis-specific configuration.
|
||||||
|
type RedisConfig struct {
|
||||||
|
// Host is the Redis server host (default: localhost)
|
||||||
|
Host string
|
||||||
|
|
||||||
|
// Port is the Redis server port (default: 6379)
|
||||||
|
Port int
|
||||||
|
|
||||||
|
// Password for Redis authentication (optional)
|
||||||
|
Password string
|
||||||
|
|
||||||
|
// DB is the Redis database number (default: 0)
|
||||||
|
DB int
|
||||||
|
|
||||||
|
// PoolSize is the maximum number of connections (default: 10)
|
||||||
|
PoolSize int
|
||||||
|
|
||||||
|
// Options contains general cache options
|
||||||
|
Options *Options
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRedisProvider creates a new Redis cache provider.
|
||||||
|
func NewRedisProvider(config *RedisConfig) (*RedisProvider, error) {
|
||||||
|
if config == nil {
|
||||||
|
config = &RedisConfig{
|
||||||
|
Host: "localhost",
|
||||||
|
Port: 6379,
|
||||||
|
DB: 0,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if config.Host == "" {
|
||||||
|
config.Host = "localhost"
|
||||||
|
}
|
||||||
|
if config.Port == 0 {
|
||||||
|
config.Port = 6379
|
||||||
|
}
|
||||||
|
if config.PoolSize == 0 {
|
||||||
|
config.PoolSize = 10
|
||||||
|
}
|
||||||
|
|
||||||
|
if config.Options == nil {
|
||||||
|
config.Options = &Options{
|
||||||
|
DefaultTTL: 5 * time.Minute,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
client := redis.NewClient(&redis.Options{
|
||||||
|
Addr: fmt.Sprintf("%s:%d", config.Host, config.Port),
|
||||||
|
Password: config.Password,
|
||||||
|
DB: config.DB,
|
||||||
|
PoolSize: config.PoolSize,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Test connection
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := client.Ping(ctx).Err(); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to connect to Redis: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &RedisProvider{
|
||||||
|
client: client,
|
||||||
|
options: config.Options,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get retrieves a value from the cache by key.
|
||||||
|
func (r *RedisProvider) Get(ctx context.Context, key string) ([]byte, bool) {
|
||||||
|
val, err := r.client.Get(ctx, key).Bytes()
|
||||||
|
if err == redis.Nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return val, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set stores a value in the cache with the specified TTL.
|
||||||
|
func (r *RedisProvider) Set(ctx context.Context, key string, value []byte, ttl time.Duration) error {
|
||||||
|
if ttl == 0 {
|
||||||
|
ttl = r.options.DefaultTTL
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.client.Set(ctx, key, value, ttl).Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetWithTags stores a value in the cache with the specified TTL and tags.
|
||||||
|
func (r *RedisProvider) SetWithTags(ctx context.Context, key string, value []byte, ttl time.Duration, tags []string) error {
|
||||||
|
if ttl == 0 {
|
||||||
|
ttl = r.options.DefaultTTL
|
||||||
|
}
|
||||||
|
|
||||||
|
pipe := r.client.Pipeline()
|
||||||
|
|
||||||
|
// Set the value
|
||||||
|
pipe.Set(ctx, key, value, ttl)
|
||||||
|
|
||||||
|
// Add key to each tag's set
|
||||||
|
for _, tag := range tags {
|
||||||
|
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||||
|
pipe.SAdd(ctx, tagKey, key)
|
||||||
|
// Set expiration on tag set (longer than cache items to ensure cleanup)
|
||||||
|
if ttl > 0 {
|
||||||
|
pipe.Expire(ctx, tagKey, ttl+time.Hour)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store tags for this key for later cleanup
|
||||||
|
if len(tags) > 0 {
|
||||||
|
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||||
|
pipe.SAdd(ctx, tagsKey, tags)
|
||||||
|
if ttl > 0 {
|
||||||
|
pipe.Expire(ctx, tagsKey, ttl)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := pipe.Exec(ctx)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete removes a key from the cache.
|
||||||
|
func (r *RedisProvider) Delete(ctx context.Context, key string) error {
|
||||||
|
pipe := r.client.Pipeline()
|
||||||
|
|
||||||
|
// Get tags for this key
|
||||||
|
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||||
|
tags, err := r.client.SMembers(ctx, tagsKey).Result()
|
||||||
|
if err == nil && len(tags) > 0 {
|
||||||
|
// Remove key from each tag set
|
||||||
|
for _, tag := range tags {
|
||||||
|
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||||
|
pipe.SRem(ctx, tagKey, key)
|
||||||
|
}
|
||||||
|
// Delete the tags key
|
||||||
|
pipe.Del(ctx, tagsKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete the actual key
|
||||||
|
pipe.Del(ctx, key)
|
||||||
|
|
||||||
|
_, err = pipe.Exec(ctx)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByTag removes all keys associated with the given tag.
|
||||||
|
func (r *RedisProvider) DeleteByTag(ctx context.Context, tag string) error {
|
||||||
|
tagKey := fmt.Sprintf("cache:tag:%s", tag)
|
||||||
|
|
||||||
|
// Get all keys associated with this tag
|
||||||
|
keys, err := r.client.SMembers(ctx, tagKey).Result()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(keys) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
pipe := r.client.Pipeline()
|
||||||
|
|
||||||
|
// Delete all keys and their tag associations
|
||||||
|
for _, key := range keys {
|
||||||
|
pipe.Del(ctx, key)
|
||||||
|
// Also delete the tags key for this cache key
|
||||||
|
tagsKey := fmt.Sprintf("cache:tags:%s", key)
|
||||||
|
pipe.Del(ctx, tagsKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete the tag set itself
|
||||||
|
pipe.Del(ctx, tagKey)
|
||||||
|
|
||||||
|
_, err = pipe.Exec(ctx)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteByPattern removes all keys matching the pattern.
|
||||||
|
func (r *RedisProvider) DeleteByPattern(ctx context.Context, pattern string) error {
|
||||||
|
iter := r.client.Scan(ctx, 0, pattern, 0).Iterator()
|
||||||
|
pipe := r.client.Pipeline()
|
||||||
|
|
||||||
|
count := 0
|
||||||
|
for iter.Next(ctx) {
|
||||||
|
pipe.Del(ctx, iter.Val())
|
||||||
|
count++
|
||||||
|
|
||||||
|
// Execute pipeline in batches of 100
|
||||||
|
if count%100 == 0 {
|
||||||
|
if _, err := pipe.Exec(ctx); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
pipe = r.client.Pipeline()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := iter.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute remaining commands
|
||||||
|
if count%100 != 0 {
|
||||||
|
_, err := pipe.Exec(ctx)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear removes all items from the cache.
|
||||||
|
func (r *RedisProvider) Clear(ctx context.Context) error {
|
||||||
|
return r.client.FlushDB(ctx).Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exists checks if a key exists in the cache.
|
||||||
|
func (r *RedisProvider) Exists(ctx context.Context, key string) bool {
|
||||||
|
result, err := r.client.Exists(ctx, key).Result()
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return result > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes the provider and releases any resources.
|
||||||
|
func (r *RedisProvider) Close() error {
|
||||||
|
return r.client.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stats returns statistics about the cache provider.
|
||||||
|
func (r *RedisProvider) Stats(ctx context.Context) (*CacheStats, error) {
|
||||||
|
info, err := r.client.Info(ctx, "stats", "keyspace").Result()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get Redis stats: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
dbSize, err := r.client.DBSize(ctx).Result()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get DB size: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse stats from INFO command
|
||||||
|
// This is a simplified version - you may want to parse more detailed stats
|
||||||
|
stats := &CacheStats{
|
||||||
|
Keys: dbSize,
|
||||||
|
ProviderType: "redis",
|
||||||
|
ProviderStats: map[string]any{
|
||||||
|
"info": info,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return stats, nil
|
||||||
|
}
|
||||||
Generated
Vendored
+218
@@ -0,0 +1,218 @@
|
|||||||
|
# Automatic Relation Loading Strategies
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
**NEW:** The database adapters now **automatically** choose the optimal loading strategy by inspecting your model's relationship tags!
|
||||||
|
|
||||||
|
Simply use `PreloadRelation()` and the system automatically:
|
||||||
|
- Detects relationship type from Bun/GORM tags
|
||||||
|
- Uses **JOIN** for many-to-one and one-to-one (efficient, no duplication)
|
||||||
|
- Uses **separate query** for one-to-many and many-to-many (avoids duplication)
|
||||||
|
|
||||||
|
## How It Works
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Just write this - the system handles the rest!
|
||||||
|
db.NewSelect().
|
||||||
|
Model(&links).
|
||||||
|
PreloadRelation("Provider"). // ✓ Auto-detects belongs-to → uses JOIN
|
||||||
|
PreloadRelation("Tags"). // ✓ Auto-detects has-many → uses separate query
|
||||||
|
Scan(ctx, &links)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Detection Logic
|
||||||
|
|
||||||
|
The system inspects your model's struct tags:
|
||||||
|
|
||||||
|
**Bun models:**
|
||||||
|
```go
|
||||||
|
type Link struct {
|
||||||
|
Provider *Provider `bun:"rel:belongs-to"` // → Detected: belongs-to → JOIN
|
||||||
|
Tags []Tag `bun:"rel:has-many"` // → Detected: has-many → Separate query
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**GORM models:**
|
||||||
|
```go
|
||||||
|
type Link struct {
|
||||||
|
ProviderID int
|
||||||
|
Provider *Provider `gorm:"foreignKey:ProviderID"` // → Detected: belongs-to → JOIN
|
||||||
|
Tags []Tag `gorm:"many2many:link_tags"` // → Detected: many-to-many → Separate query
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Type inference (fallback):**
|
||||||
|
- `[]Type` (slice) → has-many → Separate query
|
||||||
|
- `*Type` (pointer) → belongs-to → JOIN
|
||||||
|
- `Type` (struct) → belongs-to → JOIN
|
||||||
|
|
||||||
|
### What Gets Logged
|
||||||
|
|
||||||
|
Enable debug logging to see strategy selection:
|
||||||
|
|
||||||
|
```go
|
||||||
|
bunAdapter.EnableQueryDebug()
|
||||||
|
```
|
||||||
|
|
||||||
|
**Output:**
|
||||||
|
```
|
||||||
|
DEBUG: PreloadRelation 'Provider' detected as: belongs-to
|
||||||
|
INFO: Using JOIN strategy for belongs-to relation 'Provider'
|
||||||
|
DEBUG: PreloadRelation 'Links' detected as: has-many
|
||||||
|
DEBUG: Using separate query for has-many relation 'Links'
|
||||||
|
```
|
||||||
|
|
||||||
|
## Relationship Types
|
||||||
|
|
||||||
|
| Bun Tag | GORM Pattern | Field Type | Strategy | Why |
|
||||||
|
|---------|--------------|------------|----------|-----|
|
||||||
|
| `rel:has-many` | Slice field | `[]Type` | Separate Query | Avoids duplicating parent data |
|
||||||
|
| `rel:belongs-to` | `foreignKey:` | `*Type` | JOIN | Single parent, no duplication |
|
||||||
|
| `rel:has-one` | Single pointer | `*Type` | JOIN | One-to-one, no duplication |
|
||||||
|
| `rel:many-to-many` | `many2many:` | `[]Type` | Separate Query | Complex join, avoid cartesian |
|
||||||
|
|
||||||
|
## Manual Override
|
||||||
|
|
||||||
|
If you need to force a specific strategy, use `JoinRelation()`:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Force JOIN even for has-many (not recommended)
|
||||||
|
db.NewSelect().
|
||||||
|
Model(&providers).
|
||||||
|
JoinRelation("Links"). // Explicitly use JOIN
|
||||||
|
Scan(ctx, &providers)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Examples
|
||||||
|
|
||||||
|
### Automatic Strategy Selection (Recommended)
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Example 1: Loading parent provider for each link
|
||||||
|
// System detects belongs-to → uses JOIN automatically
|
||||||
|
db.NewSelect().
|
||||||
|
Model(&links).
|
||||||
|
PreloadRelation("Provider", func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
return q.Where("active = ?", true)
|
||||||
|
}).
|
||||||
|
Scan(ctx, &links)
|
||||||
|
|
||||||
|
// Generated SQL: Single query with JOIN
|
||||||
|
// SELECT links.*, providers.*
|
||||||
|
// FROM links
|
||||||
|
// LEFT JOIN providers ON links.provider_id = providers.id
|
||||||
|
// WHERE providers.active = true
|
||||||
|
|
||||||
|
// Example 2: Loading child links for each provider
|
||||||
|
// System detects has-many → uses separate query automatically
|
||||||
|
db.NewSelect().
|
||||||
|
Model(&providers).
|
||||||
|
PreloadRelation("Links", func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
return q.Where("active = ?", true)
|
||||||
|
}).
|
||||||
|
Scan(ctx, &providers)
|
||||||
|
|
||||||
|
// Generated SQL: Two queries
|
||||||
|
// Query 1: SELECT * FROM providers
|
||||||
|
// Query 2: SELECT * FROM links
|
||||||
|
// WHERE provider_id IN (1, 2, 3, ...)
|
||||||
|
// AND active = true
|
||||||
|
```
|
||||||
|
|
||||||
|
### Mixed Relationships
|
||||||
|
|
||||||
|
```go
|
||||||
|
type Order struct {
|
||||||
|
ID int
|
||||||
|
CustomerID int
|
||||||
|
Customer *Customer `bun:"rel:belongs-to"` // JOIN
|
||||||
|
Items []Item `bun:"rel:has-many"` // Separate
|
||||||
|
Invoice *Invoice `bun:"rel:has-one"` // JOIN
|
||||||
|
}
|
||||||
|
|
||||||
|
// All three handled optimally!
|
||||||
|
db.NewSelect().
|
||||||
|
Model(&orders).
|
||||||
|
PreloadRelation("Customer"). // → JOIN (many-to-one)
|
||||||
|
PreloadRelation("Items"). // → Separate (one-to-many)
|
||||||
|
PreloadRelation("Invoice"). // → JOIN (one-to-one)
|
||||||
|
Scan(ctx, &orders)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Performance Benefits
|
||||||
|
|
||||||
|
### Before (Manual Strategy Selection)
|
||||||
|
|
||||||
|
```go
|
||||||
|
// You had to remember which to use:
|
||||||
|
.PreloadRelation("Provider") // Should I use PreloadRelation or JoinRelation?
|
||||||
|
.PreloadRelation("Links") // Which is more efficient here?
|
||||||
|
```
|
||||||
|
|
||||||
|
### After (Automatic Selection)
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Just use PreloadRelation everywhere:
|
||||||
|
.PreloadRelation("Provider") // ✓ System uses JOIN automatically
|
||||||
|
.PreloadRelation("Links") // ✓ System uses separate query automatically
|
||||||
|
```
|
||||||
|
|
||||||
|
## Migration Guide
|
||||||
|
|
||||||
|
**No changes needed!** If you're already using `PreloadRelation()`, it now automatically optimizes:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Before: Always used separate query
|
||||||
|
.PreloadRelation("Provider") // Inefficient: extra round trip
|
||||||
|
|
||||||
|
// After: Automatic optimization
|
||||||
|
.PreloadRelation("Provider") // ✓ Now uses JOIN automatically!
|
||||||
|
```
|
||||||
|
|
||||||
|
## Implementation Details
|
||||||
|
|
||||||
|
### Supported Bun Tags
|
||||||
|
- `rel:has-many` → Separate query
|
||||||
|
- `rel:belongs-to` → JOIN
|
||||||
|
- `rel:has-one` → JOIN
|
||||||
|
- `rel:many-to-many` or `rel:m2m` → Separate query
|
||||||
|
|
||||||
|
### Supported GORM Patterns
|
||||||
|
- `many2many:` tag → Separate query
|
||||||
|
- `foreignKey:` tag → JOIN (belongs-to)
|
||||||
|
- `[]Type` slice without many2many → Separate query (has-many)
|
||||||
|
- `*Type` pointer with foreignKey → JOIN (belongs-to)
|
||||||
|
- `*Type` pointer without foreignKey → JOIN (has-one)
|
||||||
|
|
||||||
|
### Fallback Behavior
|
||||||
|
- `[]Type` (slice) → Separate query (safe default for collections)
|
||||||
|
- `*Type` or `Type` (single) → JOIN (safe default for single relations)
|
||||||
|
- Unknown → Separate query (safest default)
|
||||||
|
|
||||||
|
## Debugging
|
||||||
|
|
||||||
|
To see strategy selection in action:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Enable debug logging
|
||||||
|
bunAdapter.EnableQueryDebug() // or gormAdapter.EnableQueryDebug()
|
||||||
|
|
||||||
|
// Run your query
|
||||||
|
db.NewSelect().
|
||||||
|
Model(&records).
|
||||||
|
PreloadRelation("RelationName").
|
||||||
|
Scan(ctx, &records)
|
||||||
|
|
||||||
|
// Check logs for:
|
||||||
|
// - "PreloadRelation 'X' detected as: belongs-to"
|
||||||
|
// - "Using JOIN strategy for belongs-to relation 'X'"
|
||||||
|
// - Actual SQL queries executed
|
||||||
|
```
|
||||||
|
|
||||||
|
## Best Practices
|
||||||
|
|
||||||
|
1. **Use PreloadRelation() for everything** - Let the system optimize
|
||||||
|
2. **Define proper relationship tags** - Ensures correct detection
|
||||||
|
3. **Only use JoinRelation() for overrides** - When you know better than auto-detection
|
||||||
|
4. **Enable debug logging during development** - Verify optimal strategies are chosen
|
||||||
|
5. **Trust the system** - It's designed to choose correctly based on relationship type
|
||||||
+1770
File diff suppressed because it is too large
Load Diff
+1018
File diff suppressed because it is too large
Load Diff
+1600
File diff suppressed because it is too large
Load Diff
Generated
Vendored
+176
@@ -0,0 +1,176 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
_ "github.com/jackc/pgx/v5/stdlib" // PostgreSQL driver
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Example demonstrates how to use the PgSQL adapter
|
||||||
|
func ExamplePgSQLAdapter() error {
|
||||||
|
// Connect to PostgreSQL database
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to open database: %w", err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
// Create the PgSQL adapter
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
|
||||||
|
// Enable query debugging (optional)
|
||||||
|
adapter.EnableQueryDebug()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Example 1: Simple SELECT query
|
||||||
|
var results []map[string]interface{}
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Table("users").
|
||||||
|
Where("age > ?", 18).
|
||||||
|
Order("created_at DESC").
|
||||||
|
Limit(10).
|
||||||
|
Scan(ctx, &results)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("select failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 2: INSERT query
|
||||||
|
result, err := adapter.NewInsert().
|
||||||
|
Table("users").
|
||||||
|
Value("name", "John Doe").
|
||||||
|
Value("email", "john@example.com").
|
||||||
|
Value("age", 25).
|
||||||
|
Returning("id").
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("insert failed: %w", err)
|
||||||
|
}
|
||||||
|
fmt.Printf("Rows affected: %d\n", result.RowsAffected())
|
||||||
|
|
||||||
|
// Example 3: UPDATE query
|
||||||
|
result, err = adapter.NewUpdate().
|
||||||
|
Table("users").
|
||||||
|
Set("name", "Jane Doe").
|
||||||
|
Where("id = ?", 1).
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("update failed: %w", err)
|
||||||
|
}
|
||||||
|
fmt.Printf("Rows updated: %d\n", result.RowsAffected())
|
||||||
|
|
||||||
|
// Example 4: DELETE query
|
||||||
|
result, err = adapter.NewDelete().
|
||||||
|
Table("users").
|
||||||
|
Where("age < ?", 18).
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("delete failed: %w", err)
|
||||||
|
}
|
||||||
|
fmt.Printf("Rows deleted: %d\n", result.RowsAffected())
|
||||||
|
|
||||||
|
// Example 5: Using transactions
|
||||||
|
err = adapter.RunInTransaction(ctx, func(tx common.Database) error {
|
||||||
|
// Insert a new user
|
||||||
|
_, err := tx.NewInsert().
|
||||||
|
Table("users").
|
||||||
|
Value("name", "Transaction User").
|
||||||
|
Value("email", "tx@example.com").
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update another user
|
||||||
|
_, err = tx.NewUpdate().
|
||||||
|
Table("users").
|
||||||
|
Set("verified", true).
|
||||||
|
Where("email = ?", "tx@example.com").
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Both operations succeed or both rollback
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("transaction failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 6: JOIN query
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Table("users u").
|
||||||
|
Column("u.id", "u.name", "p.title as post_title").
|
||||||
|
LeftJoin("posts p ON p.user_id = u.id").
|
||||||
|
Where("u.active = ?", true).
|
||||||
|
Scan(ctx, &results)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("join query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 7: Aggregation query
|
||||||
|
count, err := adapter.NewSelect().
|
||||||
|
Table("users").
|
||||||
|
Where("active = ?", true).
|
||||||
|
Count(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("count failed: %w", err)
|
||||||
|
}
|
||||||
|
fmt.Printf("Active users: %d\n", count)
|
||||||
|
|
||||||
|
// Example 8: Raw SQL execution
|
||||||
|
_, err = adapter.Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_users_email ON users(email)")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("raw exec failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 9: Raw SQL query
|
||||||
|
var users []map[string]interface{}
|
||||||
|
err = adapter.Query(ctx, &users, "SELECT * FROM users WHERE age > $1 LIMIT $2", 18, 10)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("raw query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// User is an example model
|
||||||
|
type User struct {
|
||||||
|
ID int `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
Age int `json:"age"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// TableName implements common.TableNameProvider
|
||||||
|
func (u User) TableName() string {
|
||||||
|
return "users"
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleWithModel demonstrates using models with the PgSQL adapter
|
||||||
|
func ExampleWithModel() error {
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Use model with adapter
|
||||||
|
user := User{}
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&user).
|
||||||
|
Where("id = ?", 1).
|
||||||
|
Scan(ctx, &user)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
Generated
Vendored
+275
@@ -0,0 +1,275 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
|
||||||
|
_ "github.com/jackc/pgx/v5/stdlib"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Example models for demonstrating preload functionality
|
||||||
|
|
||||||
|
// Author model - has many Posts
|
||||||
|
type Author struct {
|
||||||
|
ID int `db:"id"`
|
||||||
|
Name string `db:"name"`
|
||||||
|
Email string `db:"email"`
|
||||||
|
Posts []*Post `bun:"rel:has-many,join:id=author_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a Author) TableName() string {
|
||||||
|
return "authors"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Post model - belongs to Author, has many Comments
|
||||||
|
type Post struct {
|
||||||
|
ID int `db:"id"`
|
||||||
|
Title string `db:"title"`
|
||||||
|
Content string `db:"content"`
|
||||||
|
AuthorID int `db:"author_id"`
|
||||||
|
Author *Author `bun:"rel:belongs-to,join:author_id=id"`
|
||||||
|
Comments []*Comment `bun:"rel:has-many,join:id=post_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p Post) TableName() string {
|
||||||
|
return "posts"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Comment model - belongs to Post
|
||||||
|
type Comment struct {
|
||||||
|
ID int `db:"id"`
|
||||||
|
Content string `db:"content"`
|
||||||
|
PostID int `db:"post_id"`
|
||||||
|
Post *Post `bun:"rel:belongs-to,join:post_id=id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c Comment) TableName() string {
|
||||||
|
return "comments"
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExamplePreload demonstrates the Preload functionality
|
||||||
|
func ExamplePreload() error {
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Example 1: Simple Preload (uses subquery for has-many)
|
||||||
|
var authors []*Author
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&Author{}).
|
||||||
|
Table("authors").
|
||||||
|
Preload("Posts"). // Load all posts for each author
|
||||||
|
Scan(ctx, &authors)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now authors[i].Posts will be populated with their posts
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExamplePreloadRelation demonstrates smart PreloadRelation with auto-detection
|
||||||
|
func ExamplePreloadRelation() error {
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Example 1: PreloadRelation auto-detects has-many (uses subquery)
|
||||||
|
var authors []*Author
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&Author{}).
|
||||||
|
Table("authors").
|
||||||
|
PreloadRelation("Posts", func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
return q.Where("published = ?", true).Order("created_at DESC")
|
||||||
|
}).
|
||||||
|
Where("active = ?", true).
|
||||||
|
Scan(ctx, &authors)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 2: PreloadRelation auto-detects belongs-to (uses JOIN)
|
||||||
|
var posts []*Post
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&Post{}).
|
||||||
|
Table("posts").
|
||||||
|
PreloadRelation("Author"). // Will use JOIN because it's belongs-to
|
||||||
|
Scan(ctx, &posts)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 3: Nested preloads
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&Author{}).
|
||||||
|
Table("authors").
|
||||||
|
PreloadRelation("Posts", func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
// First load posts, then preload comments for each post
|
||||||
|
return q.Limit(10)
|
||||||
|
}).
|
||||||
|
Scan(ctx, &authors)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Manually load nested relationships (two-level preloading)
|
||||||
|
for _, author := range authors {
|
||||||
|
if author.Posts != nil {
|
||||||
|
for _, post := range author.Posts {
|
||||||
|
var comments []*Comment
|
||||||
|
err := adapter.NewSelect().
|
||||||
|
Table("comments").
|
||||||
|
Where("post_id = ?", post.ID).
|
||||||
|
Scan(ctx, &comments)
|
||||||
|
if err == nil {
|
||||||
|
post.Comments = comments
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleJoinRelation demonstrates explicit JOIN loading
|
||||||
|
func ExampleJoinRelation() error {
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Example 1: Force JOIN for belongs-to relationship
|
||||||
|
var posts []*Post
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&Post{}).
|
||||||
|
Table("posts").
|
||||||
|
JoinRelation("Author", func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
return q.Where("active = ?", true)
|
||||||
|
}).
|
||||||
|
Scan(ctx, &posts)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 2: Multiple JOINs
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&Post{}).
|
||||||
|
Table("posts p").
|
||||||
|
Column("p.*", "a.name as author_name", "a.email as author_email").
|
||||||
|
LeftJoin("authors a ON a.id = p.author_id").
|
||||||
|
Where("p.published = ?", true).
|
||||||
|
Scan(ctx, &posts)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleScanModel demonstrates ScanModel with struct destinations
|
||||||
|
func ExampleScanModel() error {
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Example 1: Scan single struct
|
||||||
|
author := Author{}
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&author).
|
||||||
|
Table("authors").
|
||||||
|
Where("id = ?", 1).
|
||||||
|
ScanModel(ctx) // ScanModel automatically uses the model set with Model()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 2: Scan slice of structs
|
||||||
|
authors := []*Author{}
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&authors).
|
||||||
|
Table("authors").
|
||||||
|
Where("active = ?", true).
|
||||||
|
Limit(10).
|
||||||
|
ScanModel(ctx)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleCompleteWorkflow demonstrates a complete workflow with preloading
|
||||||
|
func ExampleCompleteWorkflow() error {
|
||||||
|
dsn := "postgres://username:password@localhost:5432/dbname?sslmode=disable"
|
||||||
|
db, err := sql.Open("pgx", dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
adapter := NewPgSQLAdapter(db)
|
||||||
|
adapter.EnableQueryDebug() // Enable query logging
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Step 1: Create an author
|
||||||
|
author := &Author{
|
||||||
|
Name: "John Doe",
|
||||||
|
Email: "john@example.com",
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := adapter.NewInsert().
|
||||||
|
Table("authors").
|
||||||
|
Value("name", author.Name).
|
||||||
|
Value("email", author.Email).
|
||||||
|
Returning("id").
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = result
|
||||||
|
|
||||||
|
// Step 2: Load author with all their posts
|
||||||
|
var loadedAuthor Author
|
||||||
|
err = adapter.NewSelect().
|
||||||
|
Model(&loadedAuthor).
|
||||||
|
Table("authors").
|
||||||
|
PreloadRelation("Posts", func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
return q.Order("created_at DESC").Limit(5)
|
||||||
|
}).
|
||||||
|
Where("id = ?", 1).
|
||||||
|
ScanModel(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 3: Update author name
|
||||||
|
_, err = adapter.NewUpdate().
|
||||||
|
Table("authors").
|
||||||
|
Set("name", "Jane Doe").
|
||||||
|
Where("id = ?", 1).
|
||||||
|
Exec(ctx)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
Generated
Vendored
+335
@@ -0,0 +1,335 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
const maxMetricFallbackEntityLength = 120
|
||||||
|
|
||||||
|
func recordQueryMetrics(enabled bool, operation, schema, entity, table string, startedAt time.Time, err error) {
|
||||||
|
if !enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
metrics.GetProvider().RecordDBQuery(
|
||||||
|
normalizeMetricOperation(operation),
|
||||||
|
normalizeMetricSchema(schema),
|
||||||
|
normalizeMetricEntity(entity, table),
|
||||||
|
normalizeMetricTable(table),
|
||||||
|
time.Since(startedAt),
|
||||||
|
err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricOperation(operation string) string {
|
||||||
|
operation = strings.ToUpper(strings.TrimSpace(operation))
|
||||||
|
if operation == "" {
|
||||||
|
return "UNKNOWN"
|
||||||
|
}
|
||||||
|
return operation
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricSchema(schema string) string {
|
||||||
|
schema = cleanMetricIdentifier(schema)
|
||||||
|
if schema == "" {
|
||||||
|
return "default"
|
||||||
|
}
|
||||||
|
return schema
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricEntity(entity, table string) string {
|
||||||
|
entity = cleanMetricIdentifier(entity)
|
||||||
|
if entity != "" {
|
||||||
|
return entity
|
||||||
|
}
|
||||||
|
|
||||||
|
table = cleanMetricIdentifier(table)
|
||||||
|
if table != "" {
|
||||||
|
return table
|
||||||
|
}
|
||||||
|
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMetricTable(table string) string {
|
||||||
|
table = cleanMetricIdentifier(table)
|
||||||
|
if table == "" {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
return table
|
||||||
|
}
|
||||||
|
|
||||||
|
func entityNameFromModel(model interface{}, table string) string {
|
||||||
|
if model == nil {
|
||||||
|
return cleanMetricIdentifier(table)
|
||||||
|
}
|
||||||
|
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType == nil {
|
||||||
|
return cleanMetricIdentifier(table)
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType.Kind() == reflect.Struct && modelType.Name() != "" {
|
||||||
|
return reflection.ToSnakeCase(modelType.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
return cleanMetricIdentifier(table)
|
||||||
|
}
|
||||||
|
|
||||||
|
func schemaAndTableFromModel(model interface{}, driverName string) (schema, table string) {
|
||||||
|
provider, ok := tableNameProviderFromModel(model)
|
||||||
|
if !ok {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return parseTableName(provider.TableName(), driverName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// tableNameProviderType is cached to avoid repeated reflection on every call.
|
||||||
|
var tableNameProviderType = reflect.TypeOf((*common.TableNameProvider)(nil)).Elem()
|
||||||
|
|
||||||
|
func tableNameProviderFromModel(model interface{}) (common.TableNameProvider, bool) {
|
||||||
|
if model == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider, ok := model.(common.TableNameProvider); ok {
|
||||||
|
return provider, true
|
||||||
|
}
|
||||||
|
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check whether *T implements TableNameProvider before allocating.
|
||||||
|
ptrType := reflect.PointerTo(modelType)
|
||||||
|
if !ptrType.Implements(tableNameProviderType) && !modelType.Implements(tableNameProviderType) {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
modelValue := reflect.New(modelType)
|
||||||
|
if provider, ok := modelValue.Interface().(common.TableNameProvider); ok {
|
||||||
|
return provider, true
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider, ok := modelValue.Elem().Interface().(common.TableNameProvider); ok {
|
||||||
|
return provider, true
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func metricTargetFromRawQuery(query, driverName string) (operation, schema, entity, table string) {
|
||||||
|
operation = normalizeMetricOperation(firstQueryKeyword(query))
|
||||||
|
tableRef := tableFromRawQuery(query, operation)
|
||||||
|
if tableRef == "" {
|
||||||
|
return operation, "", fallbackMetricEntityFromQuery(query), "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
schema, table = parseTableName(tableRef, driverName)
|
||||||
|
entity = cleanMetricIdentifier(table)
|
||||||
|
return operation, schema, entity, table
|
||||||
|
}
|
||||||
|
|
||||||
|
func fallbackMetricEntityFromQuery(query string) string {
|
||||||
|
query = sanitizeMetricQueryShape(query)
|
||||||
|
if query == "" {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(query) > maxMetricFallbackEntityLength {
|
||||||
|
return query[:maxMetricFallbackEntityLength-3] + "..."
|
||||||
|
}
|
||||||
|
|
||||||
|
return query
|
||||||
|
}
|
||||||
|
|
||||||
|
func sanitizeMetricQueryShape(query string) string {
|
||||||
|
query = strings.TrimSpace(query)
|
||||||
|
if query == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var out strings.Builder
|
||||||
|
for i := 0; i < len(query); {
|
||||||
|
if query[i] == '\'' {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) {
|
||||||
|
if query[i] == '\'' {
|
||||||
|
if i+1 < len(query) && query[i+1] == '\'' {
|
||||||
|
i += 2
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
break
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == '?' {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == '$' && i+1 < len(query) && isASCIIDigit(query[i+1]) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && isASCIIDigit(query[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == ':' && (i == 0 || query[i-1] != ':') && i+1 < len(query) && isIdentifierStart(query[i+1]) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && isIdentifierPart(query[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if query[i] == '@' && (i == 0 || query[i-1] != '@') && i+1 < len(query) && isIdentifierStart(query[i+1]) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && isIdentifierPart(query[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if startsNumericLiteral(query, i) {
|
||||||
|
out.WriteByte('?')
|
||||||
|
i++
|
||||||
|
for i < len(query) && (isASCIIDigit(query[i]) || query[i] == '.') {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
out.WriteByte(query[i])
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(strings.Fields(out.String()), " ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func startsNumericLiteral(query string, idx int) bool {
|
||||||
|
if idx >= len(query) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
start := idx
|
||||||
|
if query[idx] == '-' {
|
||||||
|
if idx+1 >= len(query) || !isASCIIDigit(query[idx+1]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
start++
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isASCIIDigit(query[start]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if idx > 0 && isIdentifierPart(query[idx-1]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if start+1 < len(query) && query[start] == '0' && (query[start+1] == 'x' || query[start+1] == 'X') {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func isASCIIDigit(ch byte) bool {
|
||||||
|
return ch >= '0' && ch <= '9'
|
||||||
|
}
|
||||||
|
|
||||||
|
func isIdentifierStart(ch byte) bool {
|
||||||
|
return (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || ch == '_'
|
||||||
|
}
|
||||||
|
|
||||||
|
func isIdentifierPart(ch byte) bool {
|
||||||
|
return isIdentifierStart(ch) || isASCIIDigit(ch)
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstQueryKeyword(query string) string {
|
||||||
|
query = strings.TrimSpace(query)
|
||||||
|
if query == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
fields := strings.Fields(query)
|
||||||
|
if len(fields) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return fields[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func tableFromRawQuery(query, operation string) string {
|
||||||
|
tokens := tokenizeQuery(query)
|
||||||
|
if len(tokens) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
switch operation {
|
||||||
|
case "SELECT":
|
||||||
|
return tokenAfter(tokens, "FROM")
|
||||||
|
case "INSERT":
|
||||||
|
return tokenAfter(tokens, "INTO")
|
||||||
|
case "UPDATE":
|
||||||
|
return tokenAfter(tokens, "UPDATE")
|
||||||
|
case "DELETE":
|
||||||
|
return tokenAfter(tokens, "FROM")
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func tokenAfter(tokens []string, keyword string) string {
|
||||||
|
for idx, token := range tokens {
|
||||||
|
if strings.EqualFold(token, keyword) && idx+1 < len(tokens) {
|
||||||
|
return cleanMetricIdentifier(tokens[idx+1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func tokenizeQuery(query string) []string {
|
||||||
|
replacer := strings.NewReplacer(
|
||||||
|
"\n", " ",
|
||||||
|
"\t", " ",
|
||||||
|
"(", " ",
|
||||||
|
")", " ",
|
||||||
|
",", " ",
|
||||||
|
)
|
||||||
|
return strings.Fields(replacer.Replace(query))
|
||||||
|
}
|
||||||
|
|
||||||
|
func cleanMetricIdentifier(value string) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
value = strings.Trim(value, "\"'`[]")
|
||||||
|
value = strings.TrimRight(value, ";")
|
||||||
|
return value
|
||||||
|
}
|
||||||
Generated
Vendored
+132
@@ -0,0 +1,132 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestHelper provides utilities for database testing
|
||||||
|
type TestHelper struct {
|
||||||
|
DB *sql.DB
|
||||||
|
Adapter *PgSQLAdapter
|
||||||
|
t *testing.T
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTestHelper creates a new test helper
|
||||||
|
func NewTestHelper(t *testing.T, db *sql.DB) *TestHelper {
|
||||||
|
return &TestHelper{
|
||||||
|
DB: db,
|
||||||
|
Adapter: NewPgSQLAdapter(db),
|
||||||
|
t: t,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CleanupTables truncates all test tables
|
||||||
|
func (h *TestHelper) CleanupTables() {
|
||||||
|
ctx := context.Background()
|
||||||
|
tables := []string{"comments", "posts", "users"}
|
||||||
|
|
||||||
|
for _, table := range tables {
|
||||||
|
_, err := h.DB.ExecContext(ctx, "TRUNCATE TABLE "+table+" CASCADE")
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// InsertUser inserts a test user and returns the ID
|
||||||
|
func (h *TestHelper) InsertUser(name, email string, age int) int {
|
||||||
|
ctx := context.Background()
|
||||||
|
result, err := h.Adapter.NewInsert().
|
||||||
|
Table("users").
|
||||||
|
Value("name", name).
|
||||||
|
Value("email", email).
|
||||||
|
Value("age", age).
|
||||||
|
Exec(ctx)
|
||||||
|
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
id, _ := result.LastInsertId()
|
||||||
|
return int(id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// InsertPost inserts a test post and returns the ID
|
||||||
|
func (h *TestHelper) InsertPost(userID int, title, content string, published bool) int {
|
||||||
|
ctx := context.Background()
|
||||||
|
result, err := h.Adapter.NewInsert().
|
||||||
|
Table("posts").
|
||||||
|
Value("user_id", userID).
|
||||||
|
Value("title", title).
|
||||||
|
Value("content", content).
|
||||||
|
Value("published", published).
|
||||||
|
Exec(ctx)
|
||||||
|
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
id, _ := result.LastInsertId()
|
||||||
|
return int(id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// InsertComment inserts a test comment and returns the ID
|
||||||
|
func (h *TestHelper) InsertComment(postID int, content string) int {
|
||||||
|
ctx := context.Background()
|
||||||
|
result, err := h.Adapter.NewInsert().
|
||||||
|
Table("comments").
|
||||||
|
Value("post_id", postID).
|
||||||
|
Value("content", content).
|
||||||
|
Exec(ctx)
|
||||||
|
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
id, _ := result.LastInsertId()
|
||||||
|
return int(id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AssertUserExists checks if a user exists by email
|
||||||
|
func (h *TestHelper) AssertUserExists(email string) {
|
||||||
|
ctx := context.Background()
|
||||||
|
exists, err := h.Adapter.NewSelect().
|
||||||
|
Table("users").
|
||||||
|
Where("email = ?", email).
|
||||||
|
Exists(ctx)
|
||||||
|
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
require.True(h.t, exists, "User with email %s should exist", email)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AssertUserCount asserts the number of users
|
||||||
|
func (h *TestHelper) AssertUserCount(expected int) {
|
||||||
|
ctx := context.Background()
|
||||||
|
count, err := h.Adapter.NewSelect().
|
||||||
|
Table("users").
|
||||||
|
Count(ctx)
|
||||||
|
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
require.Equal(h.t, expected, count)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserByEmail retrieves a user by email
|
||||||
|
func (h *TestHelper) GetUserByEmail(email string) map[string]interface{} {
|
||||||
|
ctx := context.Background()
|
||||||
|
var results []map[string]interface{}
|
||||||
|
err := h.Adapter.NewSelect().
|
||||||
|
Table("users").
|
||||||
|
Where("email = ?", email).
|
||||||
|
Scan(ctx, &results)
|
||||||
|
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
require.Len(h.t, results, 1, "Expected exactly one user with email %s", email)
|
||||||
|
return results[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeginTestTransaction starts a transaction for testing
|
||||||
|
func (h *TestHelper) BeginTestTransaction() (*PgSQLTxAdapter, func()) {
|
||||||
|
ctx := context.Background()
|
||||||
|
tx, err := h.DB.BeginTx(ctx, nil)
|
||||||
|
require.NoError(h.t, err)
|
||||||
|
|
||||||
|
adapter := &PgSQLTxAdapter{tx: tx}
|
||||||
|
cleanup := func() {
|
||||||
|
tx.Rollback()
|
||||||
|
}
|
||||||
|
|
||||||
|
return adapter, cleanup
|
||||||
|
}
|
||||||
+117
@@ -0,0 +1,117 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/uptrace/bun/dialect/mssqldialect"
|
||||||
|
"github.com/uptrace/bun/dialect/pgdialect"
|
||||||
|
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||||
|
"gorm.io/driver/postgres"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/driver/sqlserver"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PostgreSQL identifier length limit (63 bytes + null terminator = 64 bytes total)
|
||||||
|
const postgresIdentifierLimit = 63
|
||||||
|
|
||||||
|
// checkAliasLength checks if a preload relation path will generate aliases that exceed PostgreSQL's limit
|
||||||
|
// Returns true if the alias is likely to be truncated
|
||||||
|
func checkAliasLength(relation string) bool {
|
||||||
|
// Bun generates aliases like: parentalias__childalias__columnname
|
||||||
|
// For nested preloads, it uses the pattern: relation1__relation2__relation3__columnname
|
||||||
|
parts := strings.Split(relation, ".")
|
||||||
|
if len(parts) <= 1 {
|
||||||
|
return false // Single level relations are fine
|
||||||
|
}
|
||||||
|
|
||||||
|
// Calculate the actual alias prefix length that Bun will generate
|
||||||
|
// Bun uses double underscores (__) between each relation level
|
||||||
|
// and converts the relation names to lowercase with underscores
|
||||||
|
aliasPrefix := strings.ToLower(strings.Join(parts, "__"))
|
||||||
|
aliasPrefixLen := len(aliasPrefix)
|
||||||
|
|
||||||
|
// We need to add 2 more underscores for the column name separator plus column name length
|
||||||
|
// Column names in the error were things like "rid_mastertype_hubtype" (23 chars)
|
||||||
|
// To be safe, assume the longest column name could be around 35 chars
|
||||||
|
maxColumnNameLen := 35
|
||||||
|
estimatedMaxLen := aliasPrefixLen + 2 + maxColumnNameLen
|
||||||
|
|
||||||
|
// Check if this would exceed PostgreSQL's identifier limit
|
||||||
|
if estimatedMaxLen > postgresIdentifierLimit {
|
||||||
|
logger.Warn("Preload relation '%s' will generate aliases up to %d chars (prefix: %d + column: %d), exceeding PostgreSQL's %d char limit",
|
||||||
|
relation, estimatedMaxLen, aliasPrefixLen, maxColumnNameLen, postgresIdentifierLimit)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Also check if just the prefix is getting close (within 15 chars of limit)
|
||||||
|
// This gives room for column names
|
||||||
|
if aliasPrefixLen > (postgresIdentifierLimit - 15) {
|
||||||
|
logger.Warn("Preload relation '%s' has alias prefix of %d chars, which may cause truncation with longer column names (limit: %d)",
|
||||||
|
relation, aliasPrefixLen, postgresIdentifierLimit)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseTableName splits a table name that may contain schema into separate schema and table
|
||||||
|
// For example: "public.users" -> ("public", "users")
|
||||||
|
//
|
||||||
|
// "users" -> ("", "users")
|
||||||
|
//
|
||||||
|
// For SQLite, schema.table is translated to schema_table since SQLite doesn't support schemas
|
||||||
|
// in the same way as PostgreSQL/MSSQL
|
||||||
|
func parseTableName(fullTableName, driverName string) (schema, table string) {
|
||||||
|
if idx := strings.LastIndex(fullTableName, "."); idx != -1 {
|
||||||
|
schema = fullTableName[:idx]
|
||||||
|
table = fullTableName[idx+1:]
|
||||||
|
|
||||||
|
// For SQLite, convert schema.table to schema_table
|
||||||
|
if driverName == "sqlite" || driverName == "sqlite3" {
|
||||||
|
table = schema + "_" + table
|
||||||
|
schema = ""
|
||||||
|
}
|
||||||
|
return schema, table
|
||||||
|
}
|
||||||
|
return "", fullTableName
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPostgresDialect returns a Bun PostgreSQL dialect
|
||||||
|
func GetPostgresDialect() *pgdialect.Dialect {
|
||||||
|
return pgdialect.New()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSQLiteDialect returns a Bun SQLite dialect
|
||||||
|
func GetSQLiteDialect() *sqlitedialect.Dialect {
|
||||||
|
return sqlitedialect.New()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetMSSQLDialect returns a Bun MSSQL dialect
|
||||||
|
func GetMSSQLDialect() *mssqldialect.Dialect {
|
||||||
|
return mssqldialect.New()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPostgresDialector returns a GORM PostgreSQL dialector
|
||||||
|
func GetPostgresDialector(db *sql.DB) gorm.Dialector {
|
||||||
|
return postgres.New(postgres.Config{
|
||||||
|
Conn: db,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSQLiteDialector returns a GORM SQLite dialector
|
||||||
|
func GetSQLiteDialector(db *sql.DB) gorm.Dialector {
|
||||||
|
return sqlite.Dialector{
|
||||||
|
Conn: db,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetMSSQLDialector returns a GORM MSSQL dialector
|
||||||
|
func GetMSSQLDialector(db *sql.DB) gorm.Dialector {
|
||||||
|
return sqlserver.New(sqlserver.Config{
|
||||||
|
Conn: db,
|
||||||
|
})
|
||||||
|
}
|
||||||
+214
@@ -0,0 +1,214 @@
|
|||||||
|
package router
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/uptrace/bunrouter"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// BunRouterAdapter adapts uptrace/bunrouter to work with our Router interface
|
||||||
|
type BunRouterAdapter struct {
|
||||||
|
router *bunrouter.Router
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBunRouterAdapter creates a new bunrouter adapter
|
||||||
|
func NewBunRouterAdapter(router *bunrouter.Router) *BunRouterAdapter {
|
||||||
|
return &BunRouterAdapter{router: router}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBunRouterAdapterDefault creates a new bunrouter adapter with default router
|
||||||
|
func NewBunRouterAdapterDefault() *BunRouterAdapter {
|
||||||
|
return &BunRouterAdapter{router: bunrouter.New()}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterAdapter) HandleFunc(pattern string, handler common.HTTPHandlerFunc) common.RouteRegistration {
|
||||||
|
route := &BunRouterRegistration{
|
||||||
|
router: b.router,
|
||||||
|
pattern: pattern,
|
||||||
|
handler: handler,
|
||||||
|
}
|
||||||
|
return route
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterAdapter) ServeHTTP(w common.ResponseWriter, r common.Request) {
|
||||||
|
// This method would be used when we need to serve through our interface
|
||||||
|
// For now, we'll work directly with the underlying router
|
||||||
|
w.WriteHeader(http.StatusNotImplemented)
|
||||||
|
_, err := w.Write([]byte(`{"error":"ServeHTTP not implemented - use GetBunRouter() for direct access"}`))
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Failed to write. %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetBunRouter returns the underlying bunrouter for direct access
|
||||||
|
func (b *BunRouterAdapter) GetBunRouter() *bunrouter.Router {
|
||||||
|
return b.router
|
||||||
|
}
|
||||||
|
|
||||||
|
// BunRouterRegistration implements RouteRegistration for bunrouter
|
||||||
|
type BunRouterRegistration struct {
|
||||||
|
router *bunrouter.Router
|
||||||
|
pattern string
|
||||||
|
handler common.HTTPHandlerFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRegistration) Methods(methods ...string) common.RouteRegistration {
|
||||||
|
// bunrouter handles methods differently - we'll register for each method
|
||||||
|
for _, method := range methods {
|
||||||
|
b.router.Handle(method, b.pattern, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
// Convert bunrouter.Request to our BunRouterRequest
|
||||||
|
reqAdapter := &BunRouterRequest{req: req}
|
||||||
|
respAdapter := &HTTPResponseWriter{resp: w}
|
||||||
|
b.handler(respAdapter, reqAdapter)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRegistration) PathPrefix(prefix string) common.RouteRegistration {
|
||||||
|
// bunrouter doesn't have PathPrefix like mux, but we can modify the pattern
|
||||||
|
newPattern := prefix + b.pattern
|
||||||
|
b.pattern = newPattern
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// BunRouterRequest adapts bunrouter.Request to our Request interface
|
||||||
|
type BunRouterRequest struct {
|
||||||
|
req bunrouter.Request
|
||||||
|
body []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBunRouterRequest creates a new BunRouterRequest adapter
|
||||||
|
func NewBunRouterRequest(req bunrouter.Request) *BunRouterRequest {
|
||||||
|
return &BunRouterRequest{req: req}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) Method() string {
|
||||||
|
return b.req.Method
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) URL() string {
|
||||||
|
return b.req.URL.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) Header(key string) string {
|
||||||
|
return b.req.Header.Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) Body() ([]byte, error) {
|
||||||
|
if b.body != nil {
|
||||||
|
return b.body, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if b.req.Body == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create HTTPRequest adapter and use its Body() method
|
||||||
|
httpAdapter := NewHTTPRequest(b.req.Request)
|
||||||
|
body, err := httpAdapter.Body()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
b.body = body
|
||||||
|
return body, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) PathParam(key string) string {
|
||||||
|
return b.req.Param(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) QueryParam(key string) string {
|
||||||
|
return b.req.URL.Query().Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) AllQueryParams() map[string]string {
|
||||||
|
params := make(map[string]string)
|
||||||
|
for key, values := range b.req.URL.Query() {
|
||||||
|
if len(values) > 0 {
|
||||||
|
params[key] = values[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return params
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *BunRouterRequest) AllHeaders() map[string]string {
|
||||||
|
headers := make(map[string]string)
|
||||||
|
for key, values := range b.req.Header {
|
||||||
|
if len(values) > 0 {
|
||||||
|
headers[key] = values[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return headers
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnderlyingRequest returns the underlying *http.Request
|
||||||
|
// This is useful when you need to pass the request to other handlers
|
||||||
|
func (b *BunRouterRequest) UnderlyingRequest() *http.Request {
|
||||||
|
return b.req.Request
|
||||||
|
}
|
||||||
|
|
||||||
|
// StandardBunRouterAdapter creates routes compatible with standard bunrouter handlers
|
||||||
|
type StandardBunRouterAdapter struct {
|
||||||
|
*BunRouterAdapter
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewStandardBunRouterAdapter() *StandardBunRouterAdapter {
|
||||||
|
return &StandardBunRouterAdapter{
|
||||||
|
BunRouterAdapter: NewBunRouterAdapterDefault(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterRoute registers a route that works with the existing Handler
|
||||||
|
func (s *StandardBunRouterAdapter) RegisterRoute(method, pattern string, handler func(http.ResponseWriter, *http.Request, map[string]string)) {
|
||||||
|
s.router.Handle(method, pattern, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
// Extract path parameters
|
||||||
|
params := make(map[string]string)
|
||||||
|
|
||||||
|
// bunrouter doesn't provide a direct way to get all params
|
||||||
|
// You would typically access them individually with req.Param("name")
|
||||||
|
// For this example, we'll create the map based on the request context
|
||||||
|
|
||||||
|
handler(w, req.Request, params)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterRouteWithParams registers a route with explicit parameter extraction
|
||||||
|
func (s *StandardBunRouterAdapter) RegisterRouteWithParams(method, pattern string, paramNames []string, handler func(http.ResponseWriter, *http.Request, map[string]string)) {
|
||||||
|
s.router.Handle(method, pattern, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
// Extract specified path parameters
|
||||||
|
params := make(map[string]string)
|
||||||
|
for _, paramName := range paramNames {
|
||||||
|
params[paramName] = req.Param(paramName)
|
||||||
|
}
|
||||||
|
|
||||||
|
handler(w, req.Request, params)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// BunRouterConfig holds bunrouter-specific configuration
|
||||||
|
type BunRouterConfig struct {
|
||||||
|
UseStrictSlash bool
|
||||||
|
RedirectTrailingSlash bool
|
||||||
|
HandleMethodNotAllowed bool
|
||||||
|
HandleOPTIONS bool
|
||||||
|
GlobalOPTIONS http.Handler
|
||||||
|
GlobalMethodNotAllowed http.Handler
|
||||||
|
PanicHandler func(http.ResponseWriter, *http.Request, interface{})
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultBunRouterConfig returns default bunrouter configuration
|
||||||
|
func DefaultBunRouterConfig() *BunRouterConfig {
|
||||||
|
return &BunRouterConfig{
|
||||||
|
UseStrictSlash: false,
|
||||||
|
RedirectTrailingSlash: true,
|
||||||
|
HandleMethodNotAllowed: true,
|
||||||
|
HandleOPTIONS: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
+238
@@ -0,0 +1,238 @@
|
|||||||
|
package router
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MuxAdapter adapts Gorilla Mux to work with our Router interface
|
||||||
|
type MuxAdapter struct {
|
||||||
|
router *mux.Router
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMuxAdapter creates a new Mux adapter
|
||||||
|
func NewMuxAdapter(router *mux.Router) *MuxAdapter {
|
||||||
|
return &MuxAdapter{router: router}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MuxAdapter) HandleFunc(pattern string, handler common.HTTPHandlerFunc) common.RouteRegistration {
|
||||||
|
route := &MuxRouteRegistration{
|
||||||
|
router: m.router,
|
||||||
|
pattern: pattern,
|
||||||
|
handler: handler,
|
||||||
|
}
|
||||||
|
return route
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MuxAdapter) ServeHTTP(w common.ResponseWriter, r common.Request) {
|
||||||
|
// This method would be used when we need to serve through our interface
|
||||||
|
// For now, we'll work directly with the underlying router
|
||||||
|
w.WriteHeader(http.StatusNotImplemented)
|
||||||
|
_, err := w.Write([]byte(`{"error":"ServeHTTP not implemented - use GetMuxRouter() for direct access"}`))
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Failed to write. %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MuxRouteRegistration implements RouteRegistration for Mux
|
||||||
|
type MuxRouteRegistration struct {
|
||||||
|
router *mux.Router
|
||||||
|
pattern string
|
||||||
|
handler common.HTTPHandlerFunc
|
||||||
|
route *mux.Route
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MuxRouteRegistration) Methods(methods ...string) common.RouteRegistration {
|
||||||
|
if m.route == nil {
|
||||||
|
m.route = m.router.HandleFunc(m.pattern, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
reqAdapter := &HTTPRequest{req: r, vars: mux.Vars(r)}
|
||||||
|
respAdapter := &HTTPResponseWriter{resp: w}
|
||||||
|
m.handler(respAdapter, reqAdapter)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
m.route.Methods(methods...)
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MuxRouteRegistration) PathPrefix(prefix string) common.RouteRegistration {
|
||||||
|
if m.route == nil {
|
||||||
|
m.route = m.router.HandleFunc(m.pattern, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
reqAdapter := &HTTPRequest{req: r, vars: mux.Vars(r)}
|
||||||
|
respAdapter := &HTTPResponseWriter{resp: w}
|
||||||
|
m.handler(respAdapter, reqAdapter)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
m.route.PathPrefix(prefix)
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// HTTPRequest adapts standard http.Request to our Request interface
|
||||||
|
type HTTPRequest struct {
|
||||||
|
req *http.Request
|
||||||
|
vars map[string]string
|
||||||
|
body []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewHTTPRequest(r *http.Request) *HTTPRequest {
|
||||||
|
return &HTTPRequest{
|
||||||
|
req: r,
|
||||||
|
vars: make(map[string]string),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) Method() string {
|
||||||
|
return h.req.Method
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) URL() string {
|
||||||
|
return h.req.URL.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) Header(key string) string {
|
||||||
|
return h.req.Header.Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) Body() ([]byte, error) {
|
||||||
|
if h.body != nil {
|
||||||
|
return h.body, nil
|
||||||
|
}
|
||||||
|
if h.req.Body == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
defer h.req.Body.Close()
|
||||||
|
body, err := io.ReadAll(h.req.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
h.body = body
|
||||||
|
return body, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) PathParam(key string) string {
|
||||||
|
return h.vars[key]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) QueryParam(key string) string {
|
||||||
|
return h.req.URL.Query().Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) AllQueryParams() map[string]string {
|
||||||
|
params := make(map[string]string)
|
||||||
|
for key, values := range h.req.URL.Query() {
|
||||||
|
if len(values) > 0 {
|
||||||
|
params[key] = values[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return params
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPRequest) AllHeaders() map[string]string {
|
||||||
|
headers := make(map[string]string)
|
||||||
|
for key, values := range h.req.Header {
|
||||||
|
if len(values) > 0 {
|
||||||
|
headers[key] = values[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return headers
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnderlyingRequest returns the underlying *http.Request
|
||||||
|
// This is useful when you need to pass the request to other handlers
|
||||||
|
func (h *HTTPRequest) UnderlyingRequest() *http.Request {
|
||||||
|
return h.req
|
||||||
|
}
|
||||||
|
|
||||||
|
// HTTPResponseWriter adapts our ResponseWriter interface to standard http.ResponseWriter
|
||||||
|
type HTTPResponseWriter struct {
|
||||||
|
resp http.ResponseWriter
|
||||||
|
w common.ResponseWriter //nolint:unused
|
||||||
|
status int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewHTTPResponseWriter(w http.ResponseWriter) *HTTPResponseWriter {
|
||||||
|
return &HTTPResponseWriter{resp: w}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPResponseWriter) SetHeader(key, value string) {
|
||||||
|
h.resp.Header().Set(key, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPResponseWriter) WriteHeader(statusCode int) {
|
||||||
|
h.status = statusCode
|
||||||
|
h.resp.WriteHeader(statusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPResponseWriter) Write(data []byte) (int, error) {
|
||||||
|
return h.resp.Write(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *HTTPResponseWriter) WriteJSON(data interface{}) error {
|
||||||
|
h.SetHeader("Content-Type", "application/json")
|
||||||
|
enc := json.NewEncoder(h.resp)
|
||||||
|
enc.SetEscapeHTML(false)
|
||||||
|
return enc.Encode(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnderlyingResponseWriter returns the underlying http.ResponseWriter
|
||||||
|
// This is useful when you need to pass the response writer to other handlers
|
||||||
|
func (h *HTTPResponseWriter) UnderlyingResponseWriter() http.ResponseWriter {
|
||||||
|
return h.resp
|
||||||
|
}
|
||||||
|
|
||||||
|
// StandardMuxAdapter creates routes compatible with standard http.HandlerFunc
|
||||||
|
type StandardMuxAdapter struct {
|
||||||
|
*MuxAdapter
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewStandardMuxAdapter() *StandardMuxAdapter {
|
||||||
|
return &StandardMuxAdapter{
|
||||||
|
MuxAdapter: NewMuxAdapter(mux.NewRouter()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterRoute registers a route that works with the existing Handler
|
||||||
|
func (s *StandardMuxAdapter) RegisterRoute(pattern string, handler func(http.ResponseWriter, *http.Request, map[string]string)) *mux.Route {
|
||||||
|
return s.router.HandleFunc(pattern, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
vars := mux.Vars(r)
|
||||||
|
handler(w, r, vars)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetMuxRouter returns the underlying mux router for direct access
|
||||||
|
func (s *StandardMuxAdapter) GetMuxRouter() *mux.Router {
|
||||||
|
return s.router
|
||||||
|
}
|
||||||
|
|
||||||
|
// PathParamExtractor extracts path parameters from different router types
|
||||||
|
type PathParamExtractor interface {
|
||||||
|
ExtractParams(*http.Request) map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// MuxParamExtractor extracts parameters from Gorilla Mux
|
||||||
|
type MuxParamExtractor struct{}
|
||||||
|
|
||||||
|
func (m MuxParamExtractor) ExtractParams(r *http.Request) map[string]string {
|
||||||
|
return mux.Vars(r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RouterConfig holds router configuration
|
||||||
|
type RouterConfig struct {
|
||||||
|
PathPrefix string
|
||||||
|
Middleware []func(http.Handler) http.Handler
|
||||||
|
ParamExtractor PathParamExtractor
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultRouterConfig returns default router configuration
|
||||||
|
func DefaultRouterConfig() *RouterConfig {
|
||||||
|
return &RouterConfig{
|
||||||
|
PathPrefix: "",
|
||||||
|
Middleware: make([]func(http.Handler) http.Handler, 0),
|
||||||
|
ParamExtractor: MuxParamExtractor{},
|
||||||
|
}
|
||||||
|
}
|
||||||
+149
@@ -0,0 +1,149 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CORSConfig holds CORS configuration
|
||||||
|
type CORSConfig struct {
|
||||||
|
AllowedOrigins []string
|
||||||
|
AllowedMethods []string
|
||||||
|
AllowedHeaders []string
|
||||||
|
MaxAge int
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultCORSConfig returns a default CORS configuration suitable for HeadSpec
|
||||||
|
func DefaultCORSConfig() CORSConfig {
|
||||||
|
configManager := config.GetConfigManager()
|
||||||
|
cfg, _ := configManager.GetConfig()
|
||||||
|
hosts := make([]string, 0)
|
||||||
|
// hosts = append(hosts, "*")
|
||||||
|
|
||||||
|
_, _, ipsList := config.GetIPs()
|
||||||
|
|
||||||
|
for i := range cfg.Servers.Instances {
|
||||||
|
server := cfg.Servers.Instances[i]
|
||||||
|
if server.Port == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
hosts = append(hosts, server.ExternalURLs...)
|
||||||
|
hosts = append(hosts, fmt.Sprintf("http://%s:%d", server.Host, server.Port))
|
||||||
|
hosts = append(hosts, fmt.Sprintf("https://%s:%d", server.Host, server.Port))
|
||||||
|
hosts = append(hosts, fmt.Sprintf("http://%s:%d", "localhost", server.Port))
|
||||||
|
for _, ip := range ipsList {
|
||||||
|
hosts = append(hosts, fmt.Sprintf("http://%s:%d", ip.String(), server.Port))
|
||||||
|
hosts = append(hosts, fmt.Sprintf("https://%s:%d", ip.String(), server.Port))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return CORSConfig{
|
||||||
|
AllowedOrigins: hosts,
|
||||||
|
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
|
||||||
|
AllowedHeaders: GetHeadSpecHeaders(),
|
||||||
|
MaxAge: 86400, // 24 hours
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetHeadSpecHeaders returns all headers used by HeadSpec
|
||||||
|
func GetHeadSpecHeaders() []string {
|
||||||
|
return []string{
|
||||||
|
// Standard headers
|
||||||
|
"Content-Type",
|
||||||
|
"Authorization",
|
||||||
|
"Accept",
|
||||||
|
"Accept-Language",
|
||||||
|
"Content-Language",
|
||||||
|
|
||||||
|
// Field Selection
|
||||||
|
"X-Select-Fields",
|
||||||
|
"X-Not-Select-Fields",
|
||||||
|
"X-Clean-JSON",
|
||||||
|
|
||||||
|
// Filtering & Search
|
||||||
|
"X-FieldFilter-*",
|
||||||
|
"X-SearchFilter-*",
|
||||||
|
"X-SearchOp-*",
|
||||||
|
"X-SearchOr-*",
|
||||||
|
"X-SearchAnd-*",
|
||||||
|
"X-SearchCols",
|
||||||
|
"X-Custom-SQL-W",
|
||||||
|
"X-Custom-SQL-W-*",
|
||||||
|
"X-Custom-SQL-Or",
|
||||||
|
"X-Custom-SQL-Or-*",
|
||||||
|
|
||||||
|
// Joins & Relations
|
||||||
|
"X-Preload",
|
||||||
|
"X-Preload-*",
|
||||||
|
"X-Expand",
|
||||||
|
"X-Expand-*",
|
||||||
|
"X-Custom-SQL-Join",
|
||||||
|
"X-Custom-SQL-Join-*",
|
||||||
|
|
||||||
|
// Sorting & Pagination
|
||||||
|
"X-Sort",
|
||||||
|
"X-Sort-*",
|
||||||
|
"X-Limit",
|
||||||
|
"X-Offset",
|
||||||
|
"X-Cursor-Forward",
|
||||||
|
"X-Cursor-Backward",
|
||||||
|
|
||||||
|
// Advanced Features
|
||||||
|
"X-AdvSQL-*",
|
||||||
|
"X-CQL-Sel-*",
|
||||||
|
"X-Distinct",
|
||||||
|
"X-SkipCount",
|
||||||
|
"X-SkipCache",
|
||||||
|
"X-Fetch-RowNumber",
|
||||||
|
"X-PKRow",
|
||||||
|
|
||||||
|
// Response Format
|
||||||
|
"X-SimpleAPI",
|
||||||
|
"X-DetailAPI",
|
||||||
|
"X-Syncfusion",
|
||||||
|
"X-Single-Record-As-Object",
|
||||||
|
|
||||||
|
// Transaction Control
|
||||||
|
"X-Transaction-Atomic",
|
||||||
|
|
||||||
|
// X-Files - comprehensive JSON configuration
|
||||||
|
"X-Files",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCORSHeaders sets CORS headers on a response writer
|
||||||
|
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
||||||
|
// Set allowed origins
|
||||||
|
// if len(config.AllowedOrigins) > 0 {
|
||||||
|
// w.SetHeader("Access-Control-Allow-Origin", strings.Join(config.AllowedOrigins, ", "))
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Todo origin list parsing
|
||||||
|
w.SetHeader("Access-Control-Allow-Origin", "*")
|
||||||
|
|
||||||
|
// Set allowed methods
|
||||||
|
if len(config.AllowedMethods) > 0 {
|
||||||
|
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set allowed headers
|
||||||
|
// if len(config.AllowedHeaders) > 0 {
|
||||||
|
// w.SetHeader("Access-Control-Allow-Headers", strings.Join(config.AllowedHeaders, ", "))
|
||||||
|
// }
|
||||||
|
w.SetHeader("Access-Control-Allow-Headers", "*")
|
||||||
|
|
||||||
|
// Set max age
|
||||||
|
if config.MaxAge > 0 {
|
||||||
|
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allow credentials
|
||||||
|
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||||
|
|
||||||
|
// Expose headers that clients can read
|
||||||
|
exposeHeaders := config.AllowedHeaders
|
||||||
|
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
|
||||||
|
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
|
||||||
|
}
|
||||||
+97
@@ -0,0 +1,97 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
// Example showing how to use the common handler interfaces
|
||||||
|
// This file demonstrates the handler interface hierarchy and usage patterns
|
||||||
|
|
||||||
|
// ProcessWithAnyHandler demonstrates using the base SpecHandler interface
|
||||||
|
// which works with any handler type (resolvespec, restheadspec, or funcspec)
|
||||||
|
func ProcessWithAnyHandler(handler SpecHandler) Database {
|
||||||
|
// All handlers expose GetDatabase() through the SpecHandler interface
|
||||||
|
return handler.GetDatabase()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessCRUDRequest demonstrates using the CRUDHandler interface
|
||||||
|
// which works with resolvespec.Handler and restheadspec.Handler
|
||||||
|
func ProcessCRUDRequest(handler CRUDHandler, w ResponseWriter, r Request, params map[string]string) {
|
||||||
|
// Both resolvespec and restheadspec handlers implement Handle()
|
||||||
|
handler.Handle(w, r, params)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessMetadataRequest demonstrates getting metadata from CRUD handlers
|
||||||
|
func ProcessMetadataRequest(handler CRUDHandler, w ResponseWriter, r Request, params map[string]string) {
|
||||||
|
// Both resolvespec and restheadspec handlers implement HandleGet()
|
||||||
|
handler.HandleGet(w, r, params)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example usage patterns (not executable, just for documentation):
|
||||||
|
/*
|
||||||
|
// Example 1: Using with resolvespec.Handler
|
||||||
|
func ExampleResolveSpec() {
|
||||||
|
db := // ... get database
|
||||||
|
registry := // ... get registry
|
||||||
|
|
||||||
|
handler := resolvespec.NewHandler(db, registry)
|
||||||
|
|
||||||
|
// Can be used as SpecHandler
|
||||||
|
var specHandler SpecHandler = handler
|
||||||
|
database := specHandler.GetDatabase()
|
||||||
|
|
||||||
|
// Can be used as CRUDHandler
|
||||||
|
var crudHandler CRUDHandler = handler
|
||||||
|
crudHandler.Handle(w, r, params)
|
||||||
|
crudHandler.HandleGet(w, r, params)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 2: Using with restheadspec.Handler
|
||||||
|
func ExampleRestHeadSpec() {
|
||||||
|
db := // ... get database
|
||||||
|
registry := // ... get registry
|
||||||
|
|
||||||
|
handler := restheadspec.NewHandler(db, registry)
|
||||||
|
|
||||||
|
// Can be used as SpecHandler
|
||||||
|
var specHandler SpecHandler = handler
|
||||||
|
database := specHandler.GetDatabase()
|
||||||
|
|
||||||
|
// Can be used as CRUDHandler
|
||||||
|
var crudHandler CRUDHandler = handler
|
||||||
|
crudHandler.Handle(w, r, params)
|
||||||
|
crudHandler.HandleGet(w, r, params)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 3: Using with funcspec.Handler
|
||||||
|
func ExampleFuncSpec() {
|
||||||
|
db := // ... get database
|
||||||
|
|
||||||
|
handler := funcspec.NewHandler(db)
|
||||||
|
|
||||||
|
// Can be used as SpecHandler
|
||||||
|
var specHandler SpecHandler = handler
|
||||||
|
database := specHandler.GetDatabase()
|
||||||
|
|
||||||
|
// Can be used as QueryHandler
|
||||||
|
var queryHandler QueryHandler = handler
|
||||||
|
// funcspec has different methods: SqlQueryList() and SqlQuery()
|
||||||
|
// which return HTTP handler functions
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 4: Polymorphic handler processing
|
||||||
|
func ProcessHandlers(handlers []SpecHandler) {
|
||||||
|
for _, handler := range handlers {
|
||||||
|
// All handlers expose the database
|
||||||
|
db := handler.GetDatabase()
|
||||||
|
|
||||||
|
// Type switch for specific handler types
|
||||||
|
switch h := handler.(type) {
|
||||||
|
case CRUDHandler:
|
||||||
|
// This is resolvespec or restheadspec
|
||||||
|
// Can call Handle() and HandleGet()
|
||||||
|
_ = h
|
||||||
|
case QueryHandler:
|
||||||
|
// This is funcspec
|
||||||
|
// Can call SqlQueryList() and SqlQuery()
|
||||||
|
_ = h
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
*/
|
||||||
+309
@@ -0,0 +1,309 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ValidateAndUnwrapModelResult contains the result of model validation
|
||||||
|
type ValidateAndUnwrapModelResult struct {
|
||||||
|
ModelType reflect.Type
|
||||||
|
Model interface{}
|
||||||
|
ModelPtr interface{}
|
||||||
|
OriginalType reflect.Type
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateAndUnwrapModel validates that a model is a struct type and unwraps
|
||||||
|
// pointers, slices, and arrays to get to the base struct type.
|
||||||
|
// Returns an error if the model is not a valid struct type.
|
||||||
|
func ValidateAndUnwrapModel(model interface{}) (*ValidateAndUnwrapModelResult, error) {
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
originalType := modelType
|
||||||
|
|
||||||
|
// Unwrap pointers, slices, and arrays to get to the base struct type
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate that we have a struct type
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
return nil, fmt.Errorf("model must be a struct type, got %v. Ensure you register the struct (e.g., ModelCoreAccount{}) not a slice (e.g., []*ModelCoreAccount)", originalType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If the registered model was a pointer or slice, use the unwrapped struct type
|
||||||
|
if originalType != modelType {
|
||||||
|
model = reflect.New(modelType).Elem().Interface()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create a pointer to the model type for database operations
|
||||||
|
modelPtr := reflect.New(reflect.TypeOf(model)).Interface()
|
||||||
|
|
||||||
|
return &ValidateAndUnwrapModelResult{
|
||||||
|
ModelType: modelType,
|
||||||
|
Model: model,
|
||||||
|
ModelPtr: modelPtr,
|
||||||
|
OriginalType: originalType,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtractTagValue extracts the value for a given key from a struct tag string.
|
||||||
|
// It handles both semicolon and comma-separated tag formats (e.g., GORM and BUN tags).
|
||||||
|
// For tags like "json:name;validate:required" it will extract "name" for key "json".
|
||||||
|
// For tags like "rel:has-many,join:table" it will extract "table" for key "join".
|
||||||
|
func ExtractTagValue(tag, key string) string {
|
||||||
|
// Split by both semicolons and commas to handle different tag formats
|
||||||
|
// We need to be smart about this - commas can be part of values
|
||||||
|
// So we'll try semicolon first, then comma if needed
|
||||||
|
separators := []string{";", ","}
|
||||||
|
|
||||||
|
for _, sep := range separators {
|
||||||
|
parts := strings.Split(tag, sep)
|
||||||
|
for _, part := range parts {
|
||||||
|
part = strings.TrimSpace(part)
|
||||||
|
if strings.HasPrefix(part, key+":") {
|
||||||
|
return strings.TrimPrefix(part, key+":")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRelationshipInfo analyzes a model type and extracts relationship metadata
|
||||||
|
// for a specific relation field identified by its JSON name.
|
||||||
|
// Returns nil if the field is not found or is not a valid relationship.
|
||||||
|
func GetRelationshipInfo(modelType reflect.Type, relationName string) *RelationshipInfo {
|
||||||
|
// Ensure we have a struct type
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
logger.Warn("Cannot get relationship info from non-struct type: %v", modelType)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
|
field := modelType.Field(i)
|
||||||
|
jsonTag := field.Tag.Get("json")
|
||||||
|
jsonName := strings.Split(jsonTag, ",")[0]
|
||||||
|
|
||||||
|
if jsonName == relationName {
|
||||||
|
gormTag := field.Tag.Get("gorm")
|
||||||
|
bunTag := field.Tag.Get("bun")
|
||||||
|
info := &RelationshipInfo{
|
||||||
|
FieldName: field.Name,
|
||||||
|
JSONName: jsonName,
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(bunTag, "rel:") || strings.Contains(bunTag, "join:") {
|
||||||
|
//bun:"rel:has-many,join:rid_hub=rid_hub_division"
|
||||||
|
if strings.Contains(bunTag, "has-many") {
|
||||||
|
info.RelationType = "hasMany"
|
||||||
|
} else if strings.Contains(bunTag, "has-one") {
|
||||||
|
info.RelationType = "hasOne"
|
||||||
|
} else if strings.Contains(bunTag, "belongs-to") {
|
||||||
|
info.RelationType = "belongsTo"
|
||||||
|
} else if strings.Contains(bunTag, "many-to-many") {
|
||||||
|
info.RelationType = "many2many"
|
||||||
|
} else {
|
||||||
|
info.RelationType = "hasOne"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract join info
|
||||||
|
joinPart := ExtractTagValue(bunTag, "join")
|
||||||
|
if joinPart != "" && info.RelationType == "many2many" {
|
||||||
|
// For many2many, the join part is the join table name
|
||||||
|
info.JoinTable = joinPart
|
||||||
|
} else if joinPart != "" {
|
||||||
|
// For other relations, parse foreignKey and references
|
||||||
|
joinParts := strings.Split(joinPart, "=")
|
||||||
|
if len(joinParts) == 2 {
|
||||||
|
info.ForeignKey = joinParts[0]
|
||||||
|
info.References = joinParts[1]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get related model type
|
||||||
|
if field.Type.Kind() == reflect.Slice {
|
||||||
|
elemType := field.Type.Elem()
|
||||||
|
if elemType.Kind() == reflect.Pointer {
|
||||||
|
elemType = elemType.Elem()
|
||||||
|
}
|
||||||
|
if elemType.Kind() == reflect.Struct {
|
||||||
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
|
}
|
||||||
|
} else if field.Type.Kind() == reflect.Pointer || field.Type.Kind() == reflect.Struct {
|
||||||
|
elemType := field.Type
|
||||||
|
if elemType.Kind() == reflect.Pointer {
|
||||||
|
elemType = elemType.Elem()
|
||||||
|
}
|
||||||
|
if elemType.Kind() == reflect.Struct {
|
||||||
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return info
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse GORM tag to determine relationship type and keys
|
||||||
|
if strings.Contains(gormTag, "foreignKey") {
|
||||||
|
info.ForeignKey = ExtractTagValue(gormTag, "foreignKey")
|
||||||
|
info.References = ExtractTagValue(gormTag, "references")
|
||||||
|
|
||||||
|
// Determine if it's belongsTo or hasMany/hasOne
|
||||||
|
if field.Type.Kind() == reflect.Slice {
|
||||||
|
info.RelationType = "hasMany"
|
||||||
|
// Get the element type for slice
|
||||||
|
elemType := field.Type.Elem()
|
||||||
|
if elemType.Kind() == reflect.Pointer {
|
||||||
|
elemType = elemType.Elem()
|
||||||
|
}
|
||||||
|
if elemType.Kind() == reflect.Struct {
|
||||||
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
|
}
|
||||||
|
} else if field.Type.Kind() == reflect.Pointer || field.Type.Kind() == reflect.Struct {
|
||||||
|
info.RelationType = "belongsTo"
|
||||||
|
elemType := field.Type
|
||||||
|
if elemType.Kind() == reflect.Pointer {
|
||||||
|
elemType = elemType.Elem()
|
||||||
|
}
|
||||||
|
if elemType.Kind() == reflect.Struct {
|
||||||
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else if strings.Contains(gormTag, "many2many") {
|
||||||
|
info.RelationType = "many2many"
|
||||||
|
info.JoinTable = ExtractTagValue(gormTag, "many2many")
|
||||||
|
// Get the element type for many2many (always slice)
|
||||||
|
if field.Type.Kind() == reflect.Slice {
|
||||||
|
elemType := field.Type.Elem()
|
||||||
|
if elemType.Kind() == reflect.Pointer {
|
||||||
|
elemType = elemType.Elem()
|
||||||
|
}
|
||||||
|
if elemType.Kind() == reflect.Struct {
|
||||||
|
info.RelatedModel = reflect.New(elemType).Elem().Interface()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Field has no GORM relationship tags, so it's not a relation
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return info
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RelationPathToBunAlias converts a relation path (e.g., "Order.Customer") to a Bun alias format.
|
||||||
|
// It converts to lowercase and replaces dots with double underscores.
|
||||||
|
// For example: "Order.Customer" -> "order__customer"
|
||||||
|
func RelationPathToBunAlias(relationPath string) string {
|
||||||
|
if relationPath == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
// Convert to lowercase and replace dots with double underscores
|
||||||
|
alias := strings.ToLower(relationPath)
|
||||||
|
alias = strings.ReplaceAll(alias, ".", "__")
|
||||||
|
return alias
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReplaceTableReferencesInSQL replaces references to a base table name in a SQL expression
|
||||||
|
// with the appropriate alias for the current preload level.
|
||||||
|
// For example, if baseTableName is "mastertaskitem" and targetAlias is "mal__mal",
|
||||||
|
// it will replace "mastertaskitem.rid_mastertaskitem" with "mal__mal.rid_mastertaskitem"
|
||||||
|
func ReplaceTableReferencesInSQL(sqlExpr, baseTableName, targetAlias string) string {
|
||||||
|
if sqlExpr == "" || baseTableName == "" || targetAlias == "" {
|
||||||
|
return sqlExpr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Replace both quoted and unquoted table references
|
||||||
|
// Handle patterns like: tablename.column, "tablename".column, tablename."column", "tablename"."column"
|
||||||
|
|
||||||
|
// Pattern 1: tablename.column (unquoted)
|
||||||
|
result := strings.ReplaceAll(sqlExpr, baseTableName+".", targetAlias+".")
|
||||||
|
|
||||||
|
// Pattern 2: "tablename".column or "tablename"."column" (quoted table name)
|
||||||
|
result = strings.ReplaceAll(result, "\""+baseTableName+"\".", "\""+targetAlias+"\".")
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetTableNameFromModel extracts the table name from a model.
|
||||||
|
// It checks the bun tag first, then falls back to converting the struct name to snake_case.
|
||||||
|
func GetTableNameFromModel(model interface{}) string {
|
||||||
|
if model == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
|
||||||
|
// Unwrap pointers
|
||||||
|
for modelType != nil && modelType.Kind() == reflect.Pointer {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Look for bun tag on embedded BaseModel
|
||||||
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
|
field := modelType.Field(i)
|
||||||
|
if field.Anonymous {
|
||||||
|
bunTag := field.Tag.Get("bun")
|
||||||
|
if strings.HasPrefix(bunTag, "table:") {
|
||||||
|
return strings.TrimPrefix(bunTag, "table:")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: convert struct name to lowercase (simple heuristic)
|
||||||
|
// This handles cases like "MasterTaskItem" -> "mastertaskitem"
|
||||||
|
return strings.ToLower(modelType.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConvertSliceForBun converts []interface{} values to PostgreSQL array literal strings.
|
||||||
|
// BUN's fallback appender for []interface{} is JSON encoding, which produces "[]" —
|
||||||
|
// invalid PostgreSQL array syntax. PostgreSQL expects "{}" for empty arrays and
|
||||||
|
// "{elem1,elem2}" for non-empty ones. All other value types are returned unchanged.
|
||||||
|
func ConvertSliceForBun(value interface{}) interface{} {
|
||||||
|
arr, ok := value.([]interface{})
|
||||||
|
if !ok {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
if len(arr) == 0 {
|
||||||
|
return "{}"
|
||||||
|
}
|
||||||
|
parts := make([]string, len(arr))
|
||||||
|
for i, elem := range arr {
|
||||||
|
switch e := elem.(type) {
|
||||||
|
case string:
|
||||||
|
needsQuote := e == "" || strings.ContainsAny(e, `,"\\{}`+"\t\n\r ")
|
||||||
|
if needsQuote {
|
||||||
|
e = strings.ReplaceAll(e, `\`, `\\`)
|
||||||
|
e = strings.ReplaceAll(e, `"`, `""`)
|
||||||
|
parts[i] = `"` + e + `"`
|
||||||
|
} else {
|
||||||
|
parts[i] = e
|
||||||
|
}
|
||||||
|
case float64:
|
||||||
|
if e == float64(int64(e)) {
|
||||||
|
parts[i] = strconv.FormatInt(int64(e), 10)
|
||||||
|
} else {
|
||||||
|
parts[i] = strconv.FormatFloat(e, 'f', -1, 64)
|
||||||
|
}
|
||||||
|
case bool:
|
||||||
|
if e {
|
||||||
|
parts[i] = "t"
|
||||||
|
} else {
|
||||||
|
parts[i] = "f"
|
||||||
|
}
|
||||||
|
case nil:
|
||||||
|
parts[i] = "NULL"
|
||||||
|
default:
|
||||||
|
parts[i] = fmt.Sprintf("%v", e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "{" + strings.Join(parts, ",") + "}"
|
||||||
|
}
|
||||||
+311
@@ -0,0 +1,311 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Database interface designed to work with both GORM and Bun
|
||||||
|
type Database interface {
|
||||||
|
// Core query operations
|
||||||
|
NewSelect() SelectQuery
|
||||||
|
NewInsert() InsertQuery
|
||||||
|
NewUpdate() UpdateQuery
|
||||||
|
NewDelete() DeleteQuery
|
||||||
|
|
||||||
|
// Raw SQL execution
|
||||||
|
Exec(ctx context.Context, query string, args ...interface{}) (Result, error)
|
||||||
|
Query(ctx context.Context, dest interface{}, query string, args ...interface{}) error
|
||||||
|
|
||||||
|
// Transaction support
|
||||||
|
BeginTx(ctx context.Context) (Database, error)
|
||||||
|
CommitTx(ctx context.Context) error
|
||||||
|
RollbackTx(ctx context.Context) error
|
||||||
|
RunInTransaction(ctx context.Context, fn func(Database) error) error
|
||||||
|
|
||||||
|
// GetUnderlyingDB returns the underlying database connection
|
||||||
|
// For GORM, this returns *gorm.DB
|
||||||
|
// For Bun, this returns *bun.DB
|
||||||
|
// This is useful for provider-specific features like PostgreSQL NOTIFY/LISTEN
|
||||||
|
GetUnderlyingDB() interface{}
|
||||||
|
|
||||||
|
// DriverName returns the canonical name of the underlying database driver.
|
||||||
|
// Possible values: "postgres", "sqlite", "mssql", "mysql".
|
||||||
|
// All adapters normalise vendor-specific strings (e.g. Bun's "pg", GORM's
|
||||||
|
// "sqlserver") to the values above before returning.
|
||||||
|
DriverName() string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelectQuery interface for building SELECT queries (compatible with both GORM and Bun)
|
||||||
|
type SelectQuery interface {
|
||||||
|
Model(model interface{}) SelectQuery
|
||||||
|
Table(table string) SelectQuery
|
||||||
|
Column(columns ...string) SelectQuery
|
||||||
|
ColumnExpr(query string, args ...interface{}) SelectQuery
|
||||||
|
Where(query string, args ...interface{}) SelectQuery
|
||||||
|
WhereOr(query string, args ...interface{}) SelectQuery
|
||||||
|
Join(query string, args ...interface{}) SelectQuery
|
||||||
|
LeftJoin(query string, args ...interface{}) SelectQuery
|
||||||
|
Preload(relation string, conditions ...interface{}) SelectQuery
|
||||||
|
PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery
|
||||||
|
JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery
|
||||||
|
Order(order string) SelectQuery
|
||||||
|
OrderExpr(order string, args ...interface{}) SelectQuery
|
||||||
|
Limit(n int) SelectQuery
|
||||||
|
Offset(n int) SelectQuery
|
||||||
|
Group(group string) SelectQuery
|
||||||
|
Having(having string, args ...interface{}) SelectQuery
|
||||||
|
|
||||||
|
// Execution methods
|
||||||
|
Scan(ctx context.Context, dest interface{}) error
|
||||||
|
ScanModel(ctx context.Context) error
|
||||||
|
Count(ctx context.Context) (int, error)
|
||||||
|
Exists(ctx context.Context) (bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// InsertQuery interface for building INSERT queries
|
||||||
|
type InsertQuery interface {
|
||||||
|
Model(model interface{}) InsertQuery
|
||||||
|
Table(table string) InsertQuery
|
||||||
|
Value(column string, value interface{}) InsertQuery
|
||||||
|
OnConflict(action string) InsertQuery
|
||||||
|
Returning(columns ...string) InsertQuery
|
||||||
|
|
||||||
|
// Execution
|
||||||
|
Exec(ctx context.Context) (Result, error)
|
||||||
|
Scan(ctx context.Context, dest interface{}) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateQuery interface for building UPDATE queries
|
||||||
|
type UpdateQuery interface {
|
||||||
|
Model(model interface{}) UpdateQuery
|
||||||
|
Table(table string) UpdateQuery
|
||||||
|
Set(column string, value interface{}) UpdateQuery
|
||||||
|
SetMap(values map[string]interface{}) UpdateQuery
|
||||||
|
Where(query string, args ...interface{}) UpdateQuery
|
||||||
|
Returning(columns ...string) UpdateQuery
|
||||||
|
|
||||||
|
// Execution
|
||||||
|
Exec(ctx context.Context) (Result, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteQuery interface for building DELETE queries
|
||||||
|
type DeleteQuery interface {
|
||||||
|
Model(model interface{}) DeleteQuery
|
||||||
|
Table(table string) DeleteQuery
|
||||||
|
Where(query string, args ...interface{}) DeleteQuery
|
||||||
|
|
||||||
|
// Execution
|
||||||
|
Exec(ctx context.Context) (Result, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Result interface for query execution results
|
||||||
|
type Result interface {
|
||||||
|
RowsAffected() int64
|
||||||
|
LastInsertId() (int64, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelRegistry manages model registration and retrieval
|
||||||
|
type ModelRegistry interface {
|
||||||
|
RegisterModel(name string, model interface{}) error
|
||||||
|
GetModel(name string) (interface{}, error)
|
||||||
|
GetAllModels() map[string]interface{}
|
||||||
|
GetModelByEntity(schema, entity string) (interface{}, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Router interface for HTTP router abstraction
|
||||||
|
type Router interface {
|
||||||
|
HandleFunc(pattern string, handler HTTPHandlerFunc) RouteRegistration
|
||||||
|
ServeHTTP(w ResponseWriter, r Request)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RouteRegistration allows method chaining for route configuration
|
||||||
|
type RouteRegistration interface {
|
||||||
|
Methods(methods ...string) RouteRegistration
|
||||||
|
PathPrefix(prefix string) RouteRegistration
|
||||||
|
}
|
||||||
|
|
||||||
|
// Request interface abstracts HTTP request
|
||||||
|
type Request interface {
|
||||||
|
Method() string
|
||||||
|
URL() string
|
||||||
|
Header(key string) string
|
||||||
|
AllHeaders() map[string]string // Get all headers as a map
|
||||||
|
Body() ([]byte, error)
|
||||||
|
PathParam(key string) string
|
||||||
|
QueryParam(key string) string
|
||||||
|
AllQueryParams() map[string]string // Get all query parameters as a map
|
||||||
|
UnderlyingRequest() *http.Request // Get the underlying *http.Request for forwarding to other handlers
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResponseWriter interface abstracts HTTP response
|
||||||
|
type ResponseWriter interface {
|
||||||
|
SetHeader(key, value string)
|
||||||
|
WriteHeader(statusCode int)
|
||||||
|
Write(data []byte) (int, error)
|
||||||
|
WriteJSON(data interface{}) error
|
||||||
|
UnderlyingResponseWriter() http.ResponseWriter // Get the underlying http.ResponseWriter for forwarding to other handlers
|
||||||
|
}
|
||||||
|
|
||||||
|
// HTTPHandlerFunc type for HTTP handlers
|
||||||
|
type HTTPHandlerFunc func(ResponseWriter, Request)
|
||||||
|
|
||||||
|
// WrapHTTPRequest wraps standard http.ResponseWriter and *http.Request into common interfaces
|
||||||
|
func WrapHTTPRequest(w http.ResponseWriter, r *http.Request) (ResponseWriter, Request) {
|
||||||
|
return &StandardResponseWriter{w: w}, &StandardRequest{r: r}
|
||||||
|
}
|
||||||
|
|
||||||
|
// StandardResponseWriter adapts http.ResponseWriter to ResponseWriter interface
|
||||||
|
type StandardResponseWriter struct {
|
||||||
|
w http.ResponseWriter
|
||||||
|
status int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardResponseWriter) SetHeader(key, value string) {
|
||||||
|
s.w.Header().Set(key, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardResponseWriter) WriteHeader(statusCode int) {
|
||||||
|
s.status = statusCode
|
||||||
|
s.w.WriteHeader(statusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardResponseWriter) Write(data []byte) (int, error) {
|
||||||
|
return s.w.Write(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardResponseWriter) WriteJSON(data interface{}) error {
|
||||||
|
s.SetHeader("Content-Type", "application/json")
|
||||||
|
enc := json.NewEncoder(s.w)
|
||||||
|
enc.SetEscapeHTML(false)
|
||||||
|
return enc.Encode(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardResponseWriter) UnderlyingResponseWriter() http.ResponseWriter {
|
||||||
|
return s.w
|
||||||
|
}
|
||||||
|
|
||||||
|
// StandardRequest adapts *http.Request to Request interface
|
||||||
|
type StandardRequest struct {
|
||||||
|
r *http.Request
|
||||||
|
body []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) Method() string {
|
||||||
|
return s.r.Method
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) URL() string {
|
||||||
|
return s.r.URL.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) Header(key string) string {
|
||||||
|
return s.r.Header.Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) AllHeaders() map[string]string {
|
||||||
|
headers := make(map[string]string)
|
||||||
|
for key, values := range s.r.Header {
|
||||||
|
if len(values) > 0 {
|
||||||
|
headers[key] = values[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return headers
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) Body() ([]byte, error) {
|
||||||
|
if s.body != nil {
|
||||||
|
return s.body, nil
|
||||||
|
}
|
||||||
|
if s.r.Body == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
defer s.r.Body.Close()
|
||||||
|
body, err := io.ReadAll(s.r.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
s.body = body
|
||||||
|
return body, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) PathParam(key string) string {
|
||||||
|
// Standard http.Request doesn't have path params
|
||||||
|
// This should be set by the router
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) QueryParam(key string) string {
|
||||||
|
return s.r.URL.Query().Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) AllQueryParams() map[string]string {
|
||||||
|
params := make(map[string]string)
|
||||||
|
for key, values := range s.r.URL.Query() {
|
||||||
|
if len(values) > 0 {
|
||||||
|
params[key] = values[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return params
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandardRequest) UnderlyingRequest() *http.Request {
|
||||||
|
return s.r
|
||||||
|
}
|
||||||
|
|
||||||
|
// TableNameProvider interface for models that provide table names
|
||||||
|
type TableNameProvider interface {
|
||||||
|
TableName() string
|
||||||
|
}
|
||||||
|
|
||||||
|
type TableAliasProvider interface {
|
||||||
|
TableAlias() string
|
||||||
|
}
|
||||||
|
|
||||||
|
// PrimaryKeyNameProvider interface for models that provide primary key column names
|
||||||
|
type PrimaryKeyNameProvider interface {
|
||||||
|
GetIDName() string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SchemaProvider interface for models that provide schema names
|
||||||
|
type SchemaProvider interface {
|
||||||
|
SchemaName() string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SpecHandler interface represents common functionality across all spec handlers
|
||||||
|
// This is the base interface implemented by:
|
||||||
|
// - resolvespec.Handler: Handles CRUD operations via request body with explicit operation field
|
||||||
|
// - restheadspec.Handler: Handles CRUD operations via HTTP methods (GET/POST/PUT/DELETE)
|
||||||
|
// - funcspec.Handler: Handles custom SQL query execution with dynamic parameters
|
||||||
|
//
|
||||||
|
// The interface hierarchy is:
|
||||||
|
//
|
||||||
|
// SpecHandler (base)
|
||||||
|
// ├── CRUDHandler (resolvespec, restheadspec)
|
||||||
|
// └── QueryHandler (funcspec)
|
||||||
|
type SpecHandler interface {
|
||||||
|
// GetDatabase returns the underlying database connection
|
||||||
|
GetDatabase() Database
|
||||||
|
}
|
||||||
|
|
||||||
|
// CRUDHandler interface for handlers that support CRUD operations
|
||||||
|
// This is implemented by resolvespec.Handler and restheadspec.Handler
|
||||||
|
type CRUDHandler interface {
|
||||||
|
SpecHandler
|
||||||
|
|
||||||
|
// Handle processes API requests through router-agnostic interface
|
||||||
|
Handle(w ResponseWriter, r Request, params map[string]string)
|
||||||
|
|
||||||
|
// HandleGet processes GET requests for metadata
|
||||||
|
HandleGet(w ResponseWriter, r Request, params map[string]string)
|
||||||
|
}
|
||||||
|
|
||||||
|
// QueryHandler interface for handlers that execute SQL queries
|
||||||
|
// This is implemented by funcspec.Handler
|
||||||
|
// Note: funcspec uses standard http.ResponseWriter and *http.Request instead of common interfaces
|
||||||
|
type QueryHandler interface {
|
||||||
|
SpecHandler
|
||||||
|
// Methods are defined in funcspec package due to different function signature requirements
|
||||||
|
}
|
||||||
+645
@@ -0,0 +1,645 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CRUDRequestProvider interface for models that provide CRUD request strings
|
||||||
|
type CRUDRequestProvider interface {
|
||||||
|
GetRequest() string
|
||||||
|
}
|
||||||
|
|
||||||
|
// RelationshipInfoProvider interface for handlers that can provide relationship info
|
||||||
|
type RelationshipInfoProvider interface {
|
||||||
|
GetRelationshipInfo(modelType reflect.Type, relationName string) *RelationshipInfo
|
||||||
|
}
|
||||||
|
|
||||||
|
// NestedCUDProcessor handles recursive processing of nested object graphs
|
||||||
|
type NestedCUDProcessor struct {
|
||||||
|
db Database
|
||||||
|
registry ModelRegistry
|
||||||
|
relationshipHelper RelationshipInfoProvider
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewNestedCUDProcessor creates a new nested CUD processor
|
||||||
|
func NewNestedCUDProcessor(db Database, registry ModelRegistry, relationshipHelper RelationshipInfoProvider) *NestedCUDProcessor {
|
||||||
|
return &NestedCUDProcessor{
|
||||||
|
db: db,
|
||||||
|
registry: registry,
|
||||||
|
relationshipHelper: relationshipHelper,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessResult contains the result of processing a CUD operation
|
||||||
|
type ProcessResult struct {
|
||||||
|
ID interface{} // The ID of the processed record
|
||||||
|
AffectedRows int64 // Number of rows affected
|
||||||
|
Data map[string]interface{} // The processed data
|
||||||
|
RelationData map[string]interface{} // Data from processed relations
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessNestedCUD recursively processes nested object graphs for Create, Update, Delete operations
|
||||||
|
// with automatic foreign key resolution
|
||||||
|
func (p *NestedCUDProcessor) ProcessNestedCUD(
|
||||||
|
ctx context.Context,
|
||||||
|
operation string, // "insert", "update", or "delete"
|
||||||
|
data map[string]interface{},
|
||||||
|
model interface{},
|
||||||
|
parentIDs map[string]interface{}, // Parent IDs for foreign key resolution
|
||||||
|
tableName string,
|
||||||
|
) (*ProcessResult, error) {
|
||||||
|
logger.Info("Processing nested CUD: operation=%s, table=%s", operation, tableName)
|
||||||
|
|
||||||
|
result := &ProcessResult{
|
||||||
|
Data: make(map[string]interface{}),
|
||||||
|
RelationData: make(map[string]interface{}),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if data has a _request field that overrides the operation
|
||||||
|
if requestOp := p.extractCRUDRequest(data); requestOp != "" {
|
||||||
|
logger.Debug("Found _request override: %s", requestOp)
|
||||||
|
operation = requestOp
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get model type for reflection
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
logger.Error("Invalid model type: operation=%s, table=%s, modelType=%v, expected struct", operation, tableName, modelType)
|
||||||
|
return nil, fmt.Errorf("model must be a struct type, got %v", modelType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Separate relation fields from regular fields
|
||||||
|
relationFields := make(map[string]*RelationshipInfo)
|
||||||
|
regularData := make(map[string]interface{})
|
||||||
|
|
||||||
|
for key, value := range data {
|
||||||
|
// Skip _request field in actual data processing
|
||||||
|
if key == "_request" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this field is a relation
|
||||||
|
relInfo := p.relationshipHelper.GetRelationshipInfo(modelType, key)
|
||||||
|
if relInfo != nil {
|
||||||
|
relationFields[key] = relInfo
|
||||||
|
result.RelationData[key] = value
|
||||||
|
} else {
|
||||||
|
regularData[key] = value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Filter regularData to only include fields that exist in the model,
|
||||||
|
// and translate JSON keys to their actual database column names.
|
||||||
|
regularData = p.filterValidFields(regularData, model)
|
||||||
|
|
||||||
|
// Inject parent IDs for foreign key resolution
|
||||||
|
p.injectForeignKeys(regularData, modelType, parentIDs)
|
||||||
|
|
||||||
|
// Get the primary key name for this model
|
||||||
|
pkName := reflection.GetPrimaryKeyName(model)
|
||||||
|
|
||||||
|
// Check if we have any data to process (besides _request)
|
||||||
|
hasData := len(regularData) > 0
|
||||||
|
|
||||||
|
// Process based on operation
|
||||||
|
switch strings.ToLower(operation) {
|
||||||
|
case "insert", "create", "add":
|
||||||
|
// Only perform insert if we have data to insert
|
||||||
|
if hasData {
|
||||||
|
id, err := p.processInsert(ctx, regularData, tableName)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
|
||||||
|
return nil, fmt.Errorf("insert failed: %w", err)
|
||||||
|
}
|
||||||
|
result.ID = id
|
||||||
|
result.AffectedRows = 1
|
||||||
|
result.Data = regularData
|
||||||
|
|
||||||
|
// Re-select the inserted row so result.Data reflects DB-generated defaults.
|
||||||
|
if row, err := p.processSelect(ctx, tableName, id); err != nil {
|
||||||
|
logger.Warn("Select after insert failed: table=%s, id=%v, error=%v", tableName, id, err)
|
||||||
|
} else if len(row) > 0 {
|
||||||
|
result.Data = row
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process child relations after parent insert (to get parent ID)
|
||||||
|
if err := p.processChildRelations(ctx, "insert", id, relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
||||||
|
logger.Error("Failed to process child relations after insert: table=%s, parentID=%v, relations=%+v, error=%v", tableName, id, relationFields, err)
|
||||||
|
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.Debug("Skipping insert for %s - no data columns besides _request", tableName)
|
||||||
|
}
|
||||||
|
|
||||||
|
case "update", "change", "modify":
|
||||||
|
// Only perform update if we have data to update
|
||||||
|
if reflection.IsEmptyValue(data[pkName]) {
|
||||||
|
logger.Warn("Skipping update for %s - no primary key", tableName)
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
if hasData {
|
||||||
|
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err)
|
||||||
|
return nil, fmt.Errorf("update failed: %w", err)
|
||||||
|
}
|
||||||
|
result.ID = data[pkName]
|
||||||
|
result.AffectedRows = rows
|
||||||
|
result.Data = regularData
|
||||||
|
|
||||||
|
// Re-select the updated row so result.Data reflects current DB state.
|
||||||
|
if row, err := p.processSelect(ctx, tableName, result.ID); err != nil {
|
||||||
|
logger.Warn("Select after update failed: table=%s, id=%v, error=%v", tableName, result.ID, err)
|
||||||
|
} else if len(row) > 0 {
|
||||||
|
result.Data = row
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process child relations for update
|
||||||
|
if err := p.processChildRelations(ctx, "update", data[pkName], relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
||||||
|
logger.Error("Failed to process child relations after update: table=%s, parentID=%v, relations=%+v, error=%v", tableName, data[pkName], regularData, err)
|
||||||
|
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.Debug("Skipping update for %s - no data columns besides _request", tableName)
|
||||||
|
result.ID = data[pkName]
|
||||||
|
}
|
||||||
|
|
||||||
|
case "delete", "remove":
|
||||||
|
if reflection.IsEmptyValue(data[pkName]) {
|
||||||
|
logger.Warn("Skipping delete for %s - no primary key", tableName)
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process child relations first (for referential integrity)
|
||||||
|
if err := p.processChildRelations(ctx, "delete", data[pkName], relationFields, result.RelationData, modelType, parentIDs); err != nil {
|
||||||
|
logger.Error("Failed to process child relations before delete: table=%s, id=%v, relations=%+v, error=%v", tableName, data[pkName], relationFields, err)
|
||||||
|
return nil, fmt.Errorf("failed to process child relations: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := p.processDelete(ctx, tableName, data[pkName])
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Delete failed for table=%s, id=%v, error=%v", tableName, data[pkName], err)
|
||||||
|
return nil, fmt.Errorf("delete failed: %w", err)
|
||||||
|
}
|
||||||
|
result.ID = data[pkName]
|
||||||
|
result.AffectedRows = rows
|
||||||
|
result.Data = regularData
|
||||||
|
|
||||||
|
default:
|
||||||
|
logger.Error("Unsupported operation: %s for table=%s", operation, tableName)
|
||||||
|
return nil, fmt.Errorf("unsupported operation: %s", operation)
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Info("Nested CUD completed: operation=%s, id=%v, rows=%d", operation, result.ID, result.AffectedRows)
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractCRUDRequest extracts the request field from data if present
|
||||||
|
func (p *NestedCUDProcessor) extractCRUDRequest(data map[string]interface{}) string {
|
||||||
|
if request, ok := data["_request"]; ok {
|
||||||
|
if requestStr, ok := request.(string); ok {
|
||||||
|
return strings.ToLower(strings.TrimSpace(requestStr))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterValidFields filters input data to only include fields that exist in the model,
|
||||||
|
// and translates JSON key names to their actual database column names.
|
||||||
|
// For example, a field tagged `json:"_changed_date" bun:"changed_date"` will be
|
||||||
|
// included in the result as "changed_date", not "_changed_date".
|
||||||
|
func (p *NestedCUDProcessor) filterValidFields(data map[string]interface{}, model interface{}) map[string]interface{} {
|
||||||
|
if len(data) == 0 {
|
||||||
|
return data
|
||||||
|
}
|
||||||
|
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
return data
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build a mapping from JSON key -> DB column name for all writable fields.
|
||||||
|
// This both validates which fields belong to the model and translates their names
|
||||||
|
// to the correct column names for use in SQL insert/update queries.
|
||||||
|
jsonToDBCol := reflection.BuildJSONToDBColumnMap(modelType)
|
||||||
|
|
||||||
|
filteredData := make(map[string]interface{})
|
||||||
|
for key, value := range data {
|
||||||
|
dbColName, exists := jsonToDBCol[key]
|
||||||
|
if exists {
|
||||||
|
filteredData[dbColName] = value
|
||||||
|
} else {
|
||||||
|
logger.Debug("Skipping invalid field '%s' - not found in model %v", key, modelType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return filteredData
|
||||||
|
}
|
||||||
|
|
||||||
|
// injectForeignKeys injects parent IDs into data for foreign key fields.
|
||||||
|
// data is expected to be keyed by DB column names (as returned by filterValidFields).
|
||||||
|
func (p *NestedCUDProcessor) injectForeignKeys(data map[string]interface{}, modelType reflect.Type, parentIDs map[string]interface{}) {
|
||||||
|
if len(parentIDs) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pkCol := reflection.GetPrimaryKeyName(reflect.New(modelType).Interface())
|
||||||
|
|
||||||
|
for parentKey, parentID := range parentIDs {
|
||||||
|
dbColNames := reflection.GetForeignKeyColumn(modelType, parentKey)
|
||||||
|
|
||||||
|
if len(dbColNames) == 0 {
|
||||||
|
// No explicit tag found — fall back to naming convention by scanning scalar fields.
|
||||||
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
|
field := modelType.Field(i)
|
||||||
|
jsonName := strings.Split(field.Tag.Get("json"), ",")[0]
|
||||||
|
if strings.EqualFold(jsonName, "rid"+parentKey) ||
|
||||||
|
strings.EqualFold(jsonName, "rid_"+parentKey) ||
|
||||||
|
strings.EqualFold(jsonName, "id_"+parentKey) ||
|
||||||
|
strings.EqualFold(jsonName, parentKey+"_id") ||
|
||||||
|
strings.EqualFold(jsonName, parentKey+"id") ||
|
||||||
|
strings.EqualFold(field.Name, parentKey+"ID") {
|
||||||
|
dbColNames = []string{reflection.GetColumnName(field)}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, dbColName := range dbColNames {
|
||||||
|
if pkCol != "" && strings.EqualFold(dbColName, pkCol) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, exists := data[dbColName]; !exists {
|
||||||
|
logger.Debug("Injecting foreign key: %s = %v", dbColName, parentID)
|
||||||
|
data[dbColName] = parentID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// processInsert handles insert operation
|
||||||
|
func (p *NestedCUDProcessor) processInsert(
|
||||||
|
ctx context.Context,
|
||||||
|
data map[string]interface{},
|
||||||
|
tableName string,
|
||||||
|
) (interface{}, error) {
|
||||||
|
logger.Debug("Inserting into %s with data: %+v", tableName, data)
|
||||||
|
|
||||||
|
query := p.db.NewInsert().Table(tableName)
|
||||||
|
|
||||||
|
for key, value := range data {
|
||||||
|
query = query.Value(key, ConvertSliceForBun(value))
|
||||||
|
}
|
||||||
|
pkName := reflection.GetPrimaryKeyName(tableName)
|
||||||
|
query = query.Returning(pkName)
|
||||||
|
|
||||||
|
var id interface{}
|
||||||
|
if err := query.Scan(ctx, &id); err != nil {
|
||||||
|
logger.Error("Insert execution failed: table=%s, data=%+v, error=%v", tableName, data, err)
|
||||||
|
return nil, fmt.Errorf("insert exec failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Insert successful, ID: %v", id)
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// processSelect fetches the row identified by id from tableName into a flat map.
|
||||||
|
// Used to populate result.Data with the actual DB state after insert/update.
|
||||||
|
func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string, id interface{}) (map[string]interface{}, error) {
|
||||||
|
pkName := reflection.GetPrimaryKeyName(tableName)
|
||||||
|
var row map[string]interface{}
|
||||||
|
if err := p.db.NewSelect().
|
||||||
|
Table(tableName).
|
||||||
|
Where(fmt.Sprintf("%s = ?", QuoteIdent(pkName)), id).
|
||||||
|
Scan(ctx, &row); err != nil {
|
||||||
|
return nil, fmt.Errorf("select after write failed: %w", err)
|
||||||
|
}
|
||||||
|
return row, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// processUpdate handles update operation
|
||||||
|
func (p *NestedCUDProcessor) processUpdate(
|
||||||
|
ctx context.Context,
|
||||||
|
data map[string]interface{},
|
||||||
|
tableName string,
|
||||||
|
id interface{},
|
||||||
|
) (int64, error) {
|
||||||
|
if id == nil {
|
||||||
|
logger.Error("Update requires an ID: table=%s, data=%+v", tableName, data)
|
||||||
|
return 0, fmt.Errorf("update requires an ID")
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Updating %s with ID %v, data: %+v", tableName, id, data)
|
||||||
|
|
||||||
|
query := p.db.NewUpdate().Table(tableName).SetMap(data).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id)
|
||||||
|
|
||||||
|
result, err := query.Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Update execution failed: table=%s, id=%v, data=%+v, error=%v", tableName, id, data, err)
|
||||||
|
return 0, fmt.Errorf("update exec failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows := result.RowsAffected()
|
||||||
|
logger.Debug("Update successful, rows affected: %d", rows)
|
||||||
|
return rows, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// processDelete handles delete operation
|
||||||
|
func (p *NestedCUDProcessor) processDelete(ctx context.Context, tableName string, id interface{}) (int64, error) {
|
||||||
|
if id == nil {
|
||||||
|
logger.Error("Delete requires an ID: table=%s", tableName)
|
||||||
|
return 0, fmt.Errorf("delete requires an ID")
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Deleting from %s with ID %v", tableName, id)
|
||||||
|
|
||||||
|
query := p.db.NewDelete().Table(tableName).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id)
|
||||||
|
|
||||||
|
result, err := query.Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Delete execution failed: table=%s, id=%v, error=%v", tableName, id, err)
|
||||||
|
return 0, fmt.Errorf("delete exec failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows := result.RowsAffected()
|
||||||
|
logger.Debug("Delete successful, rows affected: %d", rows)
|
||||||
|
return rows, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// processChildRelations recursively processes child relations
|
||||||
|
func (p *NestedCUDProcessor) processChildRelations(
|
||||||
|
ctx context.Context,
|
||||||
|
operation string,
|
||||||
|
parentID interface{},
|
||||||
|
relationFields map[string]*RelationshipInfo,
|
||||||
|
relationData map[string]interface{},
|
||||||
|
parentModelType reflect.Type,
|
||||||
|
incomingParentIDs map[string]interface{}, // IDs from all ancestors
|
||||||
|
) error {
|
||||||
|
for relationName, relInfo := range relationFields {
|
||||||
|
relationValue, exists := relationData[relationName]
|
||||||
|
if !exists || relationValue == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Processing relation: %s, type: %s", relationName, relInfo.RelationType)
|
||||||
|
|
||||||
|
// Get the related model
|
||||||
|
field, found := parentModelType.FieldByName(relInfo.FieldName)
|
||||||
|
if !found {
|
||||||
|
logger.Error("Field %s not found in model type %v for relation %s", relInfo.FieldName, parentModelType, relationName)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get the model type for the relation
|
||||||
|
relatedModelType := field.Type
|
||||||
|
if relatedModelType.Kind() == reflect.Slice {
|
||||||
|
relatedModelType = relatedModelType.Elem()
|
||||||
|
}
|
||||||
|
if relatedModelType.Kind() == reflect.Pointer {
|
||||||
|
relatedModelType = relatedModelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create an instance of the related model
|
||||||
|
relatedModel := reflect.New(relatedModelType).Elem().Interface()
|
||||||
|
|
||||||
|
// Get table name for related model
|
||||||
|
relatedTableName := p.getTableNameForModel(relatedModel, relInfo.JSONName)
|
||||||
|
|
||||||
|
// Prepare parent IDs for foreign key injection
|
||||||
|
// Start by copying all incoming parent IDs (from ancestors)
|
||||||
|
parentIDs := make(map[string]interface{})
|
||||||
|
for k, v := range incomingParentIDs {
|
||||||
|
parentIDs[k] = v
|
||||||
|
}
|
||||||
|
logger.Debug("Inherited %d parent IDs from ancestors: %+v", len(incomingParentIDs), incomingParentIDs)
|
||||||
|
|
||||||
|
// Add the current parent's primary key to the parentIDs map
|
||||||
|
// This ensures nested children have access to all ancestor IDs
|
||||||
|
if parentID != nil && parentModelType != nil {
|
||||||
|
// Get the parent model's primary key field name
|
||||||
|
parentPKFieldName := reflection.GetPrimaryKeyName(parentModelType)
|
||||||
|
if parentPKFieldName != "" {
|
||||||
|
// Get the JSON name for the primary key field
|
||||||
|
parentPKJSONName := reflection.GetJSONNameForField(parentModelType, parentPKFieldName)
|
||||||
|
baseName := ""
|
||||||
|
if len(parentPKJSONName) > 1 {
|
||||||
|
baseName = parentPKJSONName
|
||||||
|
} else {
|
||||||
|
// Add parent's PK to the map using the base model name
|
||||||
|
baseName = strings.TrimSuffix(parentPKFieldName, "ID")
|
||||||
|
baseName = strings.TrimSuffix(strings.ToLower(baseName), "_id")
|
||||||
|
if baseName == "" {
|
||||||
|
baseName = "parent"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
parentIDs[baseName] = parentID
|
||||||
|
logger.Debug("Added current parent PK to parentIDs map: %s=%v (from field %s)", baseName, parentID, parentPKFieldName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Also add the foreign key reference if specified
|
||||||
|
if relInfo.ForeignKey != "" && parentID != nil {
|
||||||
|
// Extract the base name from foreign key (e.g., "DepartmentID" -> "Department")
|
||||||
|
baseName := strings.TrimSuffix(relInfo.ForeignKey, "ID")
|
||||||
|
baseName = strings.TrimSuffix(strings.ToLower(baseName), "_id")
|
||||||
|
// Only add if different from what we already added
|
||||||
|
if _, exists := parentIDs[baseName]; !exists {
|
||||||
|
parentIDs[baseName] = parentID
|
||||||
|
logger.Debug("Added foreign key to parentIDs map: %s=%v (from FK %s)", baseName, parentID, relInfo.ForeignKey)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Final parentIDs map for relation %s: %+v", relationName, parentIDs)
|
||||||
|
|
||||||
|
// Determine which field name to use for setting parent ID in child data
|
||||||
|
// Priority: Use foreign key field name if specified
|
||||||
|
var foreignKeyFieldName string
|
||||||
|
if relInfo.ForeignKey != "" {
|
||||||
|
// For has-many/has-one: join:parentCol=childCol
|
||||||
|
// ForeignKey = parent side, References = child side (where we actually set the value)
|
||||||
|
childField := relInfo.ForeignKey
|
||||||
|
if (relInfo.RelationType == "hasMany" || relInfo.RelationType == "hasOne") && relInfo.References != "" {
|
||||||
|
childField = relInfo.References
|
||||||
|
}
|
||||||
|
foreignKeyFieldName = reflection.GetJSONNameForField(relatedModelType, childField)
|
||||||
|
if foreignKeyFieldName == "" {
|
||||||
|
foreignKeyFieldName = strings.ToLower(childField)
|
||||||
|
}
|
||||||
|
logger.Debug("Using foreign key field for direct assignment: %s (from FK %s -> child %s)", foreignKeyFieldName, relInfo.ForeignKey, childField)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get the primary key name for the child model to avoid overwriting it in recursive relationships
|
||||||
|
childPKName := reflection.GetPrimaryKeyName(relatedModel)
|
||||||
|
childPKFieldName := reflection.GetJSONNameForField(relatedModelType, childPKName)
|
||||||
|
if childPKFieldName == "" {
|
||||||
|
childPKFieldName = strings.ToLower(childPKName)
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Processing relation with foreignKeyField=%s, childPK=%s", foreignKeyFieldName, childPKFieldName)
|
||||||
|
|
||||||
|
// Process based on relation type and data structure
|
||||||
|
switch v := relationValue.(type) {
|
||||||
|
case map[string]interface{}:
|
||||||
|
// Single related object - directly set foreign key if specified
|
||||||
|
// IMPORTANT: In recursive relationships, don't overwrite the primary key
|
||||||
|
if parentID != nil && foreignKeyFieldName != "" && foreignKeyFieldName != childPKFieldName {
|
||||||
|
v[foreignKeyFieldName] = parentID
|
||||||
|
logger.Debug("Set foreign key in single relation: %s=%v", foreignKeyFieldName, parentID)
|
||||||
|
} else if foreignKeyFieldName == childPKFieldName {
|
||||||
|
logger.Debug("Skipping foreign key assignment - same as primary key (recursive relationship): %s", foreignKeyFieldName)
|
||||||
|
}
|
||||||
|
_, err := p.ProcessNestedCUD(ctx, operation, v, relatedModel, parentIDs, relatedTableName)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Failed to process single relation: name=%s, table=%s, operation=%s, parentID=%v, data=%+v, error=%v",
|
||||||
|
relationName, relatedTableName, operation, parentID, v, err)
|
||||||
|
return fmt.Errorf("failed to process relation %s: %w", relationName, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
case []interface{}:
|
||||||
|
// Multiple related objects
|
||||||
|
for i, item := range v {
|
||||||
|
if itemMap, ok := item.(map[string]interface{}); ok {
|
||||||
|
// Directly set foreign key if specified
|
||||||
|
// IMPORTANT: In recursive relationships, don't overwrite the primary key
|
||||||
|
if parentID != nil && foreignKeyFieldName != "" && foreignKeyFieldName != childPKFieldName {
|
||||||
|
itemMap[foreignKeyFieldName] = parentID
|
||||||
|
logger.Debug("Set foreign key in relation array[%d]: %s=%v", i, foreignKeyFieldName, parentID)
|
||||||
|
} else if foreignKeyFieldName == childPKFieldName {
|
||||||
|
logger.Debug("Skipping foreign key assignment in array[%d] - same as primary key (recursive relationship): %s", i, foreignKeyFieldName)
|
||||||
|
}
|
||||||
|
_, err := p.ProcessNestedCUD(ctx, operation, itemMap, relatedModel, parentIDs, relatedTableName)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Failed to process relation array item: name=%s[%d], table=%s, operation=%s, parentID=%v, data=%+v, error=%v",
|
||||||
|
relationName, i, relatedTableName, operation, parentID, itemMap, err)
|
||||||
|
return fmt.Errorf("failed to process relation %s[%d]: %w", relationName, i, err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.Warn("Relation array item is not a map: name=%s[%d], type=%T", relationName, i, item)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
case []map[string]interface{}:
|
||||||
|
// Multiple related objects (typed slice)
|
||||||
|
for i, itemMap := range v {
|
||||||
|
// Directly set foreign key if specified
|
||||||
|
// IMPORTANT: In recursive relationships, don't overwrite the primary key
|
||||||
|
if parentID != nil && foreignKeyFieldName != "" && foreignKeyFieldName != childPKFieldName {
|
||||||
|
itemMap[foreignKeyFieldName] = parentID
|
||||||
|
logger.Debug("Set foreign key in relation typed array[%d]: %s=%v", i, foreignKeyFieldName, parentID)
|
||||||
|
} else if foreignKeyFieldName == childPKFieldName {
|
||||||
|
logger.Debug("Skipping foreign key assignment in typed array[%d] - same as primary key (recursive relationship): %s", i, foreignKeyFieldName)
|
||||||
|
}
|
||||||
|
_, err := p.ProcessNestedCUD(ctx, operation, itemMap, relatedModel, parentIDs, relatedTableName)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Failed to process relation typed array item: name=%s[%d], table=%s, operation=%s, parentID=%v, data=%+v, error=%v",
|
||||||
|
relationName, i, relatedTableName, operation, parentID, itemMap, err)
|
||||||
|
return fmt.Errorf("failed to process relation %s[%d]: %w", relationName, i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
default:
|
||||||
|
logger.Error("Unsupported relation data type: name=%s, type=%T, value=%+v", relationName, relationValue, relationValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// getTableNameForModel gets the table name for a model
|
||||||
|
func (p *NestedCUDProcessor) getTableNameForModel(model interface{}, defaultName string) string {
|
||||||
|
if provider, ok := model.(TableNameProvider); ok {
|
||||||
|
tableName := provider.TableName()
|
||||||
|
if tableName != "" {
|
||||||
|
return tableName
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return defaultName
|
||||||
|
}
|
||||||
|
|
||||||
|
// ShouldUseNestedProcessor determines if we should use nested CUD processing
|
||||||
|
// It recursively checks if the data contains:
|
||||||
|
// 1. A _request field at any level, OR
|
||||||
|
// 2. Nested relations that themselves contain further nested relations or _request fields
|
||||||
|
// This ensures nested processing is only used when there are deeply nested operations
|
||||||
|
func ShouldUseNestedProcessor(data map[string]interface{}, model interface{}, relationshipHelper RelationshipInfoProvider) bool {
|
||||||
|
return shouldUseNestedProcessorDepth(data, model, relationshipHelper, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
// shouldUseNestedProcessorDepth is the internal recursive implementation with depth tracking
|
||||||
|
func shouldUseNestedProcessorDepth(data map[string]interface{}, model interface{}, relationshipHelper RelationshipInfoProvider, depth int) bool {
|
||||||
|
// Check for _request field
|
||||||
|
if _, hasCRUDRequest := data["_request"]; hasCRUDRequest {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get model type
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if data contains any fields that are relations (nested objects or arrays)
|
||||||
|
for key, value := range data {
|
||||||
|
// Skip _request and regular scalar fields
|
||||||
|
if key == "_request" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this field is a relation in the model
|
||||||
|
relInfo := relationshipHelper.GetRelationshipInfo(modelType, key)
|
||||||
|
if relInfo != nil {
|
||||||
|
// Check if the value is actually nested data (object or array)
|
||||||
|
switch v := value.(type) {
|
||||||
|
case map[string]interface{}, []interface{}, []map[string]interface{}:
|
||||||
|
// If we're already at a nested level (depth > 0) and found a relation,
|
||||||
|
// that means we have multi-level nesting, so return true
|
||||||
|
if depth > 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// At depth 0, recurse to check if the nested data has further nesting
|
||||||
|
switch typedValue := v.(type) {
|
||||||
|
case map[string]interface{}:
|
||||||
|
if shouldUseNestedProcessorDepth(typedValue, relInfo.RelatedModel, relationshipHelper, depth+1) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
case []interface{}:
|
||||||
|
for _, item := range typedValue {
|
||||||
|
if itemMap, ok := item.(map[string]interface{}); ok {
|
||||||
|
if shouldUseNestedProcessorDepth(itemMap, relInfo.RelatedModel, relationshipHelper, depth+1) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case []map[string]interface{}:
|
||||||
|
for _, itemMap := range typedValue {
|
||||||
|
if shouldUseNestedProcessorDepth(itemMap, relInfo.RelatedModel, relationshipHelper, depth+1) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
+1042
File diff suppressed because it is too large
Load Diff
+153
@@ -0,0 +1,153 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
// SQLError wraps a database error together with the SQL that caused it,
|
||||||
|
// so callers can surface the query in API error responses for easier debugging.
|
||||||
|
type SQLError struct {
|
||||||
|
Err error
|
||||||
|
SQL string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *SQLError) Error() string { return e.Err.Error() }
|
||||||
|
func (e *SQLError) Unwrap() error { return e.Err }
|
||||||
|
|
||||||
|
// WrapSQLError wraps err with the given SQL. If err is nil it returns nil.
|
||||||
|
func WrapSQLError(err error, sql string) error {
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &SQLError{Err: err, SQL: sql}
|
||||||
|
}
|
||||||
|
|
||||||
|
type RequestBody struct {
|
||||||
|
Operation string `json:"operation"`
|
||||||
|
Data interface{} `json:"data"`
|
||||||
|
ID *int64 `json:"id"`
|
||||||
|
Options RequestOptions `json:"options"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type RequestOptions struct {
|
||||||
|
Preload []PreloadOption `json:"preload"`
|
||||||
|
Columns []string `json:"columns"`
|
||||||
|
OmitColumns []string `json:"omit_columns"`
|
||||||
|
Filters []FilterOption `json:"filters"`
|
||||||
|
Sort []SortOption `json:"sort"`
|
||||||
|
Limit *int `json:"limit"`
|
||||||
|
Offset *int `json:"offset"`
|
||||||
|
CustomOperators []CustomOperator `json:"customOperators"`
|
||||||
|
ComputedColumns []ComputedColumn `json:"computedColumns"`
|
||||||
|
Parameters []Parameter `json:"parameters"`
|
||||||
|
|
||||||
|
// Cursor pagination
|
||||||
|
CursorForward string `json:"cursor_forward"`
|
||||||
|
CursorBackward string `json:"cursor_backward"`
|
||||||
|
FetchRowNumber *string `json:"fetch_row_number"`
|
||||||
|
|
||||||
|
// Join table aliases (used for validation of prefixed columns in filters/sorts)
|
||||||
|
// Not serialized to JSON as it's internal validation state
|
||||||
|
JoinAliases []string `json:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type Parameter struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Value string `json:"value"`
|
||||||
|
Sequence *int `json:"sequence"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type PreloadOption struct {
|
||||||
|
Relation string `json:"relation"`
|
||||||
|
TableName string `json:"table_name"` // Actual database table name (e.g., "mastertaskitem")
|
||||||
|
Columns []string `json:"columns"`
|
||||||
|
OmitColumns []string `json:"omit_columns"`
|
||||||
|
Sort []SortOption `json:"sort"`
|
||||||
|
Filters []FilterOption `json:"filters"`
|
||||||
|
Where string `json:"where"`
|
||||||
|
Limit *int `json:"limit"`
|
||||||
|
Offset *int `json:"offset"`
|
||||||
|
Updatable *bool `json:"updateable"` // if true, the relation can be updated
|
||||||
|
ComputedQL map[string]string `json:"computed_ql"` // Computed columns as SQL expressions
|
||||||
|
Recursive bool `json:"recursive"` // if true, preload recursively up to 5 levels
|
||||||
|
|
||||||
|
// Relationship keys from XFiles - used to build proper foreign key filters
|
||||||
|
PrimaryKey string `json:"primary_key"` // Primary key of the related table
|
||||||
|
RelatedKey string `json:"related_key"` // For child tables: column in child that references parent
|
||||||
|
ForeignKey string `json:"foreign_key"` // For parent tables: column in current table that references parent
|
||||||
|
RecursiveChildKey string `json:"recursive_child_key"` // For recursive tables: FK column used for recursion (e.g., "rid_parentmastertaskitem")
|
||||||
|
|
||||||
|
// Custom SQL JOINs from XFiles - used when preload needs additional joins
|
||||||
|
SqlJoins []string `json:"sql_joins"` // Custom SQL JOIN clauses
|
||||||
|
JoinAliases []string `json:"join_aliases"` // Extracted table aliases from SqlJoins for validation
|
||||||
|
}
|
||||||
|
|
||||||
|
type FilterOption struct {
|
||||||
|
Column string `json:"column"`
|
||||||
|
Operator string `json:"operator"`
|
||||||
|
Value interface{} `json:"value"`
|
||||||
|
LogicOperator string `json:"logic_operator"` // "AND" or "OR" - how this filter combines with previous filters
|
||||||
|
}
|
||||||
|
|
||||||
|
type SortOption struct {
|
||||||
|
Column string `json:"column"`
|
||||||
|
Direction string `json:"direction"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type CustomOperator struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
SQL string `json:"sql"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ComputedColumn struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Expression string `json:"expression"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Response structures
|
||||||
|
type Response struct {
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Data interface{} `json:"data"`
|
||||||
|
Metadata *Metadata `json:"metadata,omitempty"`
|
||||||
|
Error *APIError `json:"error,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type Metadata struct {
|
||||||
|
Total int64 `json:"total"`
|
||||||
|
Count int64 `json:"count"`
|
||||||
|
Filtered int64 `json:"filtered"`
|
||||||
|
Limit int `json:"limit"`
|
||||||
|
Offset int `json:"offset"`
|
||||||
|
RowNumber *int64 `json:"row_number,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type APIError struct {
|
||||||
|
Code string `json:"code"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Details interface{} `json:"details,omitempty"`
|
||||||
|
Detail string `json:"detail,omitempty"`
|
||||||
|
SQL string `json:"sql,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type Column struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
IsNullable bool `json:"is_nullable"`
|
||||||
|
IsPrimary bool `json:"is_primary"`
|
||||||
|
IsUnique bool `json:"is_unique"`
|
||||||
|
HasIndex bool `json:"has_index"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type TableMetadata struct {
|
||||||
|
Schema string `json:"schema"`
|
||||||
|
Table string `json:"table"`
|
||||||
|
Columns []Column `json:"columns"`
|
||||||
|
Relations []string `json:"relations"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RelationshipInfo contains information about a model relationship
|
||||||
|
type RelationshipInfo struct {
|
||||||
|
FieldName string `json:"field_name"`
|
||||||
|
JSONName string `json:"json_name"`
|
||||||
|
RelationType string `json:"relation_type"` // "belongsTo", "hasMany", "hasOne", "many2many"
|
||||||
|
ForeignKey string `json:"foreign_key"`
|
||||||
|
References string `json:"references"`
|
||||||
|
JoinTable string `json:"join_table"`
|
||||||
|
RelatedModel interface{} `json:"related_model"`
|
||||||
|
}
|
||||||
+431
@@ -0,0 +1,431 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ColumnValidator validates column names against a model's fields
|
||||||
|
type ColumnValidator struct {
|
||||||
|
validColumns map[string]bool
|
||||||
|
model interface{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewColumnValidator creates a new column validator for a given model
|
||||||
|
func NewColumnValidator(model interface{}) *ColumnValidator {
|
||||||
|
validator := &ColumnValidator{
|
||||||
|
validColumns: make(map[string]bool),
|
||||||
|
model: model,
|
||||||
|
}
|
||||||
|
validator.buildValidColumns()
|
||||||
|
return validator
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildValidColumns extracts all valid column names from the model using reflection
|
||||||
|
func (v *ColumnValidator) buildValidColumns() {
|
||||||
|
modelType := reflect.TypeOf(v.model)
|
||||||
|
|
||||||
|
// Unwrap pointers, slices, and arrays to get to the base struct type
|
||||||
|
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array) {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate that we have a struct type
|
||||||
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract column names from struct fields
|
||||||
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
|
field := modelType.Field(i)
|
||||||
|
|
||||||
|
if !field.IsExported() || field.Anonymous {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get column name from bun, gorm, or json tag
|
||||||
|
columnName := v.getColumnName(field)
|
||||||
|
if columnName != "" && columnName != "-" {
|
||||||
|
v.validColumns[strings.ToLower(columnName)] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// getColumnName extracts the column name from a struct field's tags
|
||||||
|
// Supports both Bun and GORM tags
|
||||||
|
func (v *ColumnValidator) getColumnName(field reflect.StructField) string {
|
||||||
|
// First check Bun tag for column name
|
||||||
|
bunTag := field.Tag.Get("bun")
|
||||||
|
if bunTag != "" && bunTag != "-" {
|
||||||
|
parts := strings.Split(bunTag, ",")
|
||||||
|
// The first part is usually the column name
|
||||||
|
columnName := strings.TrimSpace(parts[0])
|
||||||
|
if columnName != "" && columnName != "-" {
|
||||||
|
return columnName
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check GORM tag for column name
|
||||||
|
gormTag := field.Tag.Get("gorm")
|
||||||
|
if strings.Contains(gormTag, "column:") {
|
||||||
|
parts := strings.Split(gormTag, ";")
|
||||||
|
for _, part := range parts {
|
||||||
|
part = strings.TrimSpace(part)
|
||||||
|
if strings.HasPrefix(part, "column:") {
|
||||||
|
return strings.TrimPrefix(part, "column:")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fall back to JSON tag
|
||||||
|
jsonTag := field.Tag.Get("json")
|
||||||
|
if jsonTag != "" && jsonTag != "-" {
|
||||||
|
// Extract just the name part (before any comma)
|
||||||
|
jsonName := strings.Split(jsonTag, ",")[0]
|
||||||
|
return jsonName
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fall back to field name in lowercase (snake_case conversion would be better)
|
||||||
|
return strings.ToLower(field.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateColumn validates a single column name
|
||||||
|
// Returns nil if valid, error if invalid
|
||||||
|
// Columns prefixed with "cql" (case insensitive) are always valid
|
||||||
|
// Handles PostgreSQL JSON operators (-> and ->>)
|
||||||
|
func (v *ColumnValidator) ValidateColumn(column string) error {
|
||||||
|
// Allow empty columns
|
||||||
|
if column == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allow columns prefixed with "cql" (case insensitive) for computed columns
|
||||||
|
if strings.HasPrefix(strings.ToLower(column), "cql") {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract source column name (remove JSON operators like ->> or ->)
|
||||||
|
sourceColumn := reflection.ExtractSourceColumn(column)
|
||||||
|
|
||||||
|
// Check if column exists in model
|
||||||
|
if _, exists := v.validColumns[strings.ToLower(sourceColumn)]; !exists {
|
||||||
|
return fmt.Errorf("invalid column '%s': column does not exist in model", column)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsValidColumn checks if a column is valid
|
||||||
|
// Returns true if valid, false if invalid
|
||||||
|
func (v *ColumnValidator) IsValidColumn(column string) bool {
|
||||||
|
return v.ValidateColumn(column) == nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Columns returns all valid column names known to this validator
|
||||||
|
func (v *ColumnValidator) Columns() []string {
|
||||||
|
cols := make([]string, 0, len(v.validColumns))
|
||||||
|
for col := range v.validColumns {
|
||||||
|
cols = append(cols, col)
|
||||||
|
}
|
||||||
|
sort.Strings(cols)
|
||||||
|
return cols
|
||||||
|
}
|
||||||
|
|
||||||
|
// FilterValidColumns filters a list of columns, returning only valid ones
|
||||||
|
// Logs warnings for any invalid columns
|
||||||
|
func (v *ColumnValidator) FilterValidColumns(columns []string) []string {
|
||||||
|
if len(columns) == 0 {
|
||||||
|
return columns
|
||||||
|
}
|
||||||
|
|
||||||
|
validColumns := make([]string, 0, len(columns))
|
||||||
|
for _, col := range columns {
|
||||||
|
if v.IsValidColumn(col) {
|
||||||
|
validColumns = append(validColumns, col)
|
||||||
|
} else {
|
||||||
|
logger.Warn("Invalid column '%s' filtered out: column does not exist in model", col)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return validColumns
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateColumns validates multiple column names
|
||||||
|
// Returns error with details about all invalid columns
|
||||||
|
func (v *ColumnValidator) ValidateColumns(columns []string) error {
|
||||||
|
var invalidColumns []string
|
||||||
|
|
||||||
|
for _, column := range columns {
|
||||||
|
if err := v.ValidateColumn(column); err != nil {
|
||||||
|
invalidColumns = append(invalidColumns, column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(invalidColumns) > 0 {
|
||||||
|
return fmt.Errorf("invalid columns: %s", strings.Join(invalidColumns, ", "))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateRequestOptions validates all column references in RequestOptions
|
||||||
|
func (v *ColumnValidator) ValidateRequestOptions(options RequestOptions) error {
|
||||||
|
// Validate Columns
|
||||||
|
if err := v.ValidateColumns(options.Columns); err != nil {
|
||||||
|
return fmt.Errorf("in select columns: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate OmitColumns
|
||||||
|
if err := v.ValidateColumns(options.OmitColumns); err != nil {
|
||||||
|
return fmt.Errorf("in omit columns: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate Filter columns
|
||||||
|
for _, filter := range options.Filters {
|
||||||
|
if err := v.ValidateColumn(filter.Column); err != nil {
|
||||||
|
return fmt.Errorf("in filter: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate Sort columns
|
||||||
|
for _, sort := range options.Sort {
|
||||||
|
if err := v.ValidateColumn(sort.Column); err != nil {
|
||||||
|
return fmt.Errorf("in sort: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate Preload columns (if specified)
|
||||||
|
for idx := range options.Preload {
|
||||||
|
preload := options.Preload[idx]
|
||||||
|
// Note: We don't validate the relation name itself, as it's a relationship
|
||||||
|
// Only validate columns if specified for the preload
|
||||||
|
if err := v.ValidateColumns(preload.Columns); err != nil {
|
||||||
|
return fmt.Errorf("in preload '%s' columns: %w", preload.Relation, err)
|
||||||
|
}
|
||||||
|
if err := v.ValidateColumns(preload.OmitColumns); err != nil {
|
||||||
|
return fmt.Errorf("in preload '%s' omit columns: %w", preload.Relation, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate filter columns in preload
|
||||||
|
for _, filter := range preload.Filters {
|
||||||
|
if err := v.ValidateColumn(filter.Column); err != nil {
|
||||||
|
return fmt.Errorf("in preload '%s' filter: %w", preload.Relation, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FilterRequestOptions filters all column references in RequestOptions
|
||||||
|
// Returns a new RequestOptions with only valid columns, logging warnings for invalid ones
|
||||||
|
func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOptions {
|
||||||
|
filtered := options
|
||||||
|
|
||||||
|
// Filter Columns
|
||||||
|
filtered.Columns = v.FilterValidColumns(options.Columns)
|
||||||
|
|
||||||
|
// Filter OmitColumns
|
||||||
|
filtered.OmitColumns = v.FilterValidColumns(options.OmitColumns)
|
||||||
|
|
||||||
|
// Filter Filter columns
|
||||||
|
validFilters := make([]FilterOption, 0, len(options.Filters))
|
||||||
|
for _, filter := range options.Filters {
|
||||||
|
if strings.EqualFold(filter.Column, "all") {
|
||||||
|
allCols := v.Columns()
|
||||||
|
if len(filtered.Columns) > 0 {
|
||||||
|
allCols = filtered.Columns
|
||||||
|
}
|
||||||
|
for _, col := range allCols {
|
||||||
|
expanded := filter
|
||||||
|
expanded.Column = col
|
||||||
|
expanded.LogicOperator = "OR"
|
||||||
|
|
||||||
|
validFilters = append(validFilters, expanded)
|
||||||
|
}
|
||||||
|
} else if v.IsValidColumn(filter.Column) {
|
||||||
|
validFilters = append(validFilters, filter)
|
||||||
|
} else {
|
||||||
|
logger.Warn("Invalid column in filter '%s' removed", filter.Column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
filtered.Filters = validFilters
|
||||||
|
|
||||||
|
// Filter Sort columns
|
||||||
|
validSorts := make([]SortOption, 0, len(options.Sort))
|
||||||
|
for _, sort := range options.Sort {
|
||||||
|
if v.IsValidColumn(sort.Column) {
|
||||||
|
validSorts = append(validSorts, sort)
|
||||||
|
} else {
|
||||||
|
foundJoin := false
|
||||||
|
for _, j := range options.JoinAliases {
|
||||||
|
if strings.Contains(sort.Column, j) {
|
||||||
|
foundJoin = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if foundJoin {
|
||||||
|
validSorts = append(validSorts, sort)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||||
|
// Allow sort by expression/subquery, but validate for security
|
||||||
|
if IsSafeSortExpression(sort.Column) {
|
||||||
|
validSorts = append(validSorts, sort)
|
||||||
|
} else {
|
||||||
|
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
||||||
|
}
|
||||||
|
|
||||||
|
} else {
|
||||||
|
logger.Warn("Invalid column in sort '%s' removed", sort.Column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
filtered.Sort = validSorts
|
||||||
|
|
||||||
|
// Filter Preload columns
|
||||||
|
validPreloads := make([]PreloadOption, 0, len(options.Preload))
|
||||||
|
modelType := reflect.TypeOf(v.model)
|
||||||
|
if modelType != nil && modelType.Kind() == reflect.Pointer {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
for idx := range options.Preload {
|
||||||
|
preload := options.Preload[idx]
|
||||||
|
filteredPreload := preload
|
||||||
|
|
||||||
|
// Use the related model's validator for preload columns/filters/sorts
|
||||||
|
preloadValidator := v
|
||||||
|
if modelType != nil {
|
||||||
|
if relInfo := GetRelationshipInfo(modelType, preload.Relation); relInfo != nil && relInfo.RelatedModel != nil {
|
||||||
|
preloadValidator = NewColumnValidator(relInfo.RelatedModel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
filteredPreload.Columns = preloadValidator.FilterValidColumns(preload.Columns)
|
||||||
|
filteredPreload.OmitColumns = preloadValidator.FilterValidColumns(preload.OmitColumns)
|
||||||
|
|
||||||
|
// Preserve SqlJoins and JoinAliases for preloads with custom joins
|
||||||
|
filteredPreload.SqlJoins = preload.SqlJoins
|
||||||
|
filteredPreload.JoinAliases = preload.JoinAliases
|
||||||
|
|
||||||
|
// Filter preload filters
|
||||||
|
validPreloadFilters := make([]FilterOption, 0, len(preload.Filters))
|
||||||
|
for _, filter := range preload.Filters {
|
||||||
|
if preloadValidator.IsValidColumn(filter.Column) {
|
||||||
|
validPreloadFilters = append(validPreloadFilters, filter)
|
||||||
|
} else {
|
||||||
|
// Check if the filter column references a joined table alias
|
||||||
|
foundJoin := false
|
||||||
|
for _, alias := range preload.JoinAliases {
|
||||||
|
if strings.Contains(filter.Column, alias) {
|
||||||
|
foundJoin = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if foundJoin {
|
||||||
|
validPreloadFilters = append(validPreloadFilters, filter)
|
||||||
|
} else {
|
||||||
|
logger.Warn("Invalid column in preload '%s' filter '%s' removed", preload.Relation, filter.Column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
filteredPreload.Filters = validPreloadFilters
|
||||||
|
|
||||||
|
// Filter preload sort columns
|
||||||
|
validPreloadSorts := make([]SortOption, 0, len(preload.Sort))
|
||||||
|
for _, sort := range preload.Sort {
|
||||||
|
if preloadValidator.IsValidColumn(sort.Column) {
|
||||||
|
validPreloadSorts = append(validPreloadSorts, sort)
|
||||||
|
} else if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||||
|
// Allow sort by expression/subquery, but validate for security
|
||||||
|
if IsSafeSortExpression(sort.Column) {
|
||||||
|
validPreloadSorts = append(validPreloadSorts, sort)
|
||||||
|
} else {
|
||||||
|
logger.Warn("Unsafe sort expression in preload '%s' removed: '%s'", preload.Relation, sort.Column)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
logger.Warn("Invalid column in preload '%s' sort '%s' removed", preload.Relation, sort.Column)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
filteredPreload.Sort = validPreloadSorts
|
||||||
|
|
||||||
|
validPreloads = append(validPreloads, filteredPreload)
|
||||||
|
}
|
||||||
|
filtered.Preload = validPreloads
|
||||||
|
|
||||||
|
// Clear JoinAliases - this is an internal validation field and should not be persisted
|
||||||
|
filtered.JoinAliases = nil
|
||||||
|
|
||||||
|
return filtered
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsSafeSortExpression validates that a sort expression (enclosed in brackets) is safe
|
||||||
|
// and doesn't contain SQL injection attempts or dangerous commands
|
||||||
|
func IsSafeSortExpression(expr string) bool {
|
||||||
|
if expr == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expression must be enclosed in brackets
|
||||||
|
expr = strings.TrimSpace(expr)
|
||||||
|
if !strings.HasPrefix(expr, "(") || !strings.HasSuffix(expr, ")") {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove outer brackets for content validation
|
||||||
|
expr = expr[1 : len(expr)-1]
|
||||||
|
expr = strings.TrimSpace(expr)
|
||||||
|
|
||||||
|
// Convert to lowercase for checking dangerous keywords
|
||||||
|
exprLower := strings.ToLower(expr)
|
||||||
|
|
||||||
|
// Check for dangerous SQL commands that should never be in a sort expression
|
||||||
|
dangerousKeywords := []string{
|
||||||
|
"drop ", "delete ", "insert ", "update ", "alter ", "create ",
|
||||||
|
"truncate ", "exec ", "execute ", "grant ", "revoke ",
|
||||||
|
"into ", "values ", "set ", "shutdown", "xp_",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, keyword := range dangerousKeywords {
|
||||||
|
if strings.Contains(exprLower, keyword) {
|
||||||
|
logger.Warn("Dangerous SQL keyword '%s' detected in sort expression: %s", keyword, expr)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for SQL comment attempts
|
||||||
|
if strings.Contains(expr, "--") || strings.Contains(expr, "/*") || strings.Contains(expr, "*/") {
|
||||||
|
logger.Warn("SQL comment detected in sort expression: %s", expr)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for semicolon (command separator)
|
||||||
|
if strings.Contains(expr, ";") {
|
||||||
|
logger.Warn("Command separator (;) detected in sort expression: %s", expr)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Expression appears safe
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetValidColumns returns a list of all valid column names for debugging purposes
|
||||||
|
func (v *ColumnValidator) GetValidColumns() []string {
|
||||||
|
columns := make([]string, 0, len(v.validColumns))
|
||||||
|
for col := range v.validColumns {
|
||||||
|
columns = append(columns, col)
|
||||||
|
}
|
||||||
|
return columns
|
||||||
|
}
|
||||||
|
|
||||||
|
func QuoteIdent(qualifier string) string {
|
||||||
|
return `"` + strings.ReplaceAll(qualifier, `"`, `""`) + `"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func QuoteLiteral(value string) string {
|
||||||
|
return `'` + strings.ReplaceAll(value, `'`, `''`) + `'`
|
||||||
|
}
|
||||||
+291
@@ -0,0 +1,291 @@
|
|||||||
|
# ResolveSpec Configuration System
|
||||||
|
|
||||||
|
A centralized configuration system with support for multiple configuration sources: config files (YAML, TOML, JSON), environment variables, and programmatic configuration.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- **Multiple Config Sources**: Config files, environment variables, and code
|
||||||
|
- **Priority Order**: Environment variables > Config file > Defaults
|
||||||
|
- **Multiple Formats**: YAML, TOML, JSON supported
|
||||||
|
- **Type Safety**: Strongly-typed configuration structs
|
||||||
|
- **Sensible Defaults**: Works out of the box with reasonable defaults
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
### Basic Usage
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/heinhel/ResolveSpec/pkg/config"
|
||||||
|
|
||||||
|
// Create a new config manager
|
||||||
|
mgr := config.NewManager()
|
||||||
|
|
||||||
|
// Load configuration from file and environment
|
||||||
|
if err := mgr.Load(); err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get the complete configuration
|
||||||
|
cfg, err := mgr.GetConfig()
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use the configuration
|
||||||
|
fmt.Println("Server address:", cfg.Server.Addr)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Custom Configuration Paths
|
||||||
|
|
||||||
|
```go
|
||||||
|
mgr := config.NewManagerWithOptions(
|
||||||
|
config.WithConfigFile("/path/to/config.yaml"),
|
||||||
|
config.WithEnvPrefix("MYAPP"),
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration Sources
|
||||||
|
|
||||||
|
### 1. Config Files
|
||||||
|
|
||||||
|
Place a `config.yaml` file in one of these locations:
|
||||||
|
- Current directory (`.`)
|
||||||
|
- `./config/`
|
||||||
|
- `/etc/resolvespec/`
|
||||||
|
- `$HOME/.resolvespec/`
|
||||||
|
|
||||||
|
Example `config.yaml`:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
server:
|
||||||
|
addr: ":8080"
|
||||||
|
shutdown_timeout: 30s
|
||||||
|
|
||||||
|
tracing:
|
||||||
|
enabled: true
|
||||||
|
service_name: "my-service"
|
||||||
|
|
||||||
|
cache:
|
||||||
|
provider: "redis"
|
||||||
|
redis:
|
||||||
|
host: "localhost"
|
||||||
|
port: 6379
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. Environment Variables
|
||||||
|
|
||||||
|
All configuration can be set via environment variables with the `RESOLVESPEC_` prefix:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export RESOLVESPEC_SERVER_ADDR=":9090"
|
||||||
|
export RESOLVESPEC_TRACING_ENABLED=true
|
||||||
|
export RESOLVESPEC_CACHE_PROVIDER=redis
|
||||||
|
export RESOLVESPEC_CACHE_REDIS_HOST=localhost
|
||||||
|
```
|
||||||
|
|
||||||
|
Nested configuration uses underscores:
|
||||||
|
- `server.addr` → `RESOLVESPEC_SERVER_ADDR`
|
||||||
|
- `cache.redis.host` → `RESOLVESPEC_CACHE_REDIS_HOST`
|
||||||
|
|
||||||
|
### 3. Programmatic Configuration
|
||||||
|
|
||||||
|
```go
|
||||||
|
mgr := config.NewManager()
|
||||||
|
mgr.Set("server.addr", ":9090")
|
||||||
|
mgr.Set("tracing.enabled", true)
|
||||||
|
|
||||||
|
cfg, _ := mgr.GetConfig()
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration Options
|
||||||
|
|
||||||
|
### Server Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
server:
|
||||||
|
addr: ":8080" # Server address
|
||||||
|
shutdown_timeout: 30s # Graceful shutdown timeout
|
||||||
|
drain_timeout: 25s # Connection drain timeout
|
||||||
|
read_timeout: 10s # HTTP read timeout
|
||||||
|
write_timeout: 10s # HTTP write timeout
|
||||||
|
idle_timeout: 120s # HTTP idle timeout
|
||||||
|
```
|
||||||
|
|
||||||
|
### Tracing Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
tracing:
|
||||||
|
enabled: false # Enable/disable tracing
|
||||||
|
service_name: "resolvespec" # Service name
|
||||||
|
service_version: "1.0.0" # Service version
|
||||||
|
endpoint: "http://localhost:4318/v1/traces" # OTLP endpoint
|
||||||
|
```
|
||||||
|
|
||||||
|
### Cache Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
cache:
|
||||||
|
provider: "memory" # Options: memory, redis, memcache
|
||||||
|
|
||||||
|
redis:
|
||||||
|
host: "localhost"
|
||||||
|
port: 6379
|
||||||
|
password: ""
|
||||||
|
db: 0
|
||||||
|
|
||||||
|
memcache:
|
||||||
|
servers:
|
||||||
|
- "localhost:11211"
|
||||||
|
max_idle_conns: 10
|
||||||
|
timeout: 100ms
|
||||||
|
```
|
||||||
|
|
||||||
|
### Logger Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
logger:
|
||||||
|
dev: false # Development mode (human-readable output)
|
||||||
|
path: "" # Log file path (empty = stdout)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Middleware Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
middleware:
|
||||||
|
rate_limit_rps: 100.0 # Requests per second
|
||||||
|
rate_limit_burst: 200 # Burst size
|
||||||
|
max_request_size: 10485760 # Max request size in bytes (10MB)
|
||||||
|
```
|
||||||
|
|
||||||
|
### CORS Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
cors:
|
||||||
|
allowed_origins:
|
||||||
|
- "*"
|
||||||
|
allowed_methods:
|
||||||
|
- "GET"
|
||||||
|
- "POST"
|
||||||
|
- "PUT"
|
||||||
|
- "DELETE"
|
||||||
|
- "OPTIONS"
|
||||||
|
allowed_headers:
|
||||||
|
- "*"
|
||||||
|
max_age: 3600
|
||||||
|
```
|
||||||
|
|
||||||
|
### Database Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
database:
|
||||||
|
url: "host=localhost user=postgres password=postgres dbname=mydb port=5432 sslmode=disable"
|
||||||
|
```
|
||||||
|
|
||||||
|
## Priority and Overrides
|
||||||
|
|
||||||
|
Configuration sources are applied in this order (highest priority first):
|
||||||
|
|
||||||
|
1. **Environment Variables** (highest priority)
|
||||||
|
2. **Config File**
|
||||||
|
3. **Defaults** (lowest priority)
|
||||||
|
|
||||||
|
This allows you to:
|
||||||
|
- Set defaults in code
|
||||||
|
- Override with a config file
|
||||||
|
- Override specific values with environment variables
|
||||||
|
|
||||||
|
## Examples
|
||||||
|
|
||||||
|
### Production Setup
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# config.yaml
|
||||||
|
server:
|
||||||
|
addr: ":8080"
|
||||||
|
|
||||||
|
tracing:
|
||||||
|
enabled: true
|
||||||
|
service_name: "myapi"
|
||||||
|
endpoint: "http://jaeger:4318/v1/traces"
|
||||||
|
|
||||||
|
cache:
|
||||||
|
provider: "redis"
|
||||||
|
redis:
|
||||||
|
host: "redis"
|
||||||
|
port: 6379
|
||||||
|
password: "${REDIS_PASSWORD}"
|
||||||
|
|
||||||
|
logger:
|
||||||
|
dev: false
|
||||||
|
path: "/var/log/myapi/app.log"
|
||||||
|
```
|
||||||
|
|
||||||
|
### Development Setup
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Use environment variables for development
|
||||||
|
export RESOLVESPEC_LOGGER_DEV=true
|
||||||
|
export RESOLVESPEC_TRACING_ENABLED=false
|
||||||
|
export RESOLVESPEC_CACHE_PROVIDER=memory
|
||||||
|
```
|
||||||
|
|
||||||
|
### Testing Setup
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Override config for tests
|
||||||
|
mgr := config.NewManager()
|
||||||
|
mgr.Set("cache.provider", "memory")
|
||||||
|
mgr.Set("database.url", testDBURL)
|
||||||
|
|
||||||
|
cfg, _ := mgr.GetConfig()
|
||||||
|
```
|
||||||
|
|
||||||
|
## Best Practices
|
||||||
|
|
||||||
|
1. **Use config files for base configuration** - Define your standard settings
|
||||||
|
2. **Use environment variables for secrets** - Never commit passwords/tokens
|
||||||
|
3. **Use environment variables for deployment-specific values** - Different per environment
|
||||||
|
4. **Keep defaults sensible** - Application should work with minimal configuration
|
||||||
|
5. **Document your configuration** - Comment your config.yaml files
|
||||||
|
|
||||||
|
## Integration with ResolveSpec Components
|
||||||
|
|
||||||
|
The configuration system integrates seamlessly with ResolveSpec components:
|
||||||
|
|
||||||
|
```go
|
||||||
|
cfg, _ := config.NewManager().Load().GetConfig()
|
||||||
|
|
||||||
|
// Server
|
||||||
|
srv := server.NewGracefulServer(server.Config{
|
||||||
|
Addr: cfg.Server.Addr,
|
||||||
|
ShutdownTimeout: cfg.Server.ShutdownTimeout,
|
||||||
|
// ... other fields
|
||||||
|
})
|
||||||
|
|
||||||
|
// Tracing
|
||||||
|
if cfg.Tracing.Enabled {
|
||||||
|
tracer := tracing.Init(tracing.Config{
|
||||||
|
ServiceName: cfg.Tracing.ServiceName,
|
||||||
|
ServiceVersion: cfg.Tracing.ServiceVersion,
|
||||||
|
Endpoint: cfg.Tracing.Endpoint,
|
||||||
|
})
|
||||||
|
defer tracer.Shutdown(context.Background())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cache
|
||||||
|
var cacheProvider cache.Provider
|
||||||
|
switch cfg.Cache.Provider {
|
||||||
|
case "redis":
|
||||||
|
cacheProvider = cache.NewRedisProvider(cfg.Cache.Redis.Host, cfg.Cache.Redis.Port, ...)
|
||||||
|
case "memcache":
|
||||||
|
cacheProvider = cache.NewMemcacheProvider(cfg.Cache.Memcache.Servers, ...)
|
||||||
|
default:
|
||||||
|
cacheProvider = cache.NewMemoryProvider()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Logger
|
||||||
|
logger.Init(cfg.Logger.Dev)
|
||||||
|
if cfg.Logger.Path != "" {
|
||||||
|
logger.UpdateLoggerPath(cfg.Logger.Path, cfg.Logger.Dev)
|
||||||
|
}
|
||||||
|
```
|
||||||
+196
@@ -0,0 +1,196 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import "time"
|
||||||
|
|
||||||
|
// Config represents the complete application configuration
|
||||||
|
type Config struct {
|
||||||
|
Servers ServersConfig `mapstructure:"servers"`
|
||||||
|
Tracing TracingConfig `mapstructure:"tracing"`
|
||||||
|
Cache CacheConfig `mapstructure:"cache"`
|
||||||
|
Logger LoggerConfig `mapstructure:"logger"`
|
||||||
|
ErrorTracking ErrorTrackingConfig `mapstructure:"error_tracking"`
|
||||||
|
Middleware MiddlewareConfig `mapstructure:"middleware"`
|
||||||
|
CORS CORSConfig `mapstructure:"cors"`
|
||||||
|
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
||||||
|
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
||||||
|
Paths PathsConfig `mapstructure:"paths"`
|
||||||
|
Extensions map[string]interface{} `mapstructure:"extensions"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServersConfig contains configuration for the server manager
|
||||||
|
type ServersConfig struct {
|
||||||
|
// DefaultServer is the name of the default server to use
|
||||||
|
DefaultServer string `mapstructure:"default_server"`
|
||||||
|
|
||||||
|
// Instances is a map of server name to server configuration
|
||||||
|
Instances map[string]ServerInstanceConfig `mapstructure:"instances"`
|
||||||
|
|
||||||
|
// Global timeout defaults (can be overridden per instance)
|
||||||
|
ShutdownTimeout time.Duration `mapstructure:"shutdown_timeout"`
|
||||||
|
DrainTimeout time.Duration `mapstructure:"drain_timeout"`
|
||||||
|
ReadTimeout time.Duration `mapstructure:"read_timeout"`
|
||||||
|
WriteTimeout time.Duration `mapstructure:"write_timeout"`
|
||||||
|
IdleTimeout time.Duration `mapstructure:"idle_timeout"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServerInstanceConfig defines configuration for a single server instance
|
||||||
|
type ServerInstanceConfig struct {
|
||||||
|
// Name is the unique name of this server instance
|
||||||
|
Name string `mapstructure:"name"`
|
||||||
|
|
||||||
|
// Host is the host to bind to (e.g., "localhost", "0.0.0.0", "")
|
||||||
|
Host string `mapstructure:"host"`
|
||||||
|
|
||||||
|
// Port is the port number to listen on
|
||||||
|
Port int `mapstructure:"port"`
|
||||||
|
|
||||||
|
// Description is a human-readable description of this server
|
||||||
|
Description string `mapstructure:"description"`
|
||||||
|
|
||||||
|
// GZIP enables GZIP compression middleware
|
||||||
|
GZIP bool `mapstructure:"gzip"`
|
||||||
|
|
||||||
|
// TLS/HTTPS configuration options (mutually exclusive)
|
||||||
|
// Option 1: Provide certificate and key files directly
|
||||||
|
SSLCert string `mapstructure:"ssl_cert"`
|
||||||
|
SSLKey string `mapstructure:"ssl_key"`
|
||||||
|
|
||||||
|
// Option 2: Use self-signed certificate (for development/testing)
|
||||||
|
SelfSignedSSL bool `mapstructure:"self_signed_ssl"`
|
||||||
|
|
||||||
|
// Option 3: Use Let's Encrypt / AutoTLS
|
||||||
|
AutoTLS bool `mapstructure:"auto_tls"`
|
||||||
|
AutoTLSDomains []string `mapstructure:"auto_tls_domains"`
|
||||||
|
AutoTLSCacheDir string `mapstructure:"auto_tls_cache_dir"`
|
||||||
|
AutoTLSEmail string `mapstructure:"auto_tls_email"`
|
||||||
|
|
||||||
|
// Timeout configurations (overrides global defaults)
|
||||||
|
ShutdownTimeout *time.Duration `mapstructure:"shutdown_timeout"`
|
||||||
|
DrainTimeout *time.Duration `mapstructure:"drain_timeout"`
|
||||||
|
ReadTimeout *time.Duration `mapstructure:"read_timeout"`
|
||||||
|
WriteTimeout *time.Duration `mapstructure:"write_timeout"`
|
||||||
|
IdleTimeout *time.Duration `mapstructure:"idle_timeout"`
|
||||||
|
|
||||||
|
// Tags for organization and filtering
|
||||||
|
Tags map[string]string `mapstructure:"tags"`
|
||||||
|
|
||||||
|
// ExternalURLs are additional URLs that this server instance is accessible from (for CORS) for proxy setups
|
||||||
|
ExternalURLs []string `mapstructure:"external_urls"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// TracingConfig holds OpenTelemetry tracing configuration
|
||||||
|
type TracingConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
ServiceName string `mapstructure:"service_name"`
|
||||||
|
ServiceVersion string `mapstructure:"service_version"`
|
||||||
|
Endpoint string `mapstructure:"endpoint"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// CacheConfig holds cache provider configuration
|
||||||
|
type CacheConfig struct {
|
||||||
|
Provider string `mapstructure:"provider"` // memory, redis, memcache
|
||||||
|
Redis RedisConfig `mapstructure:"redis"`
|
||||||
|
Memcache MemcacheConfig `mapstructure:"memcache"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// RedisConfig holds Redis-specific configuration
|
||||||
|
type RedisConfig struct {
|
||||||
|
Host string `mapstructure:"host"`
|
||||||
|
Port int `mapstructure:"port"`
|
||||||
|
Password string `mapstructure:"password"`
|
||||||
|
DB int `mapstructure:"db"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// MemcacheConfig holds Memcache-specific configuration
|
||||||
|
type MemcacheConfig struct {
|
||||||
|
Servers []string `mapstructure:"servers"`
|
||||||
|
MaxIdleConns int `mapstructure:"max_idle_conns"`
|
||||||
|
Timeout time.Duration `mapstructure:"timeout"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoggerConfig holds logger configuration
|
||||||
|
type LoggerConfig struct {
|
||||||
|
Dev bool `mapstructure:"dev"`
|
||||||
|
Path string `mapstructure:"path"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// MiddlewareConfig holds middleware configuration
|
||||||
|
type MiddlewareConfig struct {
|
||||||
|
RateLimitRPS float64 `mapstructure:"rate_limit_rps"`
|
||||||
|
RateLimitBurst int `mapstructure:"rate_limit_burst"`
|
||||||
|
MaxRequestSize int64 `mapstructure:"max_request_size"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// CORSConfig holds CORS configuration
|
||||||
|
type CORSConfig struct {
|
||||||
|
AllowedOrigins []string `mapstructure:"allowed_origins"`
|
||||||
|
AllowedMethods []string `mapstructure:"allowed_methods"`
|
||||||
|
AllowedHeaders []string `mapstructure:"allowed_headers"`
|
||||||
|
MaxAge int `mapstructure:"max_age"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ErrorTrackingConfig holds error tracking configuration
|
||||||
|
type ErrorTrackingConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
Provider string `mapstructure:"provider"` // sentry, noop
|
||||||
|
DSN string `mapstructure:"dsn"` // Sentry DSN
|
||||||
|
Environment string `mapstructure:"environment"` // e.g., production, staging, development
|
||||||
|
Release string `mapstructure:"release"` // Application version/release
|
||||||
|
Debug bool `mapstructure:"debug"` // Enable debug mode
|
||||||
|
SampleRate float64 `mapstructure:"sample_rate"` // Error sample rate (0.0-1.0)
|
||||||
|
TracesSampleRate float64 `mapstructure:"traces_sample_rate"` // Traces sample rate (0.0-1.0)
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventBrokerConfig contains configuration for the event broker
|
||||||
|
type EventBrokerConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
Provider string `mapstructure:"provider"` // memory, redis, nats, database
|
||||||
|
Mode string `mapstructure:"mode"` // sync, async
|
||||||
|
WorkerCount int `mapstructure:"worker_count"`
|
||||||
|
BufferSize int `mapstructure:"buffer_size"`
|
||||||
|
InstanceID string `mapstructure:"instance_id"`
|
||||||
|
Redis EventBrokerRedisConfig `mapstructure:"redis"`
|
||||||
|
NATS EventBrokerNATSConfig `mapstructure:"nats"`
|
||||||
|
Database EventBrokerDatabaseConfig `mapstructure:"database"`
|
||||||
|
RetryPolicy EventBrokerRetryPolicyConfig `mapstructure:"retry_policy"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventBrokerRedisConfig contains Redis-specific configuration
|
||||||
|
type EventBrokerRedisConfig struct {
|
||||||
|
StreamName string `mapstructure:"stream_name"`
|
||||||
|
ConsumerGroup string `mapstructure:"consumer_group"`
|
||||||
|
MaxLen int64 `mapstructure:"max_len"`
|
||||||
|
Host string `mapstructure:"host"`
|
||||||
|
Port int `mapstructure:"port"`
|
||||||
|
Password string `mapstructure:"password"`
|
||||||
|
DB int `mapstructure:"db"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventBrokerNATSConfig contains NATS-specific configuration
|
||||||
|
type EventBrokerNATSConfig struct {
|
||||||
|
URL string `mapstructure:"url"`
|
||||||
|
StreamName string `mapstructure:"stream_name"`
|
||||||
|
Subjects []string `mapstructure:"subjects"`
|
||||||
|
Storage string `mapstructure:"storage"` // file, memory
|
||||||
|
MaxAge time.Duration `mapstructure:"max_age"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventBrokerDatabaseConfig contains database provider configuration
|
||||||
|
type EventBrokerDatabaseConfig struct {
|
||||||
|
TableName string `mapstructure:"table_name"`
|
||||||
|
Channel string `mapstructure:"channel"` // PostgreSQL NOTIFY channel name
|
||||||
|
PollInterval time.Duration `mapstructure:"poll_interval"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventBrokerRetryPolicyConfig contains retry policy configuration
|
||||||
|
type EventBrokerRetryPolicyConfig struct {
|
||||||
|
MaxRetries int `mapstructure:"max_retries"`
|
||||||
|
InitialDelay time.Duration `mapstructure:"initial_delay"`
|
||||||
|
MaxDelay time.Duration `mapstructure:"max_delay"`
|
||||||
|
BackoffFactor float64 `mapstructure:"backoff_factor"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PathsConfig contains configuration for named file system paths
|
||||||
|
// This is a map of path name to file system path
|
||||||
|
// Example: "data_dir": "/var/lib/myapp/data"
|
||||||
|
type PathsConfig map[string]string
|
||||||
+264
@@ -0,0 +1,264 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/url"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DBManagerConfig contains configuration for the database connection manager
|
||||||
|
type DBManagerConfig struct {
|
||||||
|
// DefaultConnection is the name of the default connection to use
|
||||||
|
DefaultConnection string `mapstructure:"default_connection"`
|
||||||
|
|
||||||
|
// Connections is a map of connection name to connection configuration
|
||||||
|
Connections map[string]DBConnectionConfig `mapstructure:"connections"`
|
||||||
|
|
||||||
|
// Global connection pool defaults
|
||||||
|
MaxOpenConns int `mapstructure:"max_open_conns"`
|
||||||
|
MaxIdleConns int `mapstructure:"max_idle_conns"`
|
||||||
|
ConnMaxLifetime time.Duration `mapstructure:"conn_max_lifetime"`
|
||||||
|
ConnMaxIdleTime time.Duration `mapstructure:"conn_max_idle_time"`
|
||||||
|
|
||||||
|
// Retry policy
|
||||||
|
RetryAttempts int `mapstructure:"retry_attempts"`
|
||||||
|
RetryDelay time.Duration `mapstructure:"retry_delay"`
|
||||||
|
RetryMaxDelay time.Duration `mapstructure:"retry_max_delay"`
|
||||||
|
|
||||||
|
// Health checks
|
||||||
|
HealthCheckInterval time.Duration `mapstructure:"health_check_interval"`
|
||||||
|
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// DBConnectionConfig defines configuration for a single database connection
|
||||||
|
type DBConnectionConfig struct {
|
||||||
|
// Name is the unique name of this connection
|
||||||
|
Name string `mapstructure:"name"`
|
||||||
|
|
||||||
|
// Type is the database type (postgres, sqlite, mssql, mongodb)
|
||||||
|
Type string `mapstructure:"type"`
|
||||||
|
|
||||||
|
// DSN is the complete Data Source Name / connection string
|
||||||
|
// If provided, this takes precedence over individual connection parameters
|
||||||
|
DSN string `mapstructure:"dsn"`
|
||||||
|
|
||||||
|
// Connection parameters (used if DSN is not provided)
|
||||||
|
Host string `mapstructure:"host"`
|
||||||
|
Port int `mapstructure:"port"`
|
||||||
|
User string `mapstructure:"user"`
|
||||||
|
Password string `mapstructure:"password"`
|
||||||
|
Database string `mapstructure:"database"`
|
||||||
|
|
||||||
|
// PostgreSQL/MSSQL specific
|
||||||
|
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
|
||||||
|
Schema string `mapstructure:"schema"` // Default schema
|
||||||
|
|
||||||
|
// SQLite specific
|
||||||
|
FilePath string `mapstructure:"filepath"`
|
||||||
|
|
||||||
|
// MongoDB specific
|
||||||
|
AuthSource string `mapstructure:"auth_source"`
|
||||||
|
ReplicaSet string `mapstructure:"replica_set"`
|
||||||
|
ReadPreference string `mapstructure:"read_preference"` // primary, secondary, etc.
|
||||||
|
|
||||||
|
// Connection pool settings (overrides global defaults)
|
||||||
|
MaxOpenConns *int `mapstructure:"max_open_conns"`
|
||||||
|
MaxIdleConns *int `mapstructure:"max_idle_conns"`
|
||||||
|
ConnMaxLifetime *time.Duration `mapstructure:"conn_max_lifetime"`
|
||||||
|
ConnMaxIdleTime *time.Duration `mapstructure:"conn_max_idle_time"`
|
||||||
|
|
||||||
|
// Timeouts
|
||||||
|
ConnectTimeout time.Duration `mapstructure:"connect_timeout"`
|
||||||
|
QueryTimeout time.Duration `mapstructure:"query_timeout"`
|
||||||
|
|
||||||
|
// Features
|
||||||
|
EnableTracing bool `mapstructure:"enable_tracing"`
|
||||||
|
EnableMetrics bool `mapstructure:"enable_metrics"`
|
||||||
|
EnableLogging bool `mapstructure:"enable_logging"`
|
||||||
|
|
||||||
|
// DefaultORM specifies which ORM to use for the Database() method
|
||||||
|
// Options: "bun", "gorm", "native"
|
||||||
|
DefaultORM string `mapstructure:"default_orm"`
|
||||||
|
|
||||||
|
// Tags for organization and filtering
|
||||||
|
Tags map[string]string `mapstructure:"tags"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToManagerConfig converts config.DBManagerConfig to dbmanager.ManagerConfig
|
||||||
|
// This is used to avoid circular dependencies
|
||||||
|
func (c *DBManagerConfig) ToManagerConfig() interface{} {
|
||||||
|
// This will be implemented in the dbmanager package
|
||||||
|
// to convert from config types to dbmanager types
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// PopulateFromDSN parses a DSN and populates the connection fields
|
||||||
|
func (cc *DBConnectionConfig) PopulateFromDSN() error {
|
||||||
|
if cc.DSN == "" {
|
||||||
|
return nil // Nothing to populate
|
||||||
|
}
|
||||||
|
|
||||||
|
switch cc.Type {
|
||||||
|
case "postgres":
|
||||||
|
return cc.populatePostgresDSN()
|
||||||
|
case "mongodb":
|
||||||
|
return cc.populateMongoDSN()
|
||||||
|
case "mssql":
|
||||||
|
return cc.populateMSSQLDSN()
|
||||||
|
case "sqlite":
|
||||||
|
return cc.populateSQLiteDSN()
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("cannot parse DSN for unsupported database type: %s", cc.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// populatePostgresDSN parses PostgreSQL DSN format
|
||||||
|
// Example: host=localhost port=5432 user=postgres password=secret dbname=mydb sslmode=disable
|
||||||
|
func (cc *DBConnectionConfig) populatePostgresDSN() error {
|
||||||
|
parts := strings.Fields(cc.DSN)
|
||||||
|
for _, part := range parts {
|
||||||
|
kv := strings.SplitN(part, "=", 2)
|
||||||
|
if len(kv) != 2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
key, value := kv[0], kv[1]
|
||||||
|
|
||||||
|
switch key {
|
||||||
|
case "host":
|
||||||
|
cc.Host = value
|
||||||
|
case "port":
|
||||||
|
port, err := strconv.Atoi(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid port in DSN: %w", err)
|
||||||
|
}
|
||||||
|
cc.Port = port
|
||||||
|
case "user":
|
||||||
|
cc.User = value
|
||||||
|
case "password":
|
||||||
|
cc.Password = value
|
||||||
|
case "dbname":
|
||||||
|
cc.Database = value
|
||||||
|
case "sslmode":
|
||||||
|
cc.SSLMode = value
|
||||||
|
case "search_path":
|
||||||
|
cc.Schema = value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// populateMongoDSN parses MongoDB DSN format
|
||||||
|
// Example: mongodb://user:password@host:port/database?authSource=admin&replicaSet=rs0
|
||||||
|
func (cc *DBConnectionConfig) populateMongoDSN() error {
|
||||||
|
u, err := url.Parse(cc.DSN)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid MongoDB DSN: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract user and password
|
||||||
|
if u.User != nil {
|
||||||
|
cc.User = u.User.Username()
|
||||||
|
if password, ok := u.User.Password(); ok {
|
||||||
|
cc.Password = password
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract host and port
|
||||||
|
if u.Host != "" {
|
||||||
|
host := u.Host
|
||||||
|
if strings.Contains(host, ":") {
|
||||||
|
hostPort := strings.SplitN(host, ":", 2)
|
||||||
|
cc.Host = hostPort[0]
|
||||||
|
if port, err := strconv.Atoi(hostPort[1]); err == nil {
|
||||||
|
cc.Port = port
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
cc.Host = host
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract database
|
||||||
|
if u.Path != "" {
|
||||||
|
cc.Database = strings.TrimPrefix(u.Path, "/")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract query parameters
|
||||||
|
params := u.Query()
|
||||||
|
if authSource := params.Get("authSource"); authSource != "" {
|
||||||
|
cc.AuthSource = authSource
|
||||||
|
}
|
||||||
|
if replicaSet := params.Get("replicaSet"); replicaSet != "" {
|
||||||
|
cc.ReplicaSet = replicaSet
|
||||||
|
}
|
||||||
|
if readPref := params.Get("readPreference"); readPref != "" {
|
||||||
|
cc.ReadPreference = readPref
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// populateMSSQLDSN parses MSSQL DSN format
|
||||||
|
// Example: sqlserver://username:password@host:port?database=dbname&schema=dbo
|
||||||
|
func (cc *DBConnectionConfig) populateMSSQLDSN() error {
|
||||||
|
u, err := url.Parse(cc.DSN)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid MSSQL DSN: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract user and password
|
||||||
|
if u.User != nil {
|
||||||
|
cc.User = u.User.Username()
|
||||||
|
if password, ok := u.User.Password(); ok {
|
||||||
|
cc.Password = password
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract host and port
|
||||||
|
if u.Host != "" {
|
||||||
|
host := u.Host
|
||||||
|
if strings.Contains(host, ":") {
|
||||||
|
hostPort := strings.SplitN(host, ":", 2)
|
||||||
|
cc.Host = hostPort[0]
|
||||||
|
if port, err := strconv.Atoi(hostPort[1]); err == nil {
|
||||||
|
cc.Port = port
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
cc.Host = host
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract query parameters
|
||||||
|
params := u.Query()
|
||||||
|
if database := params.Get("database"); database != "" {
|
||||||
|
cc.Database = database
|
||||||
|
}
|
||||||
|
if schema := params.Get("schema"); schema != "" {
|
||||||
|
cc.Schema = schema
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// populateSQLiteDSN parses SQLite DSN format
|
||||||
|
// Example: /path/to/database.db or :memory:
|
||||||
|
func (cc *DBConnectionConfig) populateSQLiteDSN() error {
|
||||||
|
cc.FilePath = cc.DSN
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate validates the DBManager configuration
|
||||||
|
func (c *DBManagerConfig) Validate() error {
|
||||||
|
if len(c.Connections) == 0 {
|
||||||
|
return fmt.Errorf("at least one connection must be configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.DefaultConnection != "" {
|
||||||
|
if _, ok := c.Connections[c.DefaultConnection]; !ok {
|
||||||
|
return fmt.Errorf("default connection '%s' not found in connections", c.DefaultConnection)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+293
@@ -0,0 +1,293 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/spf13/viper"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Manager handles configuration loading from multiple sources
|
||||||
|
type Manager struct {
|
||||||
|
v *viper.Viper
|
||||||
|
}
|
||||||
|
|
||||||
|
var configInstance *Manager
|
||||||
|
|
||||||
|
// GetConfigManager returns a singleton configuration manager instance
|
||||||
|
func GetConfigManager() *Manager {
|
||||||
|
if configInstance == nil {
|
||||||
|
configInstance = NewManager()
|
||||||
|
}
|
||||||
|
return configInstance
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewManager creates a new configuration manager with defaults
|
||||||
|
func NewManager() *Manager {
|
||||||
|
v := viper.New()
|
||||||
|
|
||||||
|
// Set configuration file settings
|
||||||
|
v.SetConfigName("config")
|
||||||
|
v.SetConfigType("yaml")
|
||||||
|
v.AddConfigPath(".")
|
||||||
|
v.AddConfigPath("./config")
|
||||||
|
v.AddConfigPath("/etc/resolvespec")
|
||||||
|
v.AddConfigPath("$HOME/.resolvespec")
|
||||||
|
|
||||||
|
// Enable environment variable support
|
||||||
|
v.SetEnvPrefix("RESOLVESPEC")
|
||||||
|
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
|
||||||
|
v.AutomaticEnv()
|
||||||
|
|
||||||
|
// Set default values
|
||||||
|
setDefaults(v)
|
||||||
|
|
||||||
|
configInstance = &Manager{v: v}
|
||||||
|
return configInstance
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewManagerWithOptions creates a new configuration manager with custom options
|
||||||
|
func NewManagerWithOptions(opts ...Option) *Manager {
|
||||||
|
m := NewManager()
|
||||||
|
for _, opt := range opts {
|
||||||
|
opt(m)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// Option is a functional option for configuring the Manager
|
||||||
|
type Option func(*Manager)
|
||||||
|
|
||||||
|
// WithConfigFile sets a specific config file path
|
||||||
|
func WithConfigFile(path string) Option {
|
||||||
|
return func(m *Manager) {
|
||||||
|
m.v.SetConfigFile(path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithConfigName sets the config file name (without extension)
|
||||||
|
func WithConfigName(name string) Option {
|
||||||
|
return func(m *Manager) {
|
||||||
|
m.v.SetConfigName(name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithConfigPath adds a path to search for config files
|
||||||
|
func WithConfigPath(path string) Option {
|
||||||
|
return func(m *Manager) {
|
||||||
|
m.v.AddConfigPath(path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithEnvPrefix sets the environment variable prefix
|
||||||
|
func WithEnvPrefix(prefix string) Option {
|
||||||
|
return func(m *Manager) {
|
||||||
|
m.v.SetEnvPrefix(prefix)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load attempts to load configuration from file and environment
|
||||||
|
func (m *Manager) Load() error {
|
||||||
|
// Try to read config file (not an error if it doesn't exist)
|
||||||
|
if err := m.v.ReadInConfig(); err != nil {
|
||||||
|
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
|
||||||
|
return fmt.Errorf("error reading config file: %w", err)
|
||||||
|
}
|
||||||
|
// Config file not found; will rely on defaults and env vars
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetConfig returns the complete configuration
|
||||||
|
func (m *Manager) GetConfig() (*Config, error) {
|
||||||
|
var cfg Config
|
||||||
|
if err := m.v.Unmarshal(&cfg); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal config: %w", err)
|
||||||
|
}
|
||||||
|
return &cfg, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetConfig sets the complete configuration
|
||||||
|
func (m *Manager) SetConfig(cfg *Config) error {
|
||||||
|
configMap := make(map[string]interface{})
|
||||||
|
|
||||||
|
// Marshal the config to a map structure that viper can use
|
||||||
|
if err := m.v.Unmarshal(&configMap); err != nil {
|
||||||
|
return fmt.Errorf("failed to prepare config map: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use viper's merge to apply the config
|
||||||
|
m.v.Set("servers", cfg.Servers)
|
||||||
|
m.v.Set("tracing", cfg.Tracing)
|
||||||
|
m.v.Set("cache", cfg.Cache)
|
||||||
|
m.v.Set("logger", cfg.Logger)
|
||||||
|
m.v.Set("error_tracking", cfg.ErrorTracking)
|
||||||
|
m.v.Set("middleware", cfg.Middleware)
|
||||||
|
m.v.Set("cors", cfg.CORS)
|
||||||
|
m.v.Set("event_broker", cfg.EventBroker)
|
||||||
|
m.v.Set("dbmanager", cfg.DBManager)
|
||||||
|
m.v.Set("paths", cfg.Paths)
|
||||||
|
m.v.Set("extensions", cfg.Extensions)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get returns a configuration value by key
|
||||||
|
func (m *Manager) Get(key string) interface{} {
|
||||||
|
return m.v.Get(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetString returns a string configuration value
|
||||||
|
func (m *Manager) GetString(key string) string {
|
||||||
|
return m.v.GetString(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetInt returns an int configuration value
|
||||||
|
func (m *Manager) GetInt(key string) int {
|
||||||
|
return m.v.GetInt(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetBool returns a bool configuration value
|
||||||
|
func (m *Manager) GetBool(key string) bool {
|
||||||
|
return m.v.GetBool(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set sets a configuration value
|
||||||
|
func (m *Manager) Set(key string, value interface{}) {
|
||||||
|
m.v.Set(key, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveConfig writes the current configuration to the specified path
|
||||||
|
func (m *Manager) SaveConfig(path string) error {
|
||||||
|
if err := m.v.WriteConfigAs(path); err != nil {
|
||||||
|
return fmt.Errorf("failed to save config to %s: %w", path, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// setDefaults sets default configuration values
|
||||||
|
func setDefaults(v *viper.Viper) {
|
||||||
|
// Server defaults - new structure
|
||||||
|
v.SetDefault("servers.default_server", "default")
|
||||||
|
|
||||||
|
// Global server timeout defaults
|
||||||
|
v.SetDefault("servers.shutdown_timeout", "30s")
|
||||||
|
v.SetDefault("servers.drain_timeout", "25s")
|
||||||
|
v.SetDefault("servers.read_timeout", "10s")
|
||||||
|
v.SetDefault("servers.write_timeout", "10s")
|
||||||
|
v.SetDefault("servers.idle_timeout", "120s")
|
||||||
|
|
||||||
|
// Default server instance
|
||||||
|
v.SetDefault("servers.instances.default.name", "default")
|
||||||
|
v.SetDefault("servers.instances.default.host", "")
|
||||||
|
v.SetDefault("servers.instances.default.port", 8080)
|
||||||
|
v.SetDefault("servers.instances.default.description", "Default HTTP server")
|
||||||
|
v.SetDefault("servers.instances.default.gzip", false)
|
||||||
|
|
||||||
|
// Tracing defaults
|
||||||
|
v.SetDefault("tracing.enabled", false)
|
||||||
|
v.SetDefault("tracing.service_name", "resolvespec")
|
||||||
|
v.SetDefault("tracing.service_version", "1.0.0")
|
||||||
|
v.SetDefault("tracing.endpoint", "")
|
||||||
|
|
||||||
|
// Cache defaults
|
||||||
|
v.SetDefault("cache.provider", "memory")
|
||||||
|
v.SetDefault("cache.redis.host", "localhost")
|
||||||
|
v.SetDefault("cache.redis.port", 6379)
|
||||||
|
v.SetDefault("cache.redis.password", "")
|
||||||
|
v.SetDefault("cache.redis.db", 0)
|
||||||
|
v.SetDefault("cache.memcache.servers", []string{"localhost:11211"})
|
||||||
|
v.SetDefault("cache.memcache.max_idle_conns", 10)
|
||||||
|
v.SetDefault("cache.memcache.timeout", "100ms")
|
||||||
|
|
||||||
|
// Logger defaults
|
||||||
|
v.SetDefault("logger.dev", false)
|
||||||
|
v.SetDefault("logger.path", "")
|
||||||
|
|
||||||
|
// Middleware defaults
|
||||||
|
v.SetDefault("middleware.rate_limit_rps", 100.0)
|
||||||
|
v.SetDefault("middleware.rate_limit_burst", 200)
|
||||||
|
v.SetDefault("middleware.max_request_size", 10485760) // 10MB
|
||||||
|
|
||||||
|
// CORS defaults
|
||||||
|
v.SetDefault("cors.allowed_origins", []string{"*"})
|
||||||
|
v.SetDefault("cors.allowed_methods", []string{"GET", "POST", "PUT", "DELETE", "OPTIONS"})
|
||||||
|
v.SetDefault("cors.allowed_headers", []string{"*"})
|
||||||
|
v.SetDefault("cors.max_age", 3600)
|
||||||
|
|
||||||
|
// Database defaults
|
||||||
|
v.SetDefault("database.url", "")
|
||||||
|
|
||||||
|
// Database Manager defaults
|
||||||
|
v.SetDefault("dbmanager.default_connection", "default")
|
||||||
|
v.SetDefault("dbmanager.max_open_conns", 25)
|
||||||
|
v.SetDefault("dbmanager.max_idle_conns", 5)
|
||||||
|
v.SetDefault("dbmanager.conn_max_lifetime", "30m")
|
||||||
|
v.SetDefault("dbmanager.conn_max_idle_time", "5m")
|
||||||
|
v.SetDefault("dbmanager.retry_attempts", 3)
|
||||||
|
v.SetDefault("dbmanager.retry_delay", "1s")
|
||||||
|
v.SetDefault("dbmanager.retry_max_delay", "10s")
|
||||||
|
v.SetDefault("dbmanager.health_check_interval", "30s")
|
||||||
|
v.SetDefault("dbmanager.enable_auto_reconnect", true)
|
||||||
|
|
||||||
|
// Default PostgreSQL connection
|
||||||
|
v.SetDefault("dbmanager.connections.default.name", "default")
|
||||||
|
v.SetDefault("dbmanager.connections.default.type", "postgres")
|
||||||
|
v.SetDefault("dbmanager.connections.default.host", "localhost")
|
||||||
|
v.SetDefault("dbmanager.connections.default.port", 5432)
|
||||||
|
v.SetDefault("dbmanager.connections.default.user", "postgres")
|
||||||
|
v.SetDefault("dbmanager.connections.default.password", "")
|
||||||
|
v.SetDefault("dbmanager.connections.default.database", "resolvespec")
|
||||||
|
v.SetDefault("dbmanager.connections.default.sslmode", "disable")
|
||||||
|
v.SetDefault("dbmanager.connections.default.connect_timeout", "10s")
|
||||||
|
v.SetDefault("dbmanager.connections.default.query_timeout", "30s")
|
||||||
|
v.SetDefault("dbmanager.connections.default.enable_tracing", false)
|
||||||
|
v.SetDefault("dbmanager.connections.default.enable_metrics", false)
|
||||||
|
v.SetDefault("dbmanager.connections.default.enable_logging", false)
|
||||||
|
v.SetDefault("dbmanager.connections.default.default_orm", "bun")
|
||||||
|
|
||||||
|
// Event Broker defaults
|
||||||
|
v.SetDefault("event_broker.enabled", false)
|
||||||
|
v.SetDefault("event_broker.provider", "memory")
|
||||||
|
v.SetDefault("event_broker.mode", "async")
|
||||||
|
v.SetDefault("event_broker.worker_count", 10)
|
||||||
|
v.SetDefault("event_broker.buffer_size", 1000)
|
||||||
|
v.SetDefault("event_broker.instance_id", "")
|
||||||
|
|
||||||
|
// Event Broker - Redis defaults
|
||||||
|
v.SetDefault("event_broker.redis.stream_name", "resolvespec:events")
|
||||||
|
v.SetDefault("event_broker.redis.consumer_group", "resolvespec-workers")
|
||||||
|
v.SetDefault("event_broker.redis.max_len", 10000)
|
||||||
|
v.SetDefault("event_broker.redis.host", "localhost")
|
||||||
|
v.SetDefault("event_broker.redis.port", 6379)
|
||||||
|
v.SetDefault("event_broker.redis.password", "")
|
||||||
|
v.SetDefault("event_broker.redis.db", 0)
|
||||||
|
|
||||||
|
// Event Broker - NATS defaults
|
||||||
|
v.SetDefault("event_broker.nats.url", "nats://localhost:4222")
|
||||||
|
v.SetDefault("event_broker.nats.stream_name", "RESOLVESPEC_EVENTS")
|
||||||
|
v.SetDefault("event_broker.nats.subjects", []string{"events.>"})
|
||||||
|
v.SetDefault("event_broker.nats.storage", "file")
|
||||||
|
v.SetDefault("event_broker.nats.max_age", "24h")
|
||||||
|
|
||||||
|
// Event Broker - Database defaults
|
||||||
|
v.SetDefault("event_broker.database.table_name", "events")
|
||||||
|
v.SetDefault("event_broker.database.channel", "resolvespec_events")
|
||||||
|
v.SetDefault("event_broker.database.poll_interval", "1s")
|
||||||
|
|
||||||
|
// Event Broker - Retry Policy defaults
|
||||||
|
v.SetDefault("event_broker.retry_policy.max_retries", 3)
|
||||||
|
v.SetDefault("event_broker.retry_policy.initial_delay", "1s")
|
||||||
|
v.SetDefault("event_broker.retry_policy.max_delay", "30s")
|
||||||
|
v.SetDefault("event_broker.retry_policy.backoff_factor", 2.0)
|
||||||
|
|
||||||
|
// Paths defaults (common directory paths)
|
||||||
|
v.SetDefault("paths.data_dir", "./data")
|
||||||
|
v.SetDefault("paths.config_dir", "./config")
|
||||||
|
v.SetDefault("paths.logs_dir", "./logs")
|
||||||
|
v.SetDefault("paths.temp_dir", "./tmp")
|
||||||
|
|
||||||
|
// Extensions defaults (empty map)
|
||||||
|
v.SetDefault("extensions", map[string]interface{}{})
|
||||||
|
}
|
||||||
+117
@@ -0,0 +1,117 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Get retrieves a path by name
|
||||||
|
func (pc PathsConfig) Get(name string) (string, error) {
|
||||||
|
if pc == nil {
|
||||||
|
return "", fmt.Errorf("paths not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
path, ok := pc[name]
|
||||||
|
if !ok {
|
||||||
|
return "", fmt.Errorf("path '%s' not found", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
return path, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetOrDefault retrieves a path by name, returning defaultPath if not found
|
||||||
|
func (pc PathsConfig) GetOrDefault(name, defaultPath string) string {
|
||||||
|
if pc == nil {
|
||||||
|
return defaultPath
|
||||||
|
}
|
||||||
|
|
||||||
|
path, ok := pc[name]
|
||||||
|
if !ok {
|
||||||
|
return defaultPath
|
||||||
|
}
|
||||||
|
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set sets a path by name
|
||||||
|
func (pc PathsConfig) Set(name, path string) {
|
||||||
|
pc[name] = path
|
||||||
|
}
|
||||||
|
|
||||||
|
// Has checks if a path exists by name
|
||||||
|
func (pc PathsConfig) Has(name string) bool {
|
||||||
|
if pc == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_, ok := pc[name]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// EnsureDir ensures a directory exists at the specified path name
|
||||||
|
// Creates the directory if it doesn't exist with the given permissions
|
||||||
|
func (pc PathsConfig) EnsureDir(name string, perm os.FileMode) error {
|
||||||
|
path, err := pc.Get(name)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if directory exists
|
||||||
|
info, err := os.Stat(path)
|
||||||
|
if err == nil {
|
||||||
|
// Path exists, check if it's a directory
|
||||||
|
if !info.IsDir() {
|
||||||
|
return fmt.Errorf("path '%s' exists but is not a directory: %s", name, path)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Directory doesn't exist, create it
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
if err := os.MkdirAll(path, perm); err != nil {
|
||||||
|
return fmt.Errorf("failed to create directory for '%s' at %s: %w", name, path, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("failed to stat path '%s' at %s: %w", name, path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AbsPath returns the absolute path for a named path
|
||||||
|
func (pc PathsConfig) AbsPath(name string) (string, error) {
|
||||||
|
path, err := pc.Get(name)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
absPath, err := filepath.Abs(path)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to get absolute path for '%s': %w", name, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return absPath, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Join joins path segments with a named base path
|
||||||
|
func (pc PathsConfig) Join(name string, elem ...string) (string, error) {
|
||||||
|
base, err := pc.Get(name)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
parts := append([]string{base}, elem...)
|
||||||
|
return filepath.Join(parts...), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// List returns all configured path names
|
||||||
|
func (pc PathsConfig) List() []string {
|
||||||
|
if pc == nil {
|
||||||
|
return []string{}
|
||||||
|
}
|
||||||
|
|
||||||
|
names := make([]string, 0, len(pc))
|
||||||
|
for name := range pc {
|
||||||
|
names = append(names, name)
|
||||||
|
}
|
||||||
|
return names
|
||||||
|
}
|
||||||
+149
@@ -0,0 +1,149 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ApplyGlobalDefaults applies global server defaults to this instance
|
||||||
|
// Called for instances that don't specify their own timeout values
|
||||||
|
func (sic *ServerInstanceConfig) ApplyGlobalDefaults(globals ServersConfig) {
|
||||||
|
if sic.ShutdownTimeout == nil && globals.ShutdownTimeout > 0 {
|
||||||
|
t := globals.ShutdownTimeout
|
||||||
|
sic.ShutdownTimeout = &t
|
||||||
|
}
|
||||||
|
if sic.DrainTimeout == nil && globals.DrainTimeout > 0 {
|
||||||
|
t := globals.DrainTimeout
|
||||||
|
sic.DrainTimeout = &t
|
||||||
|
}
|
||||||
|
if sic.ReadTimeout == nil && globals.ReadTimeout > 0 {
|
||||||
|
t := globals.ReadTimeout
|
||||||
|
sic.ReadTimeout = &t
|
||||||
|
}
|
||||||
|
if sic.WriteTimeout == nil && globals.WriteTimeout > 0 {
|
||||||
|
t := globals.WriteTimeout
|
||||||
|
sic.WriteTimeout = &t
|
||||||
|
}
|
||||||
|
if sic.IdleTimeout == nil && globals.IdleTimeout > 0 {
|
||||||
|
t := globals.IdleTimeout
|
||||||
|
sic.IdleTimeout = &t
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate validates the ServerInstanceConfig
|
||||||
|
func (sic *ServerInstanceConfig) Validate() error {
|
||||||
|
if sic.Name == "" {
|
||||||
|
return fmt.Errorf("server instance name cannot be empty")
|
||||||
|
}
|
||||||
|
if sic.Port <= 0 || sic.Port > 65535 {
|
||||||
|
return fmt.Errorf("invalid port: %d (must be 1-65535)", sic.Port)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate TLS options are mutually exclusive
|
||||||
|
tlsCount := 0
|
||||||
|
if sic.SSLCert != "" || sic.SSLKey != "" {
|
||||||
|
tlsCount++
|
||||||
|
}
|
||||||
|
if sic.SelfSignedSSL {
|
||||||
|
tlsCount++
|
||||||
|
}
|
||||||
|
if sic.AutoTLS {
|
||||||
|
tlsCount++
|
||||||
|
}
|
||||||
|
if tlsCount > 1 {
|
||||||
|
return fmt.Errorf("server '%s': only one TLS option can be enabled", sic.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If using certificate files, both must be provided
|
||||||
|
if (sic.SSLCert != "" && sic.SSLKey == "") || (sic.SSLCert == "" && sic.SSLKey != "") {
|
||||||
|
return fmt.Errorf("server '%s': both ssl_cert and ssl_key must be provided", sic.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If using AutoTLS, domains must be specified
|
||||||
|
if sic.AutoTLS && len(sic.AutoTLSDomains) == 0 {
|
||||||
|
return fmt.Errorf("server '%s': auto_tls_domains must be specified when auto_tls is enabled", sic.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate validates the ServersConfig
|
||||||
|
func (sc *ServersConfig) Validate() error {
|
||||||
|
if len(sc.Instances) == 0 {
|
||||||
|
return fmt.Errorf("at least one server instance must be configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
if sc.DefaultServer != "" {
|
||||||
|
if _, ok := sc.Instances[sc.DefaultServer]; !ok {
|
||||||
|
return fmt.Errorf("default server '%s' not found in instances", sc.DefaultServer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate each instance
|
||||||
|
for name := range sc.Instances {
|
||||||
|
instance := sc.Instances[name]
|
||||||
|
if instance.Name != name {
|
||||||
|
return fmt.Errorf("server instance name mismatch: key='%s', name='%s'", name, instance.Name)
|
||||||
|
}
|
||||||
|
if err := instance.Validate(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDefault returns the default server instance configuration
|
||||||
|
func (sc *ServersConfig) GetDefault() (*ServerInstanceConfig, error) {
|
||||||
|
if sc.DefaultServer == "" {
|
||||||
|
return nil, fmt.Errorf("no default server configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
instance, ok := sc.Instances[sc.DefaultServer]
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("default server '%s' not found", sc.DefaultServer)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &instance, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetIPs - GetIP for pc
|
||||||
|
func GetIPs() (hostname string, ipList string, ipNetList []net.IP) {
|
||||||
|
defer func() {
|
||||||
|
if err := recover(); err != nil {
|
||||||
|
fmt.Println("Recovered in GetIPs", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
hostname, _ = os.Hostname()
|
||||||
|
ipaddrlist := make([]net.IP, 0)
|
||||||
|
iplist := ""
|
||||||
|
addrs, err := net.LookupIP(hostname)
|
||||||
|
if err != nil {
|
||||||
|
return hostname, iplist, ipaddrlist
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, a := range addrs {
|
||||||
|
// cfg.LogInfo("\nFound IP Host Address: %s", a)
|
||||||
|
if strings.Contains(a.String(), "127.0.0.1") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
iplist = fmt.Sprintf("%s,%s", iplist, a)
|
||||||
|
ipaddrlist = append(ipaddrlist, a)
|
||||||
|
}
|
||||||
|
if iplist == "" {
|
||||||
|
iff, _ := net.InterfaceAddrs()
|
||||||
|
for _, a := range iff {
|
||||||
|
// cfg.LogInfo("\nFound IP Address: %s", a)
|
||||||
|
if strings.Contains(a.String(), "127.0.0.1") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
iplist = fmt.Sprintf("%s,%s", iplist, a)
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
iplist = strings.TrimLeft(iplist, ",")
|
||||||
|
return hostname, iplist, ipaddrlist
|
||||||
|
}
|
||||||
+150
@@ -0,0 +1,150 @@
|
|||||||
|
# Error Tracking
|
||||||
|
|
||||||
|
This package provides error tracking integration for ResolveSpec, with built-in support for Sentry.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- **Provider Interface**: Flexible design supporting multiple error tracking backends
|
||||||
|
- **Sentry Integration**: Full-featured Sentry support with automatic error, warning, and panic tracking
|
||||||
|
- **Automatic Logger Integration**: All `logger.Error()` and `logger.Warn()` calls are automatically sent to the error tracker
|
||||||
|
- **Panic Tracking**: Automatic panic capture with stack traces
|
||||||
|
- **NoOp Provider**: Zero-overhead when error tracking is disabled
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
Add error tracking configuration to your config file:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
error_tracking:
|
||||||
|
enabled: true
|
||||||
|
provider: "sentry" # Currently supports: "sentry" or "noop"
|
||||||
|
dsn: "https://your-sentry-dsn@sentry.io/project-id"
|
||||||
|
environment: "production" # e.g., production, staging, development
|
||||||
|
release: "v1.0.0" # Your application version
|
||||||
|
debug: false
|
||||||
|
sample_rate: 1.0 # Error sample rate (0.0-1.0)
|
||||||
|
traces_sample_rate: 0.1 # Traces sample rate (0.0-1.0)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
### Initialization
|
||||||
|
|
||||||
|
Initialize error tracking in your application startup:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/errortracking"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Load your configuration
|
||||||
|
cfg := config.Config{
|
||||||
|
ErrorTracking: config.ErrorTrackingConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Provider: "sentry",
|
||||||
|
DSN: "https://your-sentry-dsn@sentry.io/project-id",
|
||||||
|
Environment: "production",
|
||||||
|
Release: "v1.0.0",
|
||||||
|
SampleRate: 1.0,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize logger
|
||||||
|
logger.Init(false)
|
||||||
|
|
||||||
|
// Initialize error tracking
|
||||||
|
provider, err := errortracking.NewProviderFromConfig(cfg.ErrorTracking)
|
||||||
|
if err != nil {
|
||||||
|
logger.Error("Failed to initialize error tracking: %v", err)
|
||||||
|
} else {
|
||||||
|
logger.InitErrorTracking(provider)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Your application code...
|
||||||
|
|
||||||
|
// Cleanup on shutdown
|
||||||
|
defer logger.CloseErrorTracking()
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Automatic Tracking
|
||||||
|
|
||||||
|
Once initialized, all logger errors and warnings are automatically sent to the error tracker:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// This will be logged AND sent to Sentry
|
||||||
|
logger.Error("Database connection failed: %v", err)
|
||||||
|
|
||||||
|
// This will also be logged AND sent to Sentry
|
||||||
|
logger.Warn("Cache miss for key: %s", key)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Panic Tracking
|
||||||
|
|
||||||
|
Panics are automatically captured when using the logger's panic handlers:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Using CatchPanic
|
||||||
|
defer logger.CatchPanic("MyFunction")()
|
||||||
|
|
||||||
|
// Using CatchPanicCallback
|
||||||
|
defer logger.CatchPanicCallback("MyFunction", func(err any) {
|
||||||
|
// Custom cleanup
|
||||||
|
})()
|
||||||
|
|
||||||
|
// Using HandlePanic
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = logger.HandlePanic("MyMethod", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
```
|
||||||
|
|
||||||
|
### Manual Tracking
|
||||||
|
|
||||||
|
You can also use the provider directly for custom error tracking:
|
||||||
|
|
||||||
|
```go
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/errortracking"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
func someFunction() {
|
||||||
|
tracker := logger.GetErrorTracker()
|
||||||
|
if tracker != nil {
|
||||||
|
// Capture an error
|
||||||
|
tracker.CaptureError(context.Background(), err, errortracking.SeverityError, map[string]interface{}{
|
||||||
|
"user_id": userID,
|
||||||
|
"request_id": requestID,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Capture a message
|
||||||
|
tracker.CaptureMessage(context.Background(), "Important event occurred", errortracking.SeverityInfo, map[string]interface{}{
|
||||||
|
"event_type": "user_signup",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Capture a panic
|
||||||
|
tracker.CapturePanic(context.Background(), recovered, stackTrace, map[string]interface{}{
|
||||||
|
"context": "background_job",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Severity Levels
|
||||||
|
|
||||||
|
The package supports the following severity levels:
|
||||||
|
|
||||||
|
- `SeverityError`: For errors that should be tracked and investigated
|
||||||
|
- `SeverityWarning`: For warnings that may indicate potential issues
|
||||||
|
- `SeverityInfo`: For informational messages
|
||||||
|
- `SeverityDebug`: For debug-level information
|
||||||
|
|
||||||
|
```
|
||||||
+33
@@ -0,0 +1,33 @@
|
|||||||
|
package errortracking
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NewProviderFromConfig creates an error tracking provider based on the configuration
|
||||||
|
func NewProviderFromConfig(cfg config.ErrorTrackingConfig) (Provider, error) {
|
||||||
|
if !cfg.Enabled {
|
||||||
|
return NewNoOpProvider(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
switch cfg.Provider {
|
||||||
|
case "sentry":
|
||||||
|
if cfg.DSN == "" {
|
||||||
|
return nil, fmt.Errorf("sentry DSN is required when error tracking is enabled")
|
||||||
|
}
|
||||||
|
return NewSentryProvider(SentryConfig{
|
||||||
|
DSN: cfg.DSN,
|
||||||
|
Environment: cfg.Environment,
|
||||||
|
Release: cfg.Release,
|
||||||
|
Debug: cfg.Debug,
|
||||||
|
SampleRate: cfg.SampleRate,
|
||||||
|
TracesSampleRate: cfg.TracesSampleRate,
|
||||||
|
})
|
||||||
|
case "noop", "":
|
||||||
|
return NewNoOpProvider(), nil
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("unknown error tracking provider: %s", cfg.Provider)
|
||||||
|
}
|
||||||
|
}
|
||||||
+33
@@ -0,0 +1,33 @@
|
|||||||
|
package errortracking
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Severity represents the severity level of an error
|
||||||
|
type Severity string
|
||||||
|
|
||||||
|
const (
|
||||||
|
SeverityError Severity = "error"
|
||||||
|
SeverityWarning Severity = "warning"
|
||||||
|
SeverityInfo Severity = "info"
|
||||||
|
SeverityDebug Severity = "debug"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Provider defines the interface for error tracking providers
|
||||||
|
type Provider interface {
|
||||||
|
// CaptureError captures an error with the given severity and additional context
|
||||||
|
CaptureError(ctx context.Context, err error, severity Severity, extra map[string]interface{})
|
||||||
|
|
||||||
|
// CaptureMessage captures a message with the given severity and additional context
|
||||||
|
CaptureMessage(ctx context.Context, message string, severity Severity, extra map[string]interface{})
|
||||||
|
|
||||||
|
// CapturePanic captures a panic with stack trace
|
||||||
|
CapturePanic(ctx context.Context, recovered interface{}, stackTrace []byte, extra map[string]interface{})
|
||||||
|
|
||||||
|
// Flush waits for all events to be sent (useful for graceful shutdown)
|
||||||
|
Flush(timeout int) bool
|
||||||
|
|
||||||
|
// Close closes the provider and releases resources
|
||||||
|
Close() error
|
||||||
|
}
|
||||||
+37
@@ -0,0 +1,37 @@
|
|||||||
|
package errortracking
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
// NoOpProvider is a no-op implementation of the Provider interface
|
||||||
|
// Used when error tracking is disabled
|
||||||
|
type NoOpProvider struct{}
|
||||||
|
|
||||||
|
// NewNoOpProvider creates a new NoOp provider
|
||||||
|
func NewNoOpProvider() *NoOpProvider {
|
||||||
|
return &NoOpProvider{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CaptureError does nothing
|
||||||
|
func (n *NoOpProvider) CaptureError(ctx context.Context, err error, severity Severity, extra map[string]interface{}) {
|
||||||
|
// No-op
|
||||||
|
}
|
||||||
|
|
||||||
|
// CaptureMessage does nothing
|
||||||
|
func (n *NoOpProvider) CaptureMessage(ctx context.Context, message string, severity Severity, extra map[string]interface{}) {
|
||||||
|
// No-op
|
||||||
|
}
|
||||||
|
|
||||||
|
// CapturePanic does nothing
|
||||||
|
func (n *NoOpProvider) CapturePanic(ctx context.Context, recovered interface{}, stackTrace []byte, extra map[string]interface{}) {
|
||||||
|
// No-op
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush does nothing and returns true
|
||||||
|
func (n *NoOpProvider) Flush(timeout int) bool {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close does nothing
|
||||||
|
func (n *NoOpProvider) Close() error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+157
@@ -0,0 +1,157 @@
|
|||||||
|
package errortracking
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/getsentry/sentry-go"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SentryProvider implements the Provider interface using Sentry
|
||||||
|
type SentryProvider struct {
|
||||||
|
hub *sentry.Hub
|
||||||
|
}
|
||||||
|
|
||||||
|
// SentryConfig holds the configuration for Sentry
|
||||||
|
type SentryConfig struct {
|
||||||
|
DSN string
|
||||||
|
Environment string
|
||||||
|
Release string
|
||||||
|
Debug bool
|
||||||
|
SampleRate float64
|
||||||
|
TracesSampleRate float64
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSentryProvider creates a new Sentry provider
|
||||||
|
func NewSentryProvider(config SentryConfig) (*SentryProvider, error) {
|
||||||
|
err := sentry.Init(sentry.ClientOptions{
|
||||||
|
Dsn: config.DSN,
|
||||||
|
Environment: config.Environment,
|
||||||
|
Release: config.Release,
|
||||||
|
Debug: config.Debug,
|
||||||
|
AttachStacktrace: true,
|
||||||
|
SampleRate: config.SampleRate,
|
||||||
|
TracesSampleRate: config.TracesSampleRate,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to initialize Sentry: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &SentryProvider{
|
||||||
|
hub: sentry.CurrentHub(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CaptureError captures an error with the given severity and additional context
|
||||||
|
func (s *SentryProvider) CaptureError(ctx context.Context, err error, severity Severity, extra map[string]interface{}) {
|
||||||
|
if err == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hub := sentry.GetHubFromContext(ctx)
|
||||||
|
if hub == nil {
|
||||||
|
hub = s.hub
|
||||||
|
}
|
||||||
|
|
||||||
|
event := sentry.NewEvent()
|
||||||
|
event.Level = s.convertSeverity(severity)
|
||||||
|
event.Message = err.Error()
|
||||||
|
event.Exception = []sentry.Exception{
|
||||||
|
{
|
||||||
|
Value: err.Error(),
|
||||||
|
Type: fmt.Sprintf("%T", err),
|
||||||
|
Stacktrace: sentry.ExtractStacktrace(err),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
if extra != nil {
|
||||||
|
event.Contexts["extra"] = sentry.Context(extra)
|
||||||
|
}
|
||||||
|
|
||||||
|
hub.CaptureEvent(event)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CaptureMessage captures a message with the given severity and additional context
|
||||||
|
func (s *SentryProvider) CaptureMessage(ctx context.Context, message string, severity Severity, extra map[string]interface{}) {
|
||||||
|
if message == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hub := sentry.GetHubFromContext(ctx)
|
||||||
|
if hub == nil {
|
||||||
|
hub = s.hub
|
||||||
|
}
|
||||||
|
|
||||||
|
event := sentry.NewEvent()
|
||||||
|
event.Level = s.convertSeverity(severity)
|
||||||
|
event.Message = message
|
||||||
|
|
||||||
|
if extra != nil {
|
||||||
|
event.Contexts["extra"] = sentry.Context(extra)
|
||||||
|
}
|
||||||
|
|
||||||
|
hub.CaptureEvent(event)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CapturePanic captures a panic with stack trace
|
||||||
|
func (s *SentryProvider) CapturePanic(ctx context.Context, recovered interface{}, stackTrace []byte, extra map[string]interface{}) {
|
||||||
|
if recovered == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hub := sentry.GetHubFromContext(ctx)
|
||||||
|
if hub == nil {
|
||||||
|
hub = s.hub
|
||||||
|
}
|
||||||
|
|
||||||
|
event := sentry.NewEvent()
|
||||||
|
event.Level = sentry.LevelError
|
||||||
|
event.Message = fmt.Sprintf("Panic: %v", recovered)
|
||||||
|
event.Exception = []sentry.Exception{
|
||||||
|
{
|
||||||
|
Value: fmt.Sprintf("%v", recovered),
|
||||||
|
Type: "panic",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
extraCtx := sentry.Context{}
|
||||||
|
for k, v := range extra {
|
||||||
|
extraCtx[k] = v
|
||||||
|
}
|
||||||
|
if stackTrace != nil {
|
||||||
|
extraCtx["stack_trace"] = string(stackTrace)
|
||||||
|
}
|
||||||
|
if len(extraCtx) > 0 {
|
||||||
|
event.Contexts["extra"] = extraCtx
|
||||||
|
}
|
||||||
|
|
||||||
|
hub.CaptureEvent(event)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush waits for all events to be sent (useful for graceful shutdown)
|
||||||
|
func (s *SentryProvider) Flush(timeout int) bool {
|
||||||
|
return sentry.Flush(time.Duration(timeout) * time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close closes the provider and releases resources
|
||||||
|
func (s *SentryProvider) Close() error {
|
||||||
|
sentry.Flush(2 * time.Second)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// convertSeverity converts our Severity to Sentry's Level
|
||||||
|
func (s *SentryProvider) convertSeverity(severity Severity) sentry.Level {
|
||||||
|
switch severity {
|
||||||
|
case SeverityError:
|
||||||
|
return sentry.LevelError
|
||||||
|
case SeverityWarning:
|
||||||
|
return sentry.LevelWarning
|
||||||
|
case SeverityInfo:
|
||||||
|
return sentry.LevelInfo
|
||||||
|
case SeverityDebug:
|
||||||
|
return sentry.LevelDebug
|
||||||
|
default:
|
||||||
|
return sentry.LevelError
|
||||||
|
}
|
||||||
|
}
|
||||||
+211
@@ -0,0 +1,211 @@
|
|||||||
|
package logger
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"os"
|
||||||
|
"runtime/debug"
|
||||||
|
|
||||||
|
"go.uber.org/zap"
|
||||||
|
|
||||||
|
errortracking "github.com/bitechdev/ResolveSpec/pkg/errortracking"
|
||||||
|
)
|
||||||
|
|
||||||
|
var Logger *zap.SugaredLogger
|
||||||
|
var errorTracker errortracking.Provider
|
||||||
|
|
||||||
|
func Init(dev bool) {
|
||||||
|
|
||||||
|
if dev {
|
||||||
|
cfg := zap.NewDevelopmentConfig()
|
||||||
|
UpdateLogger(&cfg)
|
||||||
|
} else {
|
||||||
|
cfg := zap.NewProductionConfig()
|
||||||
|
UpdateLogger(&cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
func UpdateLoggerPath(path string, dev bool) {
|
||||||
|
defaultConfig := zap.NewProductionConfig()
|
||||||
|
if dev {
|
||||||
|
defaultConfig = zap.NewDevelopmentConfig()
|
||||||
|
}
|
||||||
|
defaultConfig.OutputPaths = []string{path}
|
||||||
|
UpdateLogger(&defaultConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
func UpdateLogger(config *zap.Config) {
|
||||||
|
defaultConfig := zap.NewProductionConfig()
|
||||||
|
defaultConfig.OutputPaths = []string{"resolvespec.log"}
|
||||||
|
if config == nil {
|
||||||
|
config = &defaultConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
logger, err := config.Build()
|
||||||
|
if err != nil {
|
||||||
|
log.Print(err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
Logger = logger.Sugar()
|
||||||
|
Info("ResolveSpec Logger initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
// InitErrorTracking initializes the error tracking provider
|
||||||
|
func InitErrorTracking(provider errortracking.Provider) {
|
||||||
|
errorTracker = provider
|
||||||
|
if errorTracker != nil {
|
||||||
|
Info("Error tracking initialized")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetErrorTracker returns the current error tracking provider
|
||||||
|
func GetErrorTracker() errortracking.Provider {
|
||||||
|
return errorTracker
|
||||||
|
}
|
||||||
|
|
||||||
|
// CloseErrorTracking flushes and closes the error tracking provider
|
||||||
|
func CloseErrorTracking() error {
|
||||||
|
if errorTracker != nil {
|
||||||
|
errorTracker.Flush(5)
|
||||||
|
return errorTracker.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractContext attempts to find a context.Context in the given arguments.
|
||||||
|
// It returns the found context (or context.Background() if not found) and
|
||||||
|
// the remaining arguments without the context.
|
||||||
|
func extractContext(args ...interface{}) (ctx context.Context, filteredArgs []interface{}) {
|
||||||
|
ctx = context.Background()
|
||||||
|
var newArgs []interface{}
|
||||||
|
found := false
|
||||||
|
|
||||||
|
for _, arg := range args {
|
||||||
|
if c, ok := arg.(context.Context); ok {
|
||||||
|
if !found {
|
||||||
|
ctx = c
|
||||||
|
found = true
|
||||||
|
}
|
||||||
|
// Ignore any additional context.Context arguments after the first one.
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
newArgs = append(newArgs, arg)
|
||||||
|
}
|
||||||
|
return ctx, newArgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func Info(template string, args ...interface{}) {
|
||||||
|
if Logger == nil {
|
||||||
|
log.Printf(template, args...)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
Logger.Infow(fmt.Sprintf(template, args...), "process_id", os.Getpid())
|
||||||
|
}
|
||||||
|
|
||||||
|
func Warn(template string, args ...interface{}) {
|
||||||
|
ctx, remainingArgs := extractContext(args...)
|
||||||
|
message := fmt.Sprintf(template, remainingArgs...)
|
||||||
|
if Logger == nil {
|
||||||
|
log.Printf("%s", message)
|
||||||
|
} else {
|
||||||
|
Logger.Warnw(message, "process_id", os.Getpid())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send to error tracker
|
||||||
|
if errorTracker != nil {
|
||||||
|
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityWarning, map[string]interface{}{
|
||||||
|
"process_id": os.Getpid(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func Error(template string, args ...interface{}) {
|
||||||
|
ctx, remainingArgs := extractContext(args...)
|
||||||
|
message := fmt.Sprintf(template, remainingArgs...)
|
||||||
|
if Logger == nil {
|
||||||
|
log.Printf("%s", message)
|
||||||
|
} else {
|
||||||
|
Logger.Errorw(message, "process_id", os.Getpid())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send to error tracker
|
||||||
|
if errorTracker != nil {
|
||||||
|
errorTracker.CaptureMessage(ctx, message, errortracking.SeverityError, map[string]interface{}{
|
||||||
|
"process_id": os.Getpid(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func Debug(template string, args ...interface{}) {
|
||||||
|
if Logger == nil {
|
||||||
|
log.Printf(template, args...)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
Logger.Debugw(fmt.Sprintf(template, args...), "process_id", os.Getpid())
|
||||||
|
}
|
||||||
|
|
||||||
|
// CatchPanic - Handle panic
|
||||||
|
// Returns a function that should be deferred to catch panics
|
||||||
|
// Example usage: defer CatchPanicCallback("MyFunction", func(err any) { /* cleanup */ })()
|
||||||
|
func CatchPanicCallback(location string, cb func(err any), args ...interface{}) func() {
|
||||||
|
ctx, _ := extractContext(args...)
|
||||||
|
return func() {
|
||||||
|
if err := recover(); err != nil {
|
||||||
|
callstack := debug.Stack()
|
||||||
|
|
||||||
|
if Logger != nil {
|
||||||
|
Error("Panic in %s : %v", location, err, ctx) // Pass context implicitly
|
||||||
|
} else {
|
||||||
|
fmt.Printf("%s:PANIC->%+v", location, err)
|
||||||
|
debug.PrintStack()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send to error tracker
|
||||||
|
if errorTracker != nil {
|
||||||
|
errorTracker.CapturePanic(ctx, err, callstack, map[string]interface{}{
|
||||||
|
"location": location,
|
||||||
|
"process_id": os.Getpid(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if cb != nil {
|
||||||
|
cb(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CatchPanic - Handle panic
|
||||||
|
// Returns a function that should be deferred to catch panics
|
||||||
|
// Example usage: defer CatchPanic("MyFunction")()
|
||||||
|
func CatchPanic(location string, args ...interface{}) func() {
|
||||||
|
return CatchPanicCallback(location, nil, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// HandlePanic logs a panic and returns it as an error
|
||||||
|
// This should be called with the result of recover() from a deferred function
|
||||||
|
// Example usage:
|
||||||
|
//
|
||||||
|
// defer func() {
|
||||||
|
// if r := recover(); r != nil {
|
||||||
|
// err = logger.HandlePanic("MethodName", r)
|
||||||
|
// }
|
||||||
|
// }()
|
||||||
|
func HandlePanic(methodName string, r any, args ...interface{}) error {
|
||||||
|
ctx, _ := extractContext(args...)
|
||||||
|
stack := debug.Stack()
|
||||||
|
Error("Panic in %s: %v\nStack trace:\n%s", methodName, r, string(stack), ctx) // Pass context implicitly
|
||||||
|
|
||||||
|
// Send to error tracker
|
||||||
|
if errorTracker != nil {
|
||||||
|
errorTracker.CapturePanic(ctx, r, stack, map[string]interface{}{
|
||||||
|
"method": methodName,
|
||||||
|
"process_id": os.Getpid(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("panic in %s: %v", methodName, r)
|
||||||
|
}
|
||||||
+478
@@ -0,0 +1,478 @@
|
|||||||
|
# Metrics Package
|
||||||
|
|
||||||
|
A pluggable metrics collection system with Prometheus implementation.
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
|
||||||
|
// Initialize Prometheus provider with default config
|
||||||
|
provider := metrics.NewPrometheusProvider(nil)
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
|
||||||
|
// Apply middleware to your router
|
||||||
|
router.Use(provider.Middleware)
|
||||||
|
|
||||||
|
// Expose metrics endpoint
|
||||||
|
http.Handle("/metrics", provider.Handler())
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
You can customize the metrics provider using a configuration struct:
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
|
||||||
|
// Create custom configuration
|
||||||
|
config := &metrics.Config{
|
||||||
|
Enabled: true,
|
||||||
|
Provider: "prometheus",
|
||||||
|
Namespace: "myapp", // Prefix all metrics with "myapp_"
|
||||||
|
HTTPRequestBuckets: []float64{0.01, 0.05, 0.1, 0.5, 1, 2, 5},
|
||||||
|
DBQueryBuckets: []float64{0.001, 0.01, 0.05, 0.1, 0.5, 1},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize with custom config
|
||||||
|
provider := metrics.NewPrometheusProvider(config)
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Configuration Options
|
||||||
|
|
||||||
|
| Field | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `Enabled` | `bool` | `true` | Enable/disable metrics collection |
|
||||||
|
| `Provider` | `string` | `"prometheus"` | Metrics provider type |
|
||||||
|
| `Namespace` | `string` | `""` | Prefix for all metric names |
|
||||||
|
| `HTTPRequestBuckets` | `[]float64` | See below | Histogram buckets for HTTP duration (seconds) |
|
||||||
|
| `DBQueryBuckets` | `[]float64` | See below | Histogram buckets for DB query duration (seconds) |
|
||||||
|
|
||||||
|
**Default HTTP Request Buckets:** `[0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10]`
|
||||||
|
|
||||||
|
**Default DB Query Buckets:** `[0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5]`
|
||||||
|
|
||||||
|
### Pushgateway Configuration (Optional)
|
||||||
|
|
||||||
|
For batch jobs, cron tasks, or short-lived processes, you can push metrics to Prometheus Pushgateway:
|
||||||
|
|
||||||
|
| Field | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `PushgatewayURL` | `string` | `""` | URL of Pushgateway (e.g., "http://pushgateway:9091") |
|
||||||
|
| `PushgatewayJobName` | `string` | `"resolvespec"` | Job name for pushed metrics |
|
||||||
|
| `PushgatewayInterval` | `int` | `0` | Auto-push interval in seconds (0 = disabled) |
|
||||||
|
|
||||||
|
```go
|
||||||
|
config := &metrics.Config{
|
||||||
|
PushgatewayURL: "http://pushgateway:9091",
|
||||||
|
PushgatewayJobName: "batch-job",
|
||||||
|
PushgatewayInterval: 30, // Push every 30 seconds
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Provider Interface
|
||||||
|
|
||||||
|
The package uses a provider interface, allowing you to plug in different metric systems:
|
||||||
|
|
||||||
|
```go
|
||||||
|
type Provider interface {
|
||||||
|
RecordHTTPRequest(method, path, status string, duration time.Duration)
|
||||||
|
IncRequestsInFlight()
|
||||||
|
DecRequestsInFlight()
|
||||||
|
RecordDBQuery(operation, table string, duration time.Duration, err error)
|
||||||
|
RecordCacheHit(provider string)
|
||||||
|
RecordCacheMiss(provider string)
|
||||||
|
UpdateCacheSize(provider string, size int64)
|
||||||
|
Handler() http.Handler
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Recording Metrics
|
||||||
|
|
||||||
|
### HTTP Metrics (Automatic)
|
||||||
|
|
||||||
|
When using the middleware, HTTP metrics are recorded automatically:
|
||||||
|
|
||||||
|
```go
|
||||||
|
router.Use(provider.Middleware)
|
||||||
|
```
|
||||||
|
|
||||||
|
**Collected:**
|
||||||
|
- Request duration (histogram)
|
||||||
|
- Request count by method, path, and status
|
||||||
|
- Requests in flight (gauge)
|
||||||
|
|
||||||
|
### Database Metrics
|
||||||
|
|
||||||
|
```go
|
||||||
|
start := time.Now()
|
||||||
|
rows, err := db.Query("SELECT * FROM users WHERE id = ?", userID)
|
||||||
|
duration := time.Since(start)
|
||||||
|
|
||||||
|
metrics.GetProvider().RecordDBQuery("SELECT", "users", duration, err)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Cache Metrics
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Record cache hit
|
||||||
|
metrics.GetProvider().RecordCacheHit("memory")
|
||||||
|
|
||||||
|
// Record cache miss
|
||||||
|
metrics.GetProvider().RecordCacheMiss("memory")
|
||||||
|
|
||||||
|
// Update cache size
|
||||||
|
metrics.GetProvider().UpdateCacheSize("memory", 1024)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Prometheus Metrics
|
||||||
|
|
||||||
|
When using `PrometheusProvider`, the following metrics are available:
|
||||||
|
|
||||||
|
| Metric Name | Type | Labels | Description |
|
||||||
|
|-------------|------|--------|-------------|
|
||||||
|
| `http_request_duration_seconds` | Histogram | method, path, status | HTTP request duration |
|
||||||
|
| `http_requests_total` | Counter | method, path, status | Total HTTP requests |
|
||||||
|
| `http_requests_in_flight` | Gauge | - | Current in-flight requests |
|
||||||
|
| `db_query_duration_seconds` | Histogram | operation, table | Database query duration |
|
||||||
|
| `db_queries_total` | Counter | operation, table, status | Total database queries |
|
||||||
|
| `cache_hits_total` | Counter | provider | Total cache hits |
|
||||||
|
| `cache_misses_total` | Counter | provider | Total cache misses |
|
||||||
|
| `cache_size_items` | Gauge | provider | Current cache size |
|
||||||
|
| `events_published_total` | Counter | source, event_type | Total events published |
|
||||||
|
| `events_processed_total` | Counter | source, event_type, status | Total events processed |
|
||||||
|
| `event_processing_duration_seconds` | Histogram | source, event_type | Event processing duration |
|
||||||
|
| `event_queue_size` | Gauge | - | Current event queue size |
|
||||||
|
| `panics_total` | Counter | method | Total panics recovered |
|
||||||
|
|
||||||
|
**Note:** If a custom `Namespace` is configured, all metric names will be prefixed with `{namespace}_`.
|
||||||
|
|
||||||
|
## Prometheus Queries
|
||||||
|
|
||||||
|
### HTTP Request Rate
|
||||||
|
|
||||||
|
```promql
|
||||||
|
rate(http_requests_total[5m])
|
||||||
|
```
|
||||||
|
|
||||||
|
### HTTP Request Duration (95th percentile)
|
||||||
|
|
||||||
|
```promql
|
||||||
|
histogram_quantile(0.95, rate(http_request_duration_seconds_bucket[5m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
### Database Query Error Rate
|
||||||
|
|
||||||
|
```promql
|
||||||
|
rate(db_queries_total{status="error"}[5m])
|
||||||
|
```
|
||||||
|
|
||||||
|
### Cache Hit Rate
|
||||||
|
|
||||||
|
```promql
|
||||||
|
rate(cache_hits_total[5m]) / (rate(cache_hits_total[5m]) + rate(cache_misses_total[5m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
## No-Op Provider
|
||||||
|
|
||||||
|
If metrics are disabled:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// No provider set - uses no-op provider automatically
|
||||||
|
metrics.GetProvider().RecordHTTPRequest(...) // Does nothing
|
||||||
|
```
|
||||||
|
|
||||||
|
## Custom Provider
|
||||||
|
|
||||||
|
Implement your own metrics provider:
|
||||||
|
|
||||||
|
```go
|
||||||
|
type CustomProvider struct{}
|
||||||
|
|
||||||
|
func (c *CustomProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {
|
||||||
|
// Send to your metrics system
|
||||||
|
}
|
||||||
|
|
||||||
|
// Implement other Provider interface methods...
|
||||||
|
|
||||||
|
func (c *CustomProvider) Handler() http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Return your metrics format
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use it
|
||||||
|
metrics.SetProvider(&CustomProvider{})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Pushgateway Usage
|
||||||
|
|
||||||
|
### Automatic Push (Batch Jobs)
|
||||||
|
|
||||||
|
For jobs that run periodically, use automatic pushing:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Configure with automatic pushing every 30 seconds
|
||||||
|
config := &metrics.Config{
|
||||||
|
Enabled: true,
|
||||||
|
Provider: "prometheus",
|
||||||
|
Namespace: "batch_job",
|
||||||
|
PushgatewayURL: "http://pushgateway:9091",
|
||||||
|
PushgatewayJobName: "data-processor",
|
||||||
|
PushgatewayInterval: 30, // Push every 30 seconds
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := metrics.NewPrometheusProvider(config)
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
|
||||||
|
// Ensure cleanup on exit
|
||||||
|
defer provider.StopAutoPush()
|
||||||
|
|
||||||
|
// Your batch job logic here
|
||||||
|
processBatchData()
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Manual Push (Short-lived Processes)
|
||||||
|
|
||||||
|
For one-time jobs or when you want manual control:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Configure without automatic pushing
|
||||||
|
config := &metrics.Config{
|
||||||
|
Enabled: true,
|
||||||
|
Provider: "prometheus",
|
||||||
|
PushgatewayURL: "http://pushgateway:9091",
|
||||||
|
PushgatewayJobName: "migration-job",
|
||||||
|
// PushgatewayInterval: 0 (default - no auto-push)
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := metrics.NewPrometheusProvider(config)
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
|
||||||
|
// Run your job
|
||||||
|
err := runMigration()
|
||||||
|
|
||||||
|
// Push metrics at the end
|
||||||
|
if pushErr := provider.Push(); pushErr != nil {
|
||||||
|
log.Printf("Failed to push metrics: %v", pushErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Docker Compose with Pushgateway
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
version: '3'
|
||||||
|
services:
|
||||||
|
batch-job:
|
||||||
|
build: .
|
||||||
|
environment:
|
||||||
|
PUSHGATEWAY_URL: "http://pushgateway:9091"
|
||||||
|
|
||||||
|
pushgateway:
|
||||||
|
image: prom/pushgateway
|
||||||
|
ports:
|
||||||
|
- "9091:9091"
|
||||||
|
|
||||||
|
prometheus:
|
||||||
|
image: prom/prometheus
|
||||||
|
ports:
|
||||||
|
- "9090:9090"
|
||||||
|
volumes:
|
||||||
|
- ./prometheus.yml:/etc/prometheus/prometheus.yml
|
||||||
|
command:
|
||||||
|
- '--config.file=/etc/prometheus/prometheus.yml'
|
||||||
|
```
|
||||||
|
|
||||||
|
**prometheus.yml for Pushgateway:**
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
global:
|
||||||
|
scrape_interval: 15s
|
||||||
|
|
||||||
|
scrape_configs:
|
||||||
|
# Scrape the pushgateway
|
||||||
|
- job_name: 'pushgateway'
|
||||||
|
honor_labels: true # Important: preserve job labels from pushed metrics
|
||||||
|
static_configs:
|
||||||
|
- targets: ['pushgateway:9091']
|
||||||
|
```
|
||||||
|
|
||||||
|
## Complete Example
|
||||||
|
|
||||||
|
### Basic Usage
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Initialize metrics with default config
|
||||||
|
provider := metrics.NewPrometheusProvider(nil)
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
|
||||||
|
// Create router
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Apply metrics middleware
|
||||||
|
router.Use(provider.Middleware)
|
||||||
|
|
||||||
|
// Expose metrics endpoint
|
||||||
|
router.Handle("/metrics", provider.Handler())
|
||||||
|
|
||||||
|
// Your API routes
|
||||||
|
router.HandleFunc("/api/users", getUsersHandler)
|
||||||
|
|
||||||
|
log.Fatal(http.ListenAndServe(":8080", router))
|
||||||
|
}
|
||||||
|
|
||||||
|
func getUsersHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Record database query
|
||||||
|
start := time.Now()
|
||||||
|
users, err := fetchUsers()
|
||||||
|
duration := time.Since(start)
|
||||||
|
|
||||||
|
metrics.GetProvider().RecordDBQuery("SELECT", "users", duration, err)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Internal Server Error", 500)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return users...
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### With Custom Configuration
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Custom metrics configuration
|
||||||
|
metricsConfig := &metrics.Config{
|
||||||
|
Enabled: true,
|
||||||
|
Provider: "prometheus",
|
||||||
|
Namespace: "myapp",
|
||||||
|
// Custom buckets optimized for your application
|
||||||
|
HTTPRequestBuckets: []float64{0.01, 0.05, 0.1, 0.5, 1, 2, 5, 10},
|
||||||
|
DBQueryBuckets: []float64{0.001, 0.005, 0.01, 0.05, 0.1, 0.5, 1},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize with custom config
|
||||||
|
provider := metrics.NewPrometheusProvider(metricsConfig)
|
||||||
|
metrics.SetProvider(provider)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
router.Use(provider.Middleware)
|
||||||
|
router.Handle("/metrics", provider.Handler())
|
||||||
|
|
||||||
|
log.Fatal(http.ListenAndServe(":8080", router))
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Docker Compose Example
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
version: '3'
|
||||||
|
services:
|
||||||
|
app:
|
||||||
|
build: .
|
||||||
|
ports:
|
||||||
|
- "8080:8080"
|
||||||
|
|
||||||
|
prometheus:
|
||||||
|
image: prom/prometheus
|
||||||
|
ports:
|
||||||
|
- "9090:9090"
|
||||||
|
volumes:
|
||||||
|
- ./prometheus.yml:/etc/prometheus/prometheus.yml
|
||||||
|
command:
|
||||||
|
- '--config.file=/etc/prometheus/prometheus.yml'
|
||||||
|
|
||||||
|
grafana:
|
||||||
|
image: grafana/grafana
|
||||||
|
ports:
|
||||||
|
- "3000:3000"
|
||||||
|
depends_on:
|
||||||
|
- prometheus
|
||||||
|
```
|
||||||
|
|
||||||
|
**prometheus.yml:**
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
global:
|
||||||
|
scrape_interval: 15s
|
||||||
|
|
||||||
|
scrape_configs:
|
||||||
|
- job_name: 'resolvespec'
|
||||||
|
static_configs:
|
||||||
|
- targets: ['app:8080']
|
||||||
|
```
|
||||||
|
|
||||||
|
## Best Practices
|
||||||
|
|
||||||
|
1. **Label Cardinality**: Keep labels low-cardinality
|
||||||
|
- ✅ Good: `method`, `status_code`
|
||||||
|
- ❌ Bad: `user_id`, `timestamp`
|
||||||
|
|
||||||
|
2. **Path Normalization**: Normalize dynamic paths
|
||||||
|
```go
|
||||||
|
// Instead of /api/users/123
|
||||||
|
// Use /api/users/:id
|
||||||
|
```
|
||||||
|
|
||||||
|
3. **Metric Naming**: Follow Prometheus conventions
|
||||||
|
- Use `_total` suffix for counters
|
||||||
|
- Use `_seconds` suffix for durations
|
||||||
|
- Use base units (seconds, not milliseconds)
|
||||||
|
|
||||||
|
4. **Performance**: Metrics collection is lock-free and highly performant
|
||||||
|
- Safe for high-throughput applications
|
||||||
|
- Minimal overhead (<1% in most cases)
|
||||||
|
|
||||||
|
5. **Pull vs Push**:
|
||||||
|
- **Use Pull (default)**: Long-running services, web servers, microservices
|
||||||
|
- **Use Push (Pushgateway)**: Batch jobs, cron tasks, short-lived processes, serverless functions
|
||||||
|
- Pull is preferred for most applications as it allows Prometheus to detect if your service is down
|
||||||
+64
@@ -0,0 +1,64 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
// Config holds configuration for the metrics provider
|
||||||
|
type Config struct {
|
||||||
|
// Enabled determines whether metrics collection is enabled
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
|
||||||
|
// Provider specifies which metrics provider to use (prometheus, noop)
|
||||||
|
Provider string `mapstructure:"provider"`
|
||||||
|
|
||||||
|
// Namespace is an optional prefix for all metric names
|
||||||
|
Namespace string `mapstructure:"namespace"`
|
||||||
|
|
||||||
|
// HTTPRequestBuckets defines histogram buckets for HTTP request duration (in seconds)
|
||||||
|
// Default: [0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10]
|
||||||
|
HTTPRequestBuckets []float64 `mapstructure:"http_request_buckets"`
|
||||||
|
|
||||||
|
// DBQueryBuckets defines histogram buckets for database query duration (in seconds)
|
||||||
|
// Default: [0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5]
|
||||||
|
DBQueryBuckets []float64 `mapstructure:"db_query_buckets"`
|
||||||
|
|
||||||
|
// PushgatewayURL is the URL of the Prometheus Pushgateway (optional)
|
||||||
|
// If set, metrics will be pushed to this gateway instead of only being scraped
|
||||||
|
// Example: "http://pushgateway:9091"
|
||||||
|
PushgatewayURL string `mapstructure:"pushgateway_url"`
|
||||||
|
|
||||||
|
// PushgatewayJobName is the job name to use when pushing metrics to Pushgateway
|
||||||
|
// Default: "resolvespec"
|
||||||
|
PushgatewayJobName string `mapstructure:"pushgateway_job_name"`
|
||||||
|
|
||||||
|
// PushgatewayInterval is the interval at which to push metrics to Pushgateway
|
||||||
|
// Only used if PushgatewayURL is set. If 0, automatic pushing is disabled.
|
||||||
|
// Default: 0 (no automatic pushing)
|
||||||
|
PushgatewayInterval int `mapstructure:"pushgateway_interval"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultConfig returns a Config with sensible defaults
|
||||||
|
func DefaultConfig() *Config {
|
||||||
|
return &Config{
|
||||||
|
Enabled: true,
|
||||||
|
Provider: "prometheus",
|
||||||
|
// HTTP requests typically take longer than DB queries
|
||||||
|
HTTPRequestBuckets: []float64{0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10},
|
||||||
|
// DB queries are usually faster
|
||||||
|
DBQueryBuckets: []float64{0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplyDefaults fills in any missing values with defaults
|
||||||
|
func (c *Config) ApplyDefaults() {
|
||||||
|
if c.Provider == "" {
|
||||||
|
c.Provider = "prometheus"
|
||||||
|
}
|
||||||
|
if len(c.HTTPRequestBuckets) == 0 {
|
||||||
|
c.HTTPRequestBuckets = []float64{0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10}
|
||||||
|
}
|
||||||
|
if len(c.DBQueryBuckets) == 0 {
|
||||||
|
c.DBQueryBuckets = []float64{0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5}
|
||||||
|
}
|
||||||
|
// Set default job name if pushgateway is configured but job name is empty
|
||||||
|
if c.PushgatewayURL != "" && c.PushgatewayJobName == "" {
|
||||||
|
c.PushgatewayJobName = "resolvespec"
|
||||||
|
}
|
||||||
|
}
|
||||||
+98
@@ -0,0 +1,98 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Provider defines the interface for metric collection
|
||||||
|
type Provider interface {
|
||||||
|
// RecordHTTPRequest records metrics for an HTTP request
|
||||||
|
RecordHTTPRequest(method, path, status string, duration time.Duration)
|
||||||
|
|
||||||
|
// IncRequestsInFlight increments the in-flight requests counter
|
||||||
|
IncRequestsInFlight()
|
||||||
|
|
||||||
|
// DecRequestsInFlight decrements the in-flight requests counter
|
||||||
|
DecRequestsInFlight()
|
||||||
|
|
||||||
|
// RecordDBQuery records metrics for a database query
|
||||||
|
RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error)
|
||||||
|
|
||||||
|
// RecordCacheHit records a cache hit
|
||||||
|
RecordCacheHit(provider string)
|
||||||
|
|
||||||
|
// RecordCacheMiss records a cache miss
|
||||||
|
RecordCacheMiss(provider string)
|
||||||
|
|
||||||
|
// UpdateCacheSize updates the cache size metric
|
||||||
|
UpdateCacheSize(provider string, size int64)
|
||||||
|
|
||||||
|
// RecordEventPublished records an event publication
|
||||||
|
RecordEventPublished(source, eventType string)
|
||||||
|
|
||||||
|
// RecordEventProcessed records an event processing with its status
|
||||||
|
RecordEventProcessed(source, eventType, status string, duration time.Duration)
|
||||||
|
|
||||||
|
// UpdateEventQueueSize updates the event queue size metric
|
||||||
|
UpdateEventQueueSize(size int64)
|
||||||
|
|
||||||
|
// RecordPanic records a panic event
|
||||||
|
RecordPanic(methodName string)
|
||||||
|
|
||||||
|
// Handler returns an HTTP handler for exposing metrics (e.g., /metrics endpoint)
|
||||||
|
Handler() http.Handler
|
||||||
|
}
|
||||||
|
|
||||||
|
// globalProvider is the global metrics provider, protected by globalProviderMu.
|
||||||
|
var (
|
||||||
|
globalProviderMu sync.RWMutex
|
||||||
|
globalProvider Provider
|
||||||
|
)
|
||||||
|
|
||||||
|
// SetProvider sets the global metrics provider.
|
||||||
|
func SetProvider(p Provider) {
|
||||||
|
globalProviderMu.Lock()
|
||||||
|
globalProvider = p
|
||||||
|
globalProviderMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetProvider returns the current metrics provider.
|
||||||
|
func GetProvider() Provider {
|
||||||
|
globalProviderMu.RLock()
|
||||||
|
p := globalProvider
|
||||||
|
globalProviderMu.RUnlock()
|
||||||
|
if p == nil {
|
||||||
|
return &NoOpProvider{}
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// NoOpProvider is a no-op implementation of Provider
|
||||||
|
type NoOpProvider struct{}
|
||||||
|
|
||||||
|
func (n *NoOpProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {}
|
||||||
|
func (n *NoOpProvider) IncRequestsInFlight() {}
|
||||||
|
func (n *NoOpProvider) DecRequestsInFlight() {}
|
||||||
|
func (n *NoOpProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||||
|
}
|
||||||
|
func (n *NoOpProvider) RecordCacheHit(provider string) {}
|
||||||
|
func (n *NoOpProvider) RecordCacheMiss(provider string) {}
|
||||||
|
func (n *NoOpProvider) UpdateCacheSize(provider string, size int64) {}
|
||||||
|
func (n *NoOpProvider) RecordEventPublished(source, eventType string) {}
|
||||||
|
func (n *NoOpProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
|
||||||
|
}
|
||||||
|
func (n *NoOpProvider) UpdateEventQueueSize(size int64) {}
|
||||||
|
func (n *NoOpProvider) RecordPanic(methodName string) {}
|
||||||
|
func (n *NoOpProvider) Handler() http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusNotFound)
|
||||||
|
_, err := w.Write([]byte("Metrics provider not configured"))
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Failed to write. %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+312
@@ -0,0 +1,312 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
|
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||||
|
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||||
|
"github.com/prometheus/client_golang/prometheus/push"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PrometheusProvider implements the Provider interface using Prometheus
|
||||||
|
type PrometheusProvider struct {
|
||||||
|
requestDuration *prometheus.HistogramVec
|
||||||
|
requestTotal *prometheus.CounterVec
|
||||||
|
requestsInFlight prometheus.Gauge
|
||||||
|
dbQueryDuration *prometheus.HistogramVec
|
||||||
|
dbQueryTotal *prometheus.CounterVec
|
||||||
|
cacheHits *prometheus.CounterVec
|
||||||
|
cacheMisses *prometheus.CounterVec
|
||||||
|
cacheSize *prometheus.GaugeVec
|
||||||
|
eventPublished *prometheus.CounterVec
|
||||||
|
eventProcessed *prometheus.CounterVec
|
||||||
|
eventDuration *prometheus.HistogramVec
|
||||||
|
eventQueueSize prometheus.Gauge
|
||||||
|
panicsTotal *prometheus.CounterVec
|
||||||
|
|
||||||
|
// Pushgateway fields (optional)
|
||||||
|
pushgatewayURL string
|
||||||
|
pushgatewayJobName string
|
||||||
|
pusher *push.Pusher
|
||||||
|
pushTicker *time.Ticker
|
||||||
|
pushStop chan bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPrometheusProvider creates a new Prometheus metrics provider
|
||||||
|
// If cfg is nil, default configuration will be used
|
||||||
|
func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
||||||
|
// Use default config if none provided
|
||||||
|
if cfg == nil {
|
||||||
|
cfg = DefaultConfig()
|
||||||
|
} else {
|
||||||
|
// Apply defaults for any missing values
|
||||||
|
cfg.ApplyDefaults()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper to add namespace prefix if configured
|
||||||
|
metricName := func(name string) string {
|
||||||
|
if cfg.Namespace != "" {
|
||||||
|
return cfg.Namespace + "_" + name
|
||||||
|
}
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
|
||||||
|
p := &PrometheusProvider{
|
||||||
|
requestDuration: promauto.NewHistogramVec(
|
||||||
|
prometheus.HistogramOpts{
|
||||||
|
Name: metricName("http_request_duration_seconds"),
|
||||||
|
Help: "HTTP request duration in seconds",
|
||||||
|
Buckets: cfg.HTTPRequestBuckets,
|
||||||
|
},
|
||||||
|
[]string{"method", "path", "status"},
|
||||||
|
),
|
||||||
|
requestTotal: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("http_requests_total"),
|
||||||
|
Help: "Total number of HTTP requests",
|
||||||
|
},
|
||||||
|
[]string{"method", "path", "status"},
|
||||||
|
),
|
||||||
|
|
||||||
|
requestsInFlight: promauto.NewGauge(
|
||||||
|
prometheus.GaugeOpts{
|
||||||
|
Name: metricName("http_requests_in_flight"),
|
||||||
|
Help: "Current number of HTTP requests being processed",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
dbQueryDuration: promauto.NewHistogramVec(
|
||||||
|
prometheus.HistogramOpts{
|
||||||
|
Name: metricName("db_query_duration_seconds"),
|
||||||
|
Help: "Database query duration in seconds",
|
||||||
|
Buckets: cfg.DBQueryBuckets,
|
||||||
|
},
|
||||||
|
[]string{"operation", "schema", "entity", "table"},
|
||||||
|
),
|
||||||
|
dbQueryTotal: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("db_queries_total"),
|
||||||
|
Help: "Total number of database queries",
|
||||||
|
},
|
||||||
|
[]string{"operation", "schema", "entity", "table", "status"},
|
||||||
|
),
|
||||||
|
cacheHits: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("cache_hits_total"),
|
||||||
|
Help: "Total number of cache hits",
|
||||||
|
},
|
||||||
|
[]string{"provider"},
|
||||||
|
),
|
||||||
|
cacheMisses: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("cache_misses_total"),
|
||||||
|
Help: "Total number of cache misses",
|
||||||
|
},
|
||||||
|
[]string{"provider"},
|
||||||
|
),
|
||||||
|
cacheSize: promauto.NewGaugeVec(
|
||||||
|
prometheus.GaugeOpts{
|
||||||
|
Name: metricName("cache_size_items"),
|
||||||
|
Help: "Number of items in cache",
|
||||||
|
},
|
||||||
|
[]string{"provider"},
|
||||||
|
),
|
||||||
|
eventPublished: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("events_published_total"),
|
||||||
|
Help: "Total number of events published",
|
||||||
|
},
|
||||||
|
[]string{"source", "event_type"},
|
||||||
|
),
|
||||||
|
eventProcessed: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("events_processed_total"),
|
||||||
|
Help: "Total number of events processed",
|
||||||
|
},
|
||||||
|
[]string{"source", "event_type", "status"},
|
||||||
|
),
|
||||||
|
eventDuration: promauto.NewHistogramVec(
|
||||||
|
prometheus.HistogramOpts{
|
||||||
|
Name: metricName("event_processing_duration_seconds"),
|
||||||
|
Help: "Event processing duration in seconds",
|
||||||
|
Buckets: cfg.DBQueryBuckets, // Events are typically fast like DB queries
|
||||||
|
},
|
||||||
|
[]string{"source", "event_type"},
|
||||||
|
),
|
||||||
|
eventQueueSize: promauto.NewGauge(
|
||||||
|
prometheus.GaugeOpts{
|
||||||
|
Name: metricName("event_queue_size"),
|
||||||
|
Help: "Current number of events in queue",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
panicsTotal: promauto.NewCounterVec(
|
||||||
|
prometheus.CounterOpts{
|
||||||
|
Name: metricName("panics_total"),
|
||||||
|
Help: "Total number of panics",
|
||||||
|
},
|
||||||
|
[]string{"method"},
|
||||||
|
),
|
||||||
|
|
||||||
|
pushgatewayURL: cfg.PushgatewayURL,
|
||||||
|
pushgatewayJobName: cfg.PushgatewayJobName,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize pushgateway if configured
|
||||||
|
if cfg.PushgatewayURL != "" {
|
||||||
|
p.pusher = push.New(cfg.PushgatewayURL, cfg.PushgatewayJobName).
|
||||||
|
Gatherer(prometheus.DefaultGatherer)
|
||||||
|
|
||||||
|
// Start automatic pushing if interval is configured
|
||||||
|
if cfg.PushgatewayInterval > 0 {
|
||||||
|
p.pushStop = make(chan bool)
|
||||||
|
p.pushTicker = time.NewTicker(time.Duration(cfg.PushgatewayInterval) * time.Second)
|
||||||
|
go p.startAutoPush()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResponseWriter wraps http.ResponseWriter to capture status code
|
||||||
|
type ResponseWriter struct {
|
||||||
|
http.ResponseWriter
|
||||||
|
statusCode int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewResponseWriter(w http.ResponseWriter) *ResponseWriter {
|
||||||
|
return &ResponseWriter{
|
||||||
|
ResponseWriter: w,
|
||||||
|
statusCode: http.StatusOK,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rw *ResponseWriter) WriteHeader(code int) {
|
||||||
|
rw.statusCode = code
|
||||||
|
rw.ResponseWriter.WriteHeader(code)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordHTTPRequest implements Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {
|
||||||
|
p.requestDuration.WithLabelValues(method, path, status).Observe(duration.Seconds())
|
||||||
|
p.requestTotal.WithLabelValues(method, path, status).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// IncRequestsInFlight implements Provider interface
|
||||||
|
func (p *PrometheusProvider) IncRequestsInFlight() {
|
||||||
|
p.requestsInFlight.Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// DecRequestsInFlight implements Provider interface
|
||||||
|
func (p *PrometheusProvider) DecRequestsInFlight() {
|
||||||
|
p.requestsInFlight.Dec()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordDBQuery implements Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||||
|
status := "success"
|
||||||
|
if err != nil {
|
||||||
|
status = "error"
|
||||||
|
}
|
||||||
|
p.dbQueryDuration.WithLabelValues(operation, schema, entity, table).Observe(duration.Seconds())
|
||||||
|
p.dbQueryTotal.WithLabelValues(operation, schema, entity, table, status).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordCacheHit implements Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordCacheHit(provider string) {
|
||||||
|
p.cacheHits.WithLabelValues(provider).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordCacheMiss implements Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordCacheMiss(provider string) {
|
||||||
|
p.cacheMisses.WithLabelValues(provider).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateCacheSize implements Provider interface
|
||||||
|
func (p *PrometheusProvider) UpdateCacheSize(provider string, size int64) {
|
||||||
|
p.cacheSize.WithLabelValues(provider).Set(float64(size))
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordEventPublished implements Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordEventPublished(source, eventType string) {
|
||||||
|
p.eventPublished.WithLabelValues(source, eventType).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordEventProcessed implements Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
|
||||||
|
p.eventProcessed.WithLabelValues(source, eventType, status).Inc()
|
||||||
|
p.eventDuration.WithLabelValues(source, eventType).Observe(duration.Seconds())
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateEventQueueSize implements Provider interface
|
||||||
|
func (p *PrometheusProvider) UpdateEventQueueSize(size int64) {
|
||||||
|
p.eventQueueSize.Set(float64(size))
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordPanic implements the Provider interface
|
||||||
|
func (p *PrometheusProvider) RecordPanic(methodName string) {
|
||||||
|
p.panicsTotal.WithLabelValues(methodName).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handler implements Provider interface
|
||||||
|
func (p *PrometheusProvider) Handler() http.Handler {
|
||||||
|
return promhttp.Handler()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Middleware returns an HTTP middleware that collects metrics
|
||||||
|
func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
start := time.Now()
|
||||||
|
|
||||||
|
// Increment in-flight requests
|
||||||
|
p.IncRequestsInFlight()
|
||||||
|
defer p.DecRequestsInFlight()
|
||||||
|
|
||||||
|
// Wrap response writer to capture status code
|
||||||
|
rw := NewResponseWriter(w)
|
||||||
|
|
||||||
|
// Call next handler
|
||||||
|
next.ServeHTTP(rw, r)
|
||||||
|
|
||||||
|
// Record metrics
|
||||||
|
duration := time.Since(start)
|
||||||
|
status := strconv.Itoa(rw.statusCode)
|
||||||
|
|
||||||
|
p.RecordHTTPRequest(r.Method, r.URL.Path, status, duration)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Push manually pushes metrics to the configured Pushgateway
|
||||||
|
// Returns an error if pushing fails or if Pushgateway is not configured
|
||||||
|
func (p *PrometheusProvider) Push() error {
|
||||||
|
if p.pusher == nil {
|
||||||
|
return nil // Pushgateway not configured, silently skip
|
||||||
|
}
|
||||||
|
return p.pusher.Push()
|
||||||
|
}
|
||||||
|
|
||||||
|
// startAutoPush runs in a goroutine and periodically pushes metrics to Pushgateway
|
||||||
|
func (p *PrometheusProvider) startAutoPush() {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-p.pushTicker.C:
|
||||||
|
if err := p.Push(); err != nil {
|
||||||
|
// Log error but continue pushing
|
||||||
|
// Note: In production, you might want to use a proper logger
|
||||||
|
_ = err
|
||||||
|
}
|
||||||
|
case <-p.pushStop:
|
||||||
|
p.pushTicker.Stop()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// StopAutoPush stops the automatic push goroutine
|
||||||
|
// This should be called when shutting down the application
|
||||||
|
func (p *PrometheusProvider) StopAutoPush() {
|
||||||
|
if p.pushStop != nil {
|
||||||
|
close(p.pushStop)
|
||||||
|
}
|
||||||
|
}
|
||||||
+306
@@ -0,0 +1,306 @@
|
|||||||
|
package modelregistry
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ModelRules defines the permissions and security settings for a model
|
||||||
|
type ModelRules struct {
|
||||||
|
CanPublicRead bool // Whether the model can be read (GET operations)
|
||||||
|
CanPublicUpdate bool // Whether the model can be updated (PUT/PATCH operations)
|
||||||
|
CanPublicCreate bool // Whether the model can be created (POST operations)
|
||||||
|
CanPublicDelete bool // Whether the model can be deleted (DELETE operations)
|
||||||
|
CanRead bool // Whether the model can be read (GET operations)
|
||||||
|
CanUpdate bool // Whether the model can be updated (PUT/PATCH operations)
|
||||||
|
CanCreate bool // Whether the model can be created (POST operations)
|
||||||
|
CanDelete bool // Whether the model can be deleted (DELETE operations)
|
||||||
|
SecurityDisabled bool // Whether security checks are disabled for this model
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultModelRules returns the default rules for a model (all operations allowed, security enabled)
|
||||||
|
func DefaultModelRules() ModelRules {
|
||||||
|
return ModelRules{
|
||||||
|
CanRead: true,
|
||||||
|
CanUpdate: true,
|
||||||
|
CanCreate: true,
|
||||||
|
CanDelete: true,
|
||||||
|
CanPublicRead: false,
|
||||||
|
CanPublicUpdate: false,
|
||||||
|
CanPublicCreate: false,
|
||||||
|
CanPublicDelete: false,
|
||||||
|
SecurityDisabled: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultModelRegistry implements ModelRegistry interface
|
||||||
|
type DefaultModelRegistry struct {
|
||||||
|
models map[string]interface{}
|
||||||
|
rules map[string]ModelRules
|
||||||
|
mutex sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// Global default registry instance
|
||||||
|
var defaultRegistry = &DefaultModelRegistry{
|
||||||
|
models: make(map[string]interface{}),
|
||||||
|
rules: make(map[string]ModelRules),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Global list of registries (searched in order)
|
||||||
|
var registries = []*DefaultModelRegistry{defaultRegistry}
|
||||||
|
var registriesMutex sync.RWMutex
|
||||||
|
|
||||||
|
// NewModelRegistry creates a new model registry
|
||||||
|
func NewModelRegistry() *DefaultModelRegistry {
|
||||||
|
return &DefaultModelRegistry{
|
||||||
|
models: make(map[string]interface{}),
|
||||||
|
rules: make(map[string]ModelRules),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetDefaultRegistry() *DefaultModelRegistry {
|
||||||
|
return defaultRegistry
|
||||||
|
}
|
||||||
|
|
||||||
|
func SetDefaultRegistry(registry *DefaultModelRegistry) {
|
||||||
|
registriesMutex.Lock()
|
||||||
|
defer registriesMutex.Unlock()
|
||||||
|
|
||||||
|
foundAt := -1
|
||||||
|
for idx, r := range registries {
|
||||||
|
if r == defaultRegistry {
|
||||||
|
foundAt = idx
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
defaultRegistry = registry
|
||||||
|
if foundAt >= 0 {
|
||||||
|
registries[foundAt] = registry
|
||||||
|
} else {
|
||||||
|
registries = append([]*DefaultModelRegistry{registry}, registries...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddRegistry adds a registry to the global list of registries
|
||||||
|
// Registries are searched in the order they were added
|
||||||
|
func AddRegistry(registry *DefaultModelRegistry) {
|
||||||
|
registriesMutex.Lock()
|
||||||
|
defer registriesMutex.Unlock()
|
||||||
|
registries = append(registries, registry)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error {
|
||||||
|
r.mutex.Lock()
|
||||||
|
defer r.mutex.Unlock()
|
||||||
|
|
||||||
|
if _, exists := r.models[name]; exists {
|
||||||
|
return fmt.Errorf("model %s already registered", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate that model is a non-pointer struct
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
if modelType == nil {
|
||||||
|
return fmt.Errorf("model cannot be nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
originalType := modelType
|
||||||
|
|
||||||
|
// Unwrap pointers, slices, and arrays to check the underlying type
|
||||||
|
for modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate that the underlying type is a struct
|
||||||
|
if modelType.Kind() != reflect.Struct {
|
||||||
|
return fmt.Errorf("model must be a struct or pointer to struct, got %s", originalType.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// If a pointer/slice/array was passed, unwrap to the base struct
|
||||||
|
if originalType != modelType {
|
||||||
|
// Create a zero value of the struct type
|
||||||
|
model = reflect.New(modelType).Elem().Interface()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Additional check: ensure model is not a pointer
|
||||||
|
finalType := reflect.TypeOf(model)
|
||||||
|
if finalType.Kind() == reflect.Pointer {
|
||||||
|
return fmt.Errorf("model must be a non-pointer struct, got pointer to %s. Use MyModel{} instead of &MyModel{}", finalType.Elem().Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
r.models[name] = model
|
||||||
|
// Initialize with default rules if not already set
|
||||||
|
if _, exists := r.rules[name]; !exists {
|
||||||
|
r.rules[name] = DefaultModelRules()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
|
||||||
|
r.mutex.RLock()
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
|
model, exists := r.models[name]
|
||||||
|
if !exists {
|
||||||
|
return nil, fmt.Errorf("model %s not found", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
return model, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
||||||
|
r.mutex.RLock()
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
|
result := make(map[string]interface{})
|
||||||
|
for k, v := range r.models {
|
||||||
|
result[k] = v
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
||||||
|
// Try full name first
|
||||||
|
fullName := fmt.Sprintf("%s.%s", schema, entity)
|
||||||
|
if model, err := r.GetModel(fullName); err == nil {
|
||||||
|
return model, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback to entity name only
|
||||||
|
return r.GetModel(entity)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetModelRules sets the rules for a specific model
|
||||||
|
func (r *DefaultModelRegistry) SetModelRules(name string, rules ModelRules) error {
|
||||||
|
r.mutex.Lock()
|
||||||
|
defer r.mutex.Unlock()
|
||||||
|
|
||||||
|
// Check if model exists
|
||||||
|
if _, exists := r.models[name]; !exists {
|
||||||
|
return fmt.Errorf("model %s not found", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.rules[name] = rules
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelRules retrieves the rules for a specific model
|
||||||
|
// Returns default rules if model exists but rules are not set
|
||||||
|
func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
|
||||||
|
r.mutex.RLock()
|
||||||
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
|
// Check if model exists
|
||||||
|
if _, exists := r.models[name]; !exists {
|
||||||
|
return ModelRules{}, fmt.Errorf("model %s not found", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return rules if set, otherwise return default rules
|
||||||
|
if rules, exists := r.rules[name]; exists {
|
||||||
|
return rules, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return DefaultModelRules(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterModelWithRules registers a model with specific rules
|
||||||
|
func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error {
|
||||||
|
// First register the model
|
||||||
|
if err := r.RegisterModel(name, model); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Then set the rules (we need to lock again for rules)
|
||||||
|
r.mutex.Lock()
|
||||||
|
defer r.mutex.Unlock()
|
||||||
|
r.rules[name] = rules
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Global convenience functions using the default registry
|
||||||
|
|
||||||
|
// RegisterModel registers a model with the default global registry
|
||||||
|
func RegisterModel(model interface{}, name string) error {
|
||||||
|
return defaultRegistry.RegisterModel(name, model)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelByName retrieves a model by searching through all registries in order
|
||||||
|
// Returns the first match found
|
||||||
|
func GetModelByName(name string) (interface{}, error) {
|
||||||
|
registriesMutex.RLock()
|
||||||
|
defer registriesMutex.RUnlock()
|
||||||
|
|
||||||
|
for _, registry := range registries {
|
||||||
|
if model, err := registry.GetModel(name); err == nil {
|
||||||
|
return model, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("model %s not found in any registry", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IterateModels iterates over all models in the default global registry
|
||||||
|
func IterateModels(fn func(name string, model interface{})) {
|
||||||
|
defaultRegistry.mutex.RLock()
|
||||||
|
defer defaultRegistry.mutex.RUnlock()
|
||||||
|
|
||||||
|
for name, model := range defaultRegistry.models {
|
||||||
|
fn(name, model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModels returns a list of all models from all registries
|
||||||
|
// Models are collected in registry order, with duplicates included
|
||||||
|
func GetModels() []interface{} {
|
||||||
|
registriesMutex.RLock()
|
||||||
|
defer registriesMutex.RUnlock()
|
||||||
|
|
||||||
|
var models []interface{}
|
||||||
|
seen := make(map[string]bool)
|
||||||
|
|
||||||
|
for _, registry := range registries {
|
||||||
|
registry.mutex.RLock()
|
||||||
|
for name, model := range registry.models {
|
||||||
|
// Only add the first occurrence of each model name
|
||||||
|
if !seen[name] {
|
||||||
|
models = append(models, model)
|
||||||
|
seen[name] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
registry.mutex.RUnlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
return models
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetModelRules sets the rules for a specific model in the default registry
|
||||||
|
func SetModelRules(name string, rules ModelRules) error {
|
||||||
|
return defaultRegistry.SetModelRules(name, rules)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelRules retrieves the rules for a specific model from the default registry
|
||||||
|
func GetModelRules(name string) (ModelRules, error) {
|
||||||
|
return defaultRegistry.GetModelRules(name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelRulesByName retrieves the rules for a model by searching through all registries in order
|
||||||
|
// Returns the first match found
|
||||||
|
func GetModelRulesByName(name string) (ModelRules, error) {
|
||||||
|
registriesMutex.RLock()
|
||||||
|
defer registriesMutex.RUnlock()
|
||||||
|
|
||||||
|
for _, registry := range registries {
|
||||||
|
if _, err := registry.GetModel(name); err == nil {
|
||||||
|
// Model found in this registry, get its rules
|
||||||
|
return registry.GetModelRules(name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return ModelRules{}, fmt.Errorf("model %s not found in any registry", name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterModelWithRules registers a model with specific rules in the default registry
|
||||||
|
func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error {
|
||||||
|
return defaultRegistry.RegisterModelWithRules(name, model, rules)
|
||||||
|
}
|
||||||
+128
@@ -0,0 +1,128 @@
|
|||||||
|
package reflection
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ModelFieldDetail struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
DataType string `json:"datatype"`
|
||||||
|
SQLName string `json:"sqlname"`
|
||||||
|
SQLDataType string `json:"sqldatatype"`
|
||||||
|
SQLKey string `json:"sqlkey"`
|
||||||
|
Nullable bool `json:"nullable"`
|
||||||
|
FieldValue reflect.Value `json:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelColumnDetail - Get a list of columns in the SQL declaration of the model
|
||||||
|
// This function recursively processes embedded structs to include their fields
|
||||||
|
func GetModelColumnDetail(record reflect.Value) []ModelFieldDetail {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
logger.Error("Panic in GetModelColumnDetail : %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
lst := make([]ModelFieldDetail, 0)
|
||||||
|
|
||||||
|
if !record.IsValid() {
|
||||||
|
return lst
|
||||||
|
}
|
||||||
|
if record.Kind() == reflect.Pointer || record.Kind() == reflect.Interface {
|
||||||
|
record = record.Elem()
|
||||||
|
}
|
||||||
|
if record.Kind() != reflect.Struct {
|
||||||
|
return lst
|
||||||
|
}
|
||||||
|
|
||||||
|
collectFieldDetails(record, &lst)
|
||||||
|
|
||||||
|
return lst
|
||||||
|
}
|
||||||
|
|
||||||
|
// collectFieldDetails recursively collects field details from a struct value and its embedded fields
|
||||||
|
func collectFieldDetails(record reflect.Value, lst *[]ModelFieldDetail) {
|
||||||
|
modeltype := record.Type()
|
||||||
|
|
||||||
|
for i := 0; i < modeltype.NumField(); i++ {
|
||||||
|
fieldtype := modeltype.Field(i)
|
||||||
|
fieldValue := record.Field(i)
|
||||||
|
|
||||||
|
// Check if this is an embedded struct
|
||||||
|
if fieldtype.Anonymous {
|
||||||
|
// Unwrap pointer type if necessary
|
||||||
|
embeddedValue := fieldValue
|
||||||
|
if fieldValue.Kind() == reflect.Pointer {
|
||||||
|
if fieldValue.IsNil() {
|
||||||
|
// Skip nil embedded pointers
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
embeddedValue = fieldValue.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Recursively process embedded struct
|
||||||
|
if embeddedValue.Kind() == reflect.Struct {
|
||||||
|
collectFieldDetails(embeddedValue, lst)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
gormdetail := fieldtype.Tag.Get("gorm")
|
||||||
|
gormdetail = strings.Trim(gormdetail, " ")
|
||||||
|
fielddetail := ModelFieldDetail{}
|
||||||
|
fielddetail.FieldValue = fieldValue
|
||||||
|
fielddetail.Name = fieldtype.Name
|
||||||
|
fielddetail.DataType = fieldtype.Type.Name()
|
||||||
|
fielddetail.SQLName = fnFindKeyVal(gormdetail, "column:")
|
||||||
|
fielddetail.SQLDataType = fnFindKeyVal(gormdetail, "type:")
|
||||||
|
gormdetailLower := strings.ToLower(gormdetail)
|
||||||
|
switch {
|
||||||
|
case strings.Index(gormdetailLower, "identity") > 0 || strings.Index(gormdetailLower, "primary_key") > 0:
|
||||||
|
fielddetail.SQLKey = "primary_key"
|
||||||
|
case strings.Contains(gormdetailLower, "unique"):
|
||||||
|
fielddetail.SQLKey = "unique"
|
||||||
|
case strings.Contains(gormdetailLower, "uniqueindex"):
|
||||||
|
fielddetail.SQLKey = "uniqueindex"
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(strings.ToLower(gormdetail), "nullable") {
|
||||||
|
fielddetail.Nullable = true
|
||||||
|
} else if strings.Contains(strings.ToLower(gormdetail), "null") {
|
||||||
|
fielddetail.Nullable = true
|
||||||
|
}
|
||||||
|
if strings.Contains(strings.ToLower(gormdetail), "not null") {
|
||||||
|
fielddetail.Nullable = false
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(strings.ToLower(gormdetail), "foreignkey:") {
|
||||||
|
fielddetail.SQLKey = "foreign_key"
|
||||||
|
ik := strings.Index(strings.ToLower(gormdetail), "foreignkey:")
|
||||||
|
ie := strings.Index(gormdetail[ik:], ";")
|
||||||
|
if ie > ik && ik > 0 {
|
||||||
|
fielddetail.SQLName = strings.ToLower(gormdetail)[ik+11 : ik+ie]
|
||||||
|
// fmt.Printf("\r\nforeignkey: %v", fielddetail)
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
// ";foreignkey:rid_parent;association_foreignkey:id_atevent;save_associations:false;association_autocreate:false;"
|
||||||
|
|
||||||
|
*lst = append(*lst, fielddetail)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func fnFindKeyVal(src, key string) string {
|
||||||
|
icolStart := strings.Index(strings.ToLower(src), strings.ToLower(key))
|
||||||
|
val := ""
|
||||||
|
if icolStart >= 0 {
|
||||||
|
val = src[icolStart+len(key):]
|
||||||
|
icolend := strings.Index(val, ";")
|
||||||
|
if icolend > 0 {
|
||||||
|
val = val[:icolend]
|
||||||
|
}
|
||||||
|
return val
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
+137
@@ -0,0 +1,137 @@
|
|||||||
|
package reflection
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
func Len(v any) int {
|
||||||
|
val := reflect.ValueOf(v)
|
||||||
|
valKind := val.Kind()
|
||||||
|
|
||||||
|
if valKind == reflect.Pointer {
|
||||||
|
val = val.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
switch val.Kind() {
|
||||||
|
case reflect.Slice, reflect.Array, reflect.Map, reflect.String, reflect.Chan:
|
||||||
|
return val.Len()
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtractTableNameOnly extracts the table name from a fully qualified table reference.
|
||||||
|
// It removes any schema prefix (e.g., "schema.table" -> "table") and truncates at
|
||||||
|
// the first delimiter (comma, space, tab, or newline). If the input contains multiple
|
||||||
|
// dots, it returns everything after the last dot up to the first delimiter.
|
||||||
|
func ExtractTableNameOnly(fullName string) string {
|
||||||
|
// First, split by dot to remove schema prefix if present
|
||||||
|
lastDotIndex := -1
|
||||||
|
for i, char := range fullName {
|
||||||
|
if char == '.' {
|
||||||
|
lastDotIndex = i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Start from after the last dot (or from beginning if no dot)
|
||||||
|
startIndex := 0
|
||||||
|
if lastDotIndex != -1 {
|
||||||
|
startIndex = lastDotIndex + 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now find the end (first delimiter after the table name)
|
||||||
|
for i := startIndex; i < len(fullName); i++ {
|
||||||
|
char := rune(fullName[i])
|
||||||
|
if char == ',' || char == ' ' || char == '\t' || char == '\n' {
|
||||||
|
return fullName[startIndex:i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return fullName[startIndex:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsEmptyValue reports whether v is nil, an empty string, or a zero number.
|
||||||
|
func IsEmptyValue(v any) bool {
|
||||||
|
if v == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
rv := reflect.ValueOf(v)
|
||||||
|
if rv.Kind() == reflect.Pointer {
|
||||||
|
if rv.IsNil() {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
rv = rv.Elem()
|
||||||
|
}
|
||||||
|
switch rv.Kind() {
|
||||||
|
case reflect.String:
|
||||||
|
return rv.String() == ""
|
||||||
|
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||||
|
return rv.Int() == 0
|
||||||
|
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
|
||||||
|
return rv.Uint() == 0
|
||||||
|
case reflect.Float32, reflect.Float64:
|
||||||
|
return rv.Float() == 0
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetPointerElement returns the element type if the provided reflect.Type is a pointer.
|
||||||
|
// If the type is a slice of pointers, it returns the element type of the pointer within the slice.
|
||||||
|
// If neither condition is met, it returns the original type.
|
||||||
|
func GetPointerElement(v reflect.Type) reflect.Type {
|
||||||
|
if v.Kind() == reflect.Pointer {
|
||||||
|
return v.Elem()
|
||||||
|
}
|
||||||
|
if v.Kind() == reflect.Slice && v.Elem().Kind() == reflect.Pointer {
|
||||||
|
subElem := v.Elem()
|
||||||
|
if subElem.Elem().Kind() == reflect.Pointer {
|
||||||
|
return subElem.Elem().Elem()
|
||||||
|
}
|
||||||
|
return v.Elem()
|
||||||
|
}
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetJSONNameForField gets the JSON tag name for a struct field.
|
||||||
|
// Returns the JSON field name from the json struct tag, or an empty string if not found.
|
||||||
|
// Handles the "json" tag format: "name", "name,omitempty", etc.
|
||||||
|
func GetJSONNameForField(modelType reflect.Type, fieldName string) string {
|
||||||
|
if modelType == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unwrap pointer and slice indirections to reach the struct type
|
||||||
|
for {
|
||||||
|
switch modelType.Kind() {
|
||||||
|
case reflect.Pointer, reflect.Slice:
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
if modelType.Kind() != reflect.Struct {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find the field
|
||||||
|
field, found := modelType.FieldByName(fieldName)
|
||||||
|
if !found {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get the JSON tag
|
||||||
|
jsonTag := field.Tag.Get("json")
|
||||||
|
if jsonTag == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse the tag (format: "name,omitempty" or just "name")
|
||||||
|
parts := strings.Split(jsonTag, ",")
|
||||||
|
if len(parts) > 0 && parts[0] != "" && parts[0] != "-" {
|
||||||
|
return parts[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
return ""
|
||||||
|
}
|
||||||
+1743
File diff suppressed because it is too large
Load Diff
+572
@@ -0,0 +1,572 @@
|
|||||||
|
# ResolveSpec Query Features Examples
|
||||||
|
|
||||||
|
This document provides examples of using the advanced query features in ResolveSpec, including OR logic filters, Custom Operators, and FetchRowNumber.
|
||||||
|
|
||||||
|
## OR Logic in Filters (SearchOr)
|
||||||
|
|
||||||
|
### Basic OR Filter Example
|
||||||
|
|
||||||
|
Find all users with status "active" OR "pending":
|
||||||
|
|
||||||
|
```json
|
||||||
|
POST /users
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "pending",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Combined AND/OR Filters
|
||||||
|
|
||||||
|
Find users with (status="active" OR status="pending") AND age >= 18:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "pending",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "age",
|
||||||
|
"operator": "gte",
|
||||||
|
"value": 18
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**SQL Generated:** `WHERE (status = 'active' OR status = 'pending') AND age >= 18`
|
||||||
|
|
||||||
|
**Important Notes:**
|
||||||
|
- By default, filters use AND logic
|
||||||
|
- Consecutive filters with `"logic_operator": "OR"` are automatically grouped with parentheses
|
||||||
|
- This grouping ensures OR conditions don't interfere with AND conditions
|
||||||
|
- You don't need to specify `"logic_operator": "AND"` as it's the default
|
||||||
|
|
||||||
|
### Multiple OR Groups
|
||||||
|
|
||||||
|
You can have multiple separate OR groups:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "pending",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "priority",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "high"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "priority",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "urgent",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**SQL Generated:** `WHERE (status = 'active' OR status = 'pending') AND (priority = 'high' OR priority = 'urgent')`
|
||||||
|
|
||||||
|
## Custom Operators
|
||||||
|
|
||||||
|
### Simple Custom SQL Condition
|
||||||
|
|
||||||
|
Filter by email domain using custom SQL:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "company_emails",
|
||||||
|
"sql": "email LIKE '%@company.com'"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Multiple Custom Operators
|
||||||
|
|
||||||
|
Combine multiple custom SQL conditions:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "recent_active",
|
||||||
|
"sql": "last_login > NOW() - INTERVAL '30 days'"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "high_score",
|
||||||
|
"sql": "score > 1000"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Complex Custom Operator
|
||||||
|
|
||||||
|
Use complex SQL expressions:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "priority_users",
|
||||||
|
"sql": "(subscription = 'premium' AND points > 500) OR (subscription = 'enterprise')"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Combining Custom Operators with Regular Filters
|
||||||
|
|
||||||
|
Mix custom operators with standard filters:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "country",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "USA"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "active_last_month",
|
||||||
|
"sql": "last_activity > NOW() - INTERVAL '1 month'"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Row Numbers
|
||||||
|
|
||||||
|
### Two Ways to Get Row Numbers
|
||||||
|
|
||||||
|
There are two different features for row numbers:
|
||||||
|
|
||||||
|
1. **`fetch_row_number`** - Get the position of ONE specific record in a sorted/filtered set
|
||||||
|
2. **`RowNumber` field in models** - Automatically number all records in the response
|
||||||
|
|
||||||
|
### 1. FetchRowNumber - Get Position of Specific Record
|
||||||
|
|
||||||
|
Get the rank/position of a specific user in a leaderboard. **Important:** When `fetch_row_number` is specified, the response contains **ONLY that specific record**, not all records.
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "score",
|
||||||
|
"direction": "desc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"fetch_row_number": "12345"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Response - Contains ONLY the specified user:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": true,
|
||||||
|
"data": {
|
||||||
|
"id": 12345,
|
||||||
|
"name": "Alice Smith",
|
||||||
|
"score": 9850,
|
||||||
|
"level": 42
|
||||||
|
},
|
||||||
|
"metadata": {
|
||||||
|
"total": 10000,
|
||||||
|
"count": 1,
|
||||||
|
"filtered": 10000,
|
||||||
|
"row_number": 42
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Result:** User "12345" is ranked #42 out of 10,000 users. The response includes only Alice's data, not the other 9,999 users.
|
||||||
|
|
||||||
|
### Row Number with Filters
|
||||||
|
|
||||||
|
Find position within a filtered subset (e.g., "What's my rank in my country?"):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "country",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "USA"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "score",
|
||||||
|
"direction": "desc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"fetch_row_number": "12345"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Response:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": true,
|
||||||
|
"data": {
|
||||||
|
"id": 12345,
|
||||||
|
"name": "Bob Johnson",
|
||||||
|
"country": "USA",
|
||||||
|
"score": 7200,
|
||||||
|
"status": "active"
|
||||||
|
},
|
||||||
|
"metadata": {
|
||||||
|
"total": 2500,
|
||||||
|
"count": 1,
|
||||||
|
"filtered": 2500,
|
||||||
|
"row_number": 156
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Result:** Bob is ranked #156 out of 2,500 active USA users. Only Bob's record is returned.
|
||||||
|
|
||||||
|
### 2. RowNumber Field - Auto-Number All Records
|
||||||
|
|
||||||
|
If your model has a `RowNumber int64` field, restheadspec will automatically populate it for paginated results.
|
||||||
|
|
||||||
|
**Model Definition:**
|
||||||
|
```go
|
||||||
|
type Player struct {
|
||||||
|
ID int64 `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Score int64 `json:"score"`
|
||||||
|
RowNumber int64 `json:"row_number"` // Will be auto-populated
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Request (with pagination):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"sort": [{"column": "score", "direction": "desc"}],
|
||||||
|
"limit": 10,
|
||||||
|
"offset": 20
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Response - RowNumber automatically set:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": true,
|
||||||
|
"data": [
|
||||||
|
{
|
||||||
|
"id": 456,
|
||||||
|
"name": "Player21",
|
||||||
|
"score": 8900,
|
||||||
|
"row_number": 21
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 789,
|
||||||
|
"name": "Player22",
|
||||||
|
"score": 8850,
|
||||||
|
"row_number": 22
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 123,
|
||||||
|
"name": "Player23",
|
||||||
|
"score": 8800,
|
||||||
|
"row_number": 23
|
||||||
|
}
|
||||||
|
// ... records 24-30 ...
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**How It Works:**
|
||||||
|
- `row_number = offset + index + 1` (1-based)
|
||||||
|
- With offset=20, first record gets row_number=21
|
||||||
|
- With offset=20, second record gets row_number=22
|
||||||
|
- Perfect for displaying "Rank" in paginated tables
|
||||||
|
|
||||||
|
**Use Case:** Displaying leaderboards with rank numbers:
|
||||||
|
```
|
||||||
|
Rank | Player | Score
|
||||||
|
-----|-----------|-------
|
||||||
|
21 | Player21 | 8900
|
||||||
|
22 | Player22 | 8850
|
||||||
|
23 | Player23 | 8800
|
||||||
|
```
|
||||||
|
|
||||||
|
**Note:** This feature is available in all three packages: resolvespec, restheadspec, and websocketspec.
|
||||||
|
|
||||||
|
### When to Use Each Feature
|
||||||
|
|
||||||
|
| Feature | Use Case | Returns | Performance |
|
||||||
|
|---------|----------|---------|-------------|
|
||||||
|
| `fetch_row_number` | "What's my rank?" | 1 record with position | Fast - 1 record |
|
||||||
|
| `RowNumber` field | "Show top 10 with ranks" | Many records numbered | Fast - simple math |
|
||||||
|
|
||||||
|
**Combined Example - Full Leaderboard UI:**
|
||||||
|
|
||||||
|
```javascript
|
||||||
|
// Request 1: Get current user's rank
|
||||||
|
const userRank = await api.read({
|
||||||
|
fetch_row_number: currentUserId,
|
||||||
|
sort: [{column: "score", direction: "desc"}]
|
||||||
|
});
|
||||||
|
// Returns: {id: 123, name: "You", score: 7500, row_number: 156}
|
||||||
|
|
||||||
|
// Request 2: Get top 10 with rank numbers
|
||||||
|
const top10 = await api.read({
|
||||||
|
sort: [{column: "score", direction: "desc"}],
|
||||||
|
limit: 10,
|
||||||
|
offset: 0
|
||||||
|
});
|
||||||
|
// Returns: [{row_number: 1, ...}, {row_number: 2, ...}, ...]
|
||||||
|
|
||||||
|
// Display:
|
||||||
|
// "Your Rank: #156"
|
||||||
|
// "Top Players:"
|
||||||
|
// "#1 - Alice - 9999"
|
||||||
|
// "#2 - Bob - 9876"
|
||||||
|
// ...
|
||||||
|
```
|
||||||
|
|
||||||
|
## Complete Example: Advanced Query
|
||||||
|
|
||||||
|
Combine all features for a complex query:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"columns": ["id", "name", "email", "score", "status"],
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "trial",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "score",
|
||||||
|
"operator": "gte",
|
||||||
|
"value": 100
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "recent_activity",
|
||||||
|
"sql": "last_login > NOW() - INTERVAL '7 days'"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "verified_email",
|
||||||
|
"sql": "email_verified = true"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "score",
|
||||||
|
"direction": "desc"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "created_at",
|
||||||
|
"direction": "asc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"fetch_row_number": "12345",
|
||||||
|
"limit": 50,
|
||||||
|
"offset": 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
This query:
|
||||||
|
- Selects specific columns
|
||||||
|
- Filters for users with status "active" OR "trial"
|
||||||
|
- AND score >= 100
|
||||||
|
- Applies custom SQL conditions for recent activity and verified emails
|
||||||
|
- Sorts by score (descending) then creation date (ascending)
|
||||||
|
- Returns the row number of user "12345" in this filtered/sorted set
|
||||||
|
- Returns 50 records starting from the first one
|
||||||
|
|
||||||
|
## Use Cases
|
||||||
|
|
||||||
|
### 1. Leaderboards - Get Current User's Rank
|
||||||
|
|
||||||
|
Get the current user's position and data (returns only their record):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "game_id",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "game123"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "score",
|
||||||
|
"direction": "desc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"fetch_row_number": "current_user_id"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Tip:** For full leaderboards, make two requests:
|
||||||
|
1. One with `fetch_row_number` to get user's rank
|
||||||
|
2. One with `limit` and `offset` to get top players list
|
||||||
|
|
||||||
|
### 2. Multi-Status Search
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "order_status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "pending"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "order_status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "processing",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "order_status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "shipped",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. Advanced Date Filtering
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "this_month",
|
||||||
|
"sql": "created_at >= DATE_TRUNC('month', CURRENT_DATE)"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "business_hours",
|
||||||
|
"sql": "EXTRACT(HOUR FROM created_at) BETWEEN 9 AND 17"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Security Considerations
|
||||||
|
|
||||||
|
**Warning:** Custom operators allow raw SQL, which can be a security risk if not properly handled:
|
||||||
|
|
||||||
|
1. **Never** directly interpolate user input into custom operator SQL
|
||||||
|
2. Always validate and sanitize custom operator SQL on the backend
|
||||||
|
3. Consider using a whitelist of allowed custom operators
|
||||||
|
4. Use prepared statements or parameterized queries when possible
|
||||||
|
5. Implement proper authorization checks before executing queries
|
||||||
|
|
||||||
|
Example of safe custom operator handling in Go:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Whitelist of allowed custom operators
|
||||||
|
allowedOperators := map[string]string{
|
||||||
|
"recent_week": "created_at > NOW() - INTERVAL '7 days'",
|
||||||
|
"active_users": "status = 'active' AND last_login > NOW() - INTERVAL '30 days'",
|
||||||
|
"premium_only": "subscription_level = 'premium'",
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate custom operators from request
|
||||||
|
for _, op := range req.Options.CustomOperators {
|
||||||
|
if sql, ok := allowedOperators[op.Name]; ok {
|
||||||
|
op.SQL = sql // Use whitelisted SQL
|
||||||
|
} else {
|
||||||
|
return errors.New("custom operator not allowed: " + op.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
+853
@@ -0,0 +1,853 @@
|
|||||||
|
# ResolveSpec - Body-Based REST API
|
||||||
|
|
||||||
|
ResolveSpec provides a REST API where query options are passed in the JSON request body. This approach offers GraphQL-like flexibility while maintaining RESTful principles, making it ideal for complex queries and operations.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
* **Body-Based Querying**: All query options passed via JSON request body
|
||||||
|
* **Lifecycle Hooks**: Before/after hooks for create, read, update, delete operations
|
||||||
|
* **Cursor Pagination**: Efficient cursor-based pagination with complex sorting
|
||||||
|
* **Offset Pagination**: Traditional limit/offset pagination support
|
||||||
|
* **Advanced Filtering**: Multiple operators, AND/OR logic, and custom SQL
|
||||||
|
* **Relationship Preloading**: Load related entities with custom column selection and filters
|
||||||
|
* **Recursive CRUD**: Automatically handle nested object graphs with foreign key resolution
|
||||||
|
* **Computed Columns**: Define virtual columns with SQL expressions
|
||||||
|
* **Database-Agnostic**: Works with GORM, Bun, or custom database adapters
|
||||||
|
* **Router-Agnostic**: Integrates with any HTTP router through standard interfaces
|
||||||
|
* **Type-Safe**: Strong type validation and conversion
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
### Setup with GORM
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/resolvespec"
|
||||||
|
import "github.com/gorilla/mux"
|
||||||
|
|
||||||
|
// Create handler
|
||||||
|
handler := resolvespec.NewHandlerWithGORM(db)
|
||||||
|
|
||||||
|
// IMPORTANT: Register models BEFORE setting up routes
|
||||||
|
handler.registry.RegisterModel("core.users", &User{})
|
||||||
|
handler.registry.RegisterModel("core.posts", &Post{})
|
||||||
|
|
||||||
|
// Setup routes
|
||||||
|
router := mux.NewRouter()
|
||||||
|
resolvespec.SetupMuxRoutes(router, handler, nil)
|
||||||
|
|
||||||
|
// Start server
|
||||||
|
http.ListenAndServe(":8080", router)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Setup with Bun ORM
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/resolvespec"
|
||||||
|
import "github.com/uptrace/bun"
|
||||||
|
|
||||||
|
// Create handler with Bun
|
||||||
|
handler := resolvespec.NewHandlerWithBun(bunDB)
|
||||||
|
|
||||||
|
// Register models
|
||||||
|
handler.registry.RegisterModel("core.users", &User{})
|
||||||
|
|
||||||
|
// Setup routes (same as GORM)
|
||||||
|
router := mux.NewRouter()
|
||||||
|
resolvespec.SetupMuxRoutes(router, handler, nil)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Basic Usage
|
||||||
|
|
||||||
|
### Simple Read Request
|
||||||
|
|
||||||
|
```http
|
||||||
|
POST /core/users HTTP/1.1
|
||||||
|
Content-Type: application/json
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
### With Preloading
|
||||||
|
|
||||||
|
```http
|
||||||
|
POST /core/users HTTP/1.1
|
||||||
|
Content-Type: application/json
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
## Request Structure
|
||||||
|
|
||||||
|
### Request Format
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read|create|update|delete",
|
||||||
|
"data": {
|
||||||
|
// For create/update operations
|
||||||
|
},
|
||||||
|
"options": {
|
||||||
|
"columns": [...],
|
||||||
|
"preload": [...],
|
||||||
|
"filters": [...],
|
||||||
|
"sort": [...],
|
||||||
|
"limit": number,
|
||||||
|
"offset": number,
|
||||||
|
"cursor_forward": "string",
|
||||||
|
"cursor_backward": "string",
|
||||||
|
"customOperators": [...],
|
||||||
|
"computedColumns": [...]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Operations
|
||||||
|
|
||||||
|
| Operation | Description | Requires Data | Requires ID |
|
||||||
|
|-----------|-------------|---------------|-------------|
|
||||||
|
| `read` | Fetch records | No | Optional (single record) |
|
||||||
|
| `create` | Create new record(s) | Yes | No |
|
||||||
|
| `update` | Update existing record(s) | Yes | Yes (in URL) |
|
||||||
|
| `delete` | Delete record(s) | No | Yes (in URL) |
|
||||||
|
|
||||||
|
### Options Fields
|
||||||
|
|
||||||
|
| Field | Type | Description | Example |
|
||||||
|
|-------|------|-------------|---------|
|
||||||
|
| `columns` | `[]string` | Columns to select | `["id", "name", "email"]` |
|
||||||
|
| `preload` | `[]PreloadConfig` | Relations to load | See [Preloading](#preloading) |
|
||||||
|
| `filters` | `[]Filter` | Filter conditions | See [Filtering](#filtering) |
|
||||||
|
| `sort` | `[]Sort` | Sort criteria | `[{"column": "created_at", "direction": "desc"}]` |
|
||||||
|
| `limit` | `int` | Max records to return | `50` |
|
||||||
|
| `offset` | `int` | Number of records to skip | `100` |
|
||||||
|
| `cursor_forward` | `string` | Cursor for next page | `"12345"` |
|
||||||
|
| `cursor_backward` | `string` | Cursor for previous page | `"12300"` |
|
||||||
|
| `customOperators` | `[]CustomOperator` | Custom SQL conditions | See [Custom Operators](#custom-operators) |
|
||||||
|
| `computedColumns` | `[]ComputedColumn` | Virtual columns | See [Computed Columns](#computed-columns) |
|
||||||
|
|
||||||
|
## Filtering
|
||||||
|
|
||||||
|
### Available Operators
|
||||||
|
|
||||||
|
| Operator | Description | Example |
|
||||||
|
|----------|-------------|---------|
|
||||||
|
| `eq` | Equal | `{"column": "status", "operator": "eq", "value": "active"}` |
|
||||||
|
| `neq` | Not Equal | `{"column": "status", "operator": "neq", "value": "deleted"}` |
|
||||||
|
| `gt` | Greater Than | `{"column": "age", "operator": "gt", "value": 18}` |
|
||||||
|
| `gte` | Greater Than or Equal | `{"column": "age", "operator": "gte", "value": 18}` |
|
||||||
|
| `lt` | Less Than | `{"column": "price", "operator": "lt", "value": 100}` |
|
||||||
|
| `lte` | Less Than or Equal | `{"column": "price", "operator": "lte", "value": 100}` |
|
||||||
|
| `like` | LIKE pattern | `{"column": "name", "operator": "like", "value": "%john%"}` |
|
||||||
|
| `ilike` | Case-insensitive LIKE | `{"column": "email", "operator": "ilike", "value": "%@example.com"}` |
|
||||||
|
| `in` | IN clause | `{"column": "status", "operator": "in", "value": ["active", "pending"]}` |
|
||||||
|
| `contains` | Contains string | `{"column": "description", "operator": "contains", "value": "important"}` |
|
||||||
|
| `startswith` | Starts with string | `{"column": "name", "operator": "startswith", "value": "John"}` |
|
||||||
|
| `endswith` | Ends with string | `{"column": "email", "operator": "endswith", "value": "@example.com"}` |
|
||||||
|
| `between` | Between (exclusive) | `{"column": "age", "operator": "between", "value": [18, 65]}` |
|
||||||
|
| `betweeninclusive` | Between (inclusive) | `{"column": "price", "operator": "betweeninclusive", "value": [10, 100]}` |
|
||||||
|
| `empty` | IS NULL or empty | `{"column": "deleted_at", "operator": "empty"}` |
|
||||||
|
| `notempty` | IS NOT NULL | `{"column": "email", "operator": "notempty"}` |
|
||||||
|
|
||||||
|
### Complex Filtering Example
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "age",
|
||||||
|
"operator": "gte",
|
||||||
|
"value": 18
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "email",
|
||||||
|
"operator": "ilike",
|
||||||
|
"value": "%@company.com"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### OR Logic in Filters (SearchOr)
|
||||||
|
|
||||||
|
Use the `logic_operator` field to combine filters with OR logic instead of the default AND:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "pending",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "priority",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "high",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
This will produce: `WHERE (status = 'active' OR status = 'pending' OR priority = 'high')`
|
||||||
|
|
||||||
|
**Important:** Consecutive OR filters are automatically grouped together with parentheses to ensure proper query logic.
|
||||||
|
|
||||||
|
#### Mixing AND and OR
|
||||||
|
|
||||||
|
Consecutive OR filters are grouped, then combined with AND filters:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "pending",
|
||||||
|
"logic_operator": "OR"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "age",
|
||||||
|
"operator": "gte",
|
||||||
|
"value": 18
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Produces: `WHERE (status = 'active' OR status = 'pending') AND age >= 18`
|
||||||
|
|
||||||
|
This grouping ensures OR conditions don't interfere with other AND conditions in the query.
|
||||||
|
|
||||||
|
### Custom Operators
|
||||||
|
|
||||||
|
Add custom SQL conditions when needed:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "email_domain_filter",
|
||||||
|
"sql": "LOWER(email) LIKE '%@example.com'"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "recent_records",
|
||||||
|
"sql": "created_at > NOW() - INTERVAL '7 days'"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Custom operators are applied as additional WHERE conditions to your query.
|
||||||
|
|
||||||
|
### Fetch Row Number
|
||||||
|
|
||||||
|
Get the row number (position) of a specific record in the filtered and sorted result set. **When `fetch_row_number` is specified, only that specific record is returned** (not all records).
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "active"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "score",
|
||||||
|
"direction": "desc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"fetch_row_number": "12345"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Response - Returns ONLY the specified record with its position:**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": true,
|
||||||
|
"data": {
|
||||||
|
"id": 12345,
|
||||||
|
"name": "John Doe",
|
||||||
|
"score": 850,
|
||||||
|
"status": "active"
|
||||||
|
},
|
||||||
|
"metadata": {
|
||||||
|
"total": 1000,
|
||||||
|
"count": 1,
|
||||||
|
"filtered": 1000,
|
||||||
|
"row_number": 42
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Use Case:** Perfect for "Show me this user and their ranking" - you get just that one user with their position in the leaderboard.
|
||||||
|
|
||||||
|
**Note:** This is different from the `RowNumber` field feature, which automatically numbers all records in a paginated response based on offset. That feature uses simple math (`offset + index + 1`), while `fetch_row_number` uses SQL window functions to calculate the actual position in a sorted/filtered set. To use the `RowNumber` field feature, simply add a `RowNumber int64` field to your model - it will be automatically populated with the row position based on pagination.
|
||||||
|
|
||||||
|
## Preloading
|
||||||
|
|
||||||
|
Load related entities with custom configuration:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"columns": ["id", "name", "email"],
|
||||||
|
"preload": [
|
||||||
|
{
|
||||||
|
"relation": "posts",
|
||||||
|
"columns": ["id", "title", "created_at"],
|
||||||
|
"filters": [
|
||||||
|
{
|
||||||
|
"column": "status",
|
||||||
|
"operator": "eq",
|
||||||
|
"value": "published"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "created_at",
|
||||||
|
"direction": "desc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"limit": 5
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"relation": "profile",
|
||||||
|
"columns": ["bio", "website"]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Cursor Pagination
|
||||||
|
|
||||||
|
Efficient pagination for large datasets:
|
||||||
|
|
||||||
|
### First Request (No Cursor)
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "created_at",
|
||||||
|
"direction": "desc"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "id",
|
||||||
|
"direction": "asc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"limit": 50
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Next Page (Forward Cursor)
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "created_at",
|
||||||
|
"direction": "desc"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "id",
|
||||||
|
"direction": "asc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"limit": 50,
|
||||||
|
"cursor_forward": "12345"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Previous Page (Backward Cursor)
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"sort": [
|
||||||
|
{
|
||||||
|
"column": "created_at",
|
||||||
|
"direction": "desc"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"column": "id",
|
||||||
|
"direction": "asc"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"limit": 50,
|
||||||
|
"cursor_backward": "12300"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Benefits over offset pagination**:
|
||||||
|
* Consistent results when data changes
|
||||||
|
* Better performance for large offsets
|
||||||
|
* Prevents "skipped" or duplicate records
|
||||||
|
* Works with complex sort expressions
|
||||||
|
|
||||||
|
## Recursive CRUD Operations
|
||||||
|
|
||||||
|
Automatically handle nested object graphs with intelligent foreign key resolution.
|
||||||
|
|
||||||
|
### Creating Nested Objects
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "create",
|
||||||
|
"data": {
|
||||||
|
"name": "John Doe",
|
||||||
|
"email": "john@example.com",
|
||||||
|
"posts": [
|
||||||
|
{
|
||||||
|
"title": "My First Post",
|
||||||
|
"content": "Hello World",
|
||||||
|
"tags": [
|
||||||
|
{"name": "tech"},
|
||||||
|
{"name": "programming"}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"title": "Second Post",
|
||||||
|
"content": "More content"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"profile": {
|
||||||
|
"bio": "Software Developer",
|
||||||
|
"website": "https://example.com"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Per-Record Operation Control with `_request`
|
||||||
|
|
||||||
|
Control individual operations for each nested record:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "update",
|
||||||
|
"data": {
|
||||||
|
"name": "John Updated",
|
||||||
|
"posts": [
|
||||||
|
{
|
||||||
|
"_request": "insert",
|
||||||
|
"title": "New Post",
|
||||||
|
"content": "Fresh content"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"_request": "update",
|
||||||
|
"id": 456,
|
||||||
|
"title": "Updated Post Title"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"_request": "delete",
|
||||||
|
"id": 789
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Supported `_request` values**:
|
||||||
|
* `insert` - Create a new related record
|
||||||
|
* `update` - Update an existing related record
|
||||||
|
* `delete` - Delete a related record
|
||||||
|
* `upsert` - Create if doesn't exist, update if exists
|
||||||
|
|
||||||
|
**How It Works**:
|
||||||
|
1. Automatic foreign key resolution - parent IDs propagate to children
|
||||||
|
2. Recursive processing - handles nested relationships at any depth
|
||||||
|
3. Transaction safety - all operations execute atomically
|
||||||
|
4. Relationship detection - automatically detects belongsTo, hasMany, hasOne, many2many
|
||||||
|
5. Flexible operations - mix create, update, and delete in one request
|
||||||
|
|
||||||
|
## Computed Columns
|
||||||
|
|
||||||
|
Define virtual columns using SQL expressions:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"columns": ["id", "first_name", "last_name"],
|
||||||
|
"computedColumns": [
|
||||||
|
{
|
||||||
|
"name": "full_name",
|
||||||
|
"expression": "CONCAT(first_name, ' ', last_name)"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "age_years",
|
||||||
|
"expression": "EXTRACT(YEAR FROM AGE(birth_date))"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Custom Operators
|
||||||
|
|
||||||
|
Add custom SQL conditions when standard filters aren't sufficient:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"operation": "read",
|
||||||
|
"options": {
|
||||||
|
"customOperators": [
|
||||||
|
{
|
||||||
|
"name": "email_domain_filter",
|
||||||
|
"sql": "LOWER(email) LIKE '%@example.com'"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "recent_records",
|
||||||
|
"sql": "created_at > NOW() - INTERVAL '7 days'"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "complex_condition",
|
||||||
|
"sql": "(status = 'active' AND score > 100) OR (status = 'pending' AND priority = 'high')"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Note:** Custom operators are applied as WHERE conditions. Make sure to properly escape and sanitize any user input to prevent SQL injection.
|
||||||
|
|
||||||
|
## Lifecycle Hooks
|
||||||
|
|
||||||
|
Register hooks for all CRUD operations:
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/resolvespec"
|
||||||
|
|
||||||
|
// Create handler
|
||||||
|
handler := resolvespec.NewHandlerWithGORM(db)
|
||||||
|
|
||||||
|
// Register a before-read hook (e.g., for authorization)
|
||||||
|
handler.Hooks().Register(resolvespec.BeforeRead, func(ctx *resolvespec.HookContext) error {
|
||||||
|
// Check permissions
|
||||||
|
if !userHasPermission(ctx.Context, ctx.Entity) {
|
||||||
|
return fmt.Errorf("unauthorized access to %s", ctx.Entity)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Modify query options
|
||||||
|
if ctx.Options.Limit == nil || *ctx.Options.Limit > 100 {
|
||||||
|
ctx.Options.Limit = ptr(100) // Enforce max limit
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
// Register an after-read hook (e.g., for data transformation)
|
||||||
|
handler.Hooks().Register(resolvespec.AfterRead, func(ctx *resolvespec.HookContext) error {
|
||||||
|
// Transform or filter results
|
||||||
|
if users, ok := ctx.Result.([]User); ok {
|
||||||
|
for i := range users {
|
||||||
|
users[i].Email = maskEmail(users[i].Email)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
// Register a before-create hook (e.g., for validation)
|
||||||
|
handler.Hooks().Register(resolvespec.BeforeCreate, func(ctx *resolvespec.HookContext) error {
|
||||||
|
// Validate data
|
||||||
|
if user, ok := ctx.Data.(*User); ok {
|
||||||
|
if user.Email == "" {
|
||||||
|
return fmt.Errorf("email is required")
|
||||||
|
}
|
||||||
|
// Add timestamps
|
||||||
|
user.CreatedAt = time.Now()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
**Available Hook Types**:
|
||||||
|
* `BeforeHandle` — fires after model resolution, before operation dispatch (auth checks)
|
||||||
|
* `BeforeRead`, `AfterRead`
|
||||||
|
* `BeforeCreate`, `AfterCreate`
|
||||||
|
* `BeforeUpdate`, `AfterUpdate`
|
||||||
|
* `BeforeDelete`, `AfterDelete`
|
||||||
|
|
||||||
|
**HookContext** provides:
|
||||||
|
* `Context`: Request context
|
||||||
|
* `Handler`: Access to handler, database, and registry
|
||||||
|
* `Schema`, `Entity`, `TableName`: Request info
|
||||||
|
* `Model`: The registered model type
|
||||||
|
* `Operation`: Current operation string (`"read"`, `"create"`, `"update"`, `"delete"`)
|
||||||
|
* `Options`: Parsed request options (filters, sorting, etc.)
|
||||||
|
* `ID`: Record ID (for single-record operations)
|
||||||
|
* `Data`: Request data (for create/update)
|
||||||
|
* `Result`: Operation result (for after hooks)
|
||||||
|
* `Writer`: Response writer (allows hooks to modify response)
|
||||||
|
* `Abort`, `AbortMessage`, `AbortCode`: Set in hook to abort with an error response
|
||||||
|
|
||||||
|
## Model Registration
|
||||||
|
|
||||||
|
```go
|
||||||
|
type User struct {
|
||||||
|
ID uint `json:"id" gorm:"primaryKey"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
Posts []Post `json:"posts,omitempty" gorm:"foreignKey:UserID"`
|
||||||
|
Profile *Profile `json:"profile,omitempty" gorm:"foreignKey:UserID"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type Post struct {
|
||||||
|
ID uint `json:"id" gorm:"primaryKey"`
|
||||||
|
UserID uint `json:"user_id"`
|
||||||
|
Title string `json:"title"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
Tags []Tag `json:"tags,omitempty" gorm:"many2many:post_tags"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Schema.Table format
|
||||||
|
handler.registry.RegisterModel("core.users", &User{})
|
||||||
|
handler.registry.RegisterModel("core.posts", &Post{})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Complete Example
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/resolvespec"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
"gorm.io/driver/postgres"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
type User struct {
|
||||||
|
ID uint `json:"id" gorm:"primaryKey"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
Posts []Post `json:"posts,omitempty" gorm:"foreignKey:UserID"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type Post struct {
|
||||||
|
ID uint `json:"id" gorm:"primaryKey"`
|
||||||
|
UserID uint `json:"user_id"`
|
||||||
|
Title string `json:"title"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Connect to database
|
||||||
|
db, err := gorm.Open(postgres.Open("your-connection-string"), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create handler
|
||||||
|
handler := resolvespec.NewHandlerWithGORM(db)
|
||||||
|
|
||||||
|
// Register models
|
||||||
|
handler.registry.RegisterModel("core.users", &User{})
|
||||||
|
handler.registry.RegisterModel("core.posts", &Post{})
|
||||||
|
|
||||||
|
// Add hooks
|
||||||
|
handler.Hooks().Register(resolvespec.BeforeRead, func(ctx *resolvespec.HookContext) error {
|
||||||
|
log.Printf("Reading %s", ctx.Entity)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
// Setup routes
|
||||||
|
router := mux.NewRouter()
|
||||||
|
resolvespec.SetupMuxRoutes(router, handler, nil)
|
||||||
|
|
||||||
|
// Start server
|
||||||
|
log.Println("Server starting on :8080")
|
||||||
|
log.Fatal(http.ListenAndServe(":8080", router))
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
ResolveSpec is designed for testability:
|
||||||
|
|
||||||
|
```go
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestUserRead(t *testing.T) {
|
||||||
|
handler := resolvespec.NewHandlerWithGORM(testDB)
|
||||||
|
handler.registry.RegisterModel("core.users", &User{})
|
||||||
|
|
||||||
|
reqBody := map[string]interface{}{
|
||||||
|
"operation": "read",
|
||||||
|
"options": map[string]interface{}{
|
||||||
|
"columns": []string{"id", "name"},
|
||||||
|
"limit": 10,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
body, _ := json.Marshal(reqBody)
|
||||||
|
req := httptest.NewRequest("POST", "/core/users", bytes.NewReader(body))
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
|
// Test your handler...
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Router Integration
|
||||||
|
|
||||||
|
### Gorilla Mux
|
||||||
|
|
||||||
|
```go
|
||||||
|
router := mux.NewRouter()
|
||||||
|
resolvespec.SetupMuxRoutes(router, handler, nil)
|
||||||
|
```
|
||||||
|
|
||||||
|
### BunRouter
|
||||||
|
|
||||||
|
```go
|
||||||
|
router := bunrouter.New()
|
||||||
|
resolvespec.SetupBunRouterWithResolveSpec(router, handler)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Custom Routers
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Implement custom integration using common.Request and common.ResponseWriter
|
||||||
|
router.POST("/:schema/:entity", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
params := extractParams(r) // Your param extraction logic
|
||||||
|
reqAdapter := router.NewHTTPRequest(r)
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
handler.Handle(respAdapter, reqAdapter, params)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Response Format
|
||||||
|
|
||||||
|
### Success Response
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": true,
|
||||||
|
"data": [...],
|
||||||
|
"metadata": {
|
||||||
|
"total": 100,
|
||||||
|
"filtered": 50,
|
||||||
|
"limit": 10,
|
||||||
|
"offset": 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Error Response
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": false,
|
||||||
|
"error": {
|
||||||
|
"code": "validation_error",
|
||||||
|
"message": "Invalid request",
|
||||||
|
"details": "..."
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## See Also
|
||||||
|
|
||||||
|
* [Main README](../../README.md) - ResolveSpec overview
|
||||||
|
* [RestHeadSpec Package](../restheadspec/README.md) - Header-based API
|
||||||
|
* [StaticWeb Package](../server/staticweb/README.md) - Static file server
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
This package is part of ResolveSpec and is licensed under the MIT License.
|
||||||
|
```
|
||||||
|
|
||||||
|
## Response Format
|
||||||
|
|
||||||
|
### Success Response
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": true,
|
||||||
|
"data": [...],
|
||||||
|
"metadata": {
|
||||||
|
"total": 100,
|
||||||
|
"filtered": 50,
|
||||||
|
"limit": 10,
|
||||||
|
"offset": 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Error Response
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"success": false,
|
||||||
|
"error": {
|
||||||
|
"code": "validation_error",
|
||||||
|
"message": "Invalid request",
|
||||||
|
"details": "..."
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## See Also
|
||||||
|
|
||||||
|
* [Main README](../../README.md) - ResolveSpec overview
|
||||||
|
* [RestHeadSpec Package](../restheadspec/README.md) - Header-based API
|
||||||
|
* [StaticWeb Package](../server/staticweb/README.md) - Static file server
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
This package is part of ResolveSpec and is licensed under the MIT License.
|
||||||
+118
@@ -0,0 +1,118 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
// queryCacheKey represents the components used to build a cache key for query total count
|
||||||
|
type queryCacheKey struct {
|
||||||
|
TableName string `json:"table_name"`
|
||||||
|
Filters []common.FilterOption `json:"filters"`
|
||||||
|
Sort []common.SortOption `json:"sort"`
|
||||||
|
CustomSQLWhere string `json:"custom_sql_where,omitempty"`
|
||||||
|
CustomSQLOr string `json:"custom_sql_or,omitempty"`
|
||||||
|
CursorForward string `json:"cursor_forward,omitempty"`
|
||||||
|
CursorBackward string `json:"cursor_backward,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// cachedTotal represents a cached total count
|
||||||
|
type cachedTotal struct {
|
||||||
|
Total int `json:"total"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildQueryCacheKey builds a cache key from query parameters for total count caching
|
||||||
|
func buildQueryCacheKey(tableName string, filters []common.FilterOption, sort []common.SortOption, customWhere, customOr string) string {
|
||||||
|
key := queryCacheKey{
|
||||||
|
TableName: tableName,
|
||||||
|
Filters: filters,
|
||||||
|
Sort: sort,
|
||||||
|
CustomSQLWhere: customWhere,
|
||||||
|
CustomSQLOr: customOr,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serialize to JSON for consistent hashing
|
||||||
|
jsonData, err := json.Marshal(key)
|
||||||
|
if err != nil {
|
||||||
|
// Fallback to simple string concatenation if JSON fails
|
||||||
|
return hashString(fmt.Sprintf("%s_%v_%v_%s_%s", tableName, filters, sort, customWhere, customOr))
|
||||||
|
}
|
||||||
|
|
||||||
|
return hashString(string(jsonData))
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildExtendedQueryCacheKey builds a cache key for extended query options with cursor pagination
|
||||||
|
func buildExtendedQueryCacheKey(tableName string, filters []common.FilterOption, sort []common.SortOption,
|
||||||
|
customWhere, customOr string, cursorFwd, cursorBwd string) string {
|
||||||
|
|
||||||
|
key := queryCacheKey{
|
||||||
|
TableName: tableName,
|
||||||
|
Filters: filters,
|
||||||
|
Sort: sort,
|
||||||
|
CustomSQLWhere: customWhere,
|
||||||
|
CustomSQLOr: customOr,
|
||||||
|
CursorForward: cursorFwd,
|
||||||
|
CursorBackward: cursorBwd,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Serialize to JSON for consistent hashing
|
||||||
|
jsonData, err := json.Marshal(key)
|
||||||
|
if err != nil {
|
||||||
|
// Fallback to simple string concatenation if JSON fails
|
||||||
|
return hashString(fmt.Sprintf("%s_%v_%v_%s_%s_%s_%s",
|
||||||
|
tableName, filters, sort, customWhere, customOr, cursorFwd, cursorBwd))
|
||||||
|
}
|
||||||
|
|
||||||
|
return hashString(string(jsonData))
|
||||||
|
}
|
||||||
|
|
||||||
|
// hashString computes SHA256 hash of a string
|
||||||
|
func hashString(s string) string {
|
||||||
|
h := sha256.New()
|
||||||
|
h.Write([]byte(s))
|
||||||
|
return hex.EncodeToString(h.Sum(nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// getQueryTotalCacheKey returns a formatted cache key for storing/retrieving total count
|
||||||
|
func getQueryTotalCacheKey(hash string) string {
|
||||||
|
return fmt.Sprintf("query_total:%s", hash)
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildCacheTags creates cache tags from schema and table name
|
||||||
|
func buildCacheTags(schema, tableName string) []string {
|
||||||
|
return []string{
|
||||||
|
fmt.Sprintf("schema:%s", strings.ToLower(schema)),
|
||||||
|
fmt.Sprintf("table:%s", strings.ToLower(tableName)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// setQueryTotalCache stores a query total in the cache with schema and table tags
|
||||||
|
func setQueryTotalCache(ctx context.Context, cacheKey string, total int, schema, tableName string, ttl time.Duration) error {
|
||||||
|
c := cache.GetDefaultCache()
|
||||||
|
cacheData := cachedTotal{Total: total}
|
||||||
|
tags := buildCacheTags(schema, tableName)
|
||||||
|
|
||||||
|
return c.SetWithTags(ctx, cacheKey, cacheData, ttl, tags)
|
||||||
|
}
|
||||||
|
|
||||||
|
// invalidateCacheForTags removes all cached items matching the specified tags
|
||||||
|
func invalidateCacheForTags(ctx context.Context, tags []string) error {
|
||||||
|
c := cache.GetDefaultCache()
|
||||||
|
|
||||||
|
// Invalidate for each tag
|
||||||
|
for _, tag := range tags {
|
||||||
|
if err := c.DeleteByTag(ctx, tag); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+85
@@ -0,0 +1,85 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Context keys for request-scoped data
|
||||||
|
type contextKey string
|
||||||
|
|
||||||
|
const (
|
||||||
|
contextKeySchema contextKey = "schema"
|
||||||
|
contextKeyEntity contextKey = "entity"
|
||||||
|
contextKeyTableName contextKey = "tableName"
|
||||||
|
contextKeyModel contextKey = "model"
|
||||||
|
contextKeyModelPtr contextKey = "modelPtr"
|
||||||
|
)
|
||||||
|
|
||||||
|
// WithSchema adds schema to context
|
||||||
|
func WithSchema(ctx context.Context, schema string) context.Context {
|
||||||
|
return context.WithValue(ctx, contextKeySchema, schema)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSchema retrieves schema from context
|
||||||
|
func GetSchema(ctx context.Context) string {
|
||||||
|
if v := ctx.Value(contextKeySchema); v != nil {
|
||||||
|
return v.(string)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithEntity adds entity to context
|
||||||
|
func WithEntity(ctx context.Context, entity string) context.Context {
|
||||||
|
return context.WithValue(ctx, contextKeyEntity, entity)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetEntity retrieves entity from context
|
||||||
|
func GetEntity(ctx context.Context) string {
|
||||||
|
if v := ctx.Value(contextKeyEntity); v != nil {
|
||||||
|
return v.(string)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithTableName adds table name to context
|
||||||
|
func WithTableName(ctx context.Context, tableName string) context.Context {
|
||||||
|
return context.WithValue(ctx, contextKeyTableName, tableName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetTableName retrieves table name from context
|
||||||
|
func GetTableName(ctx context.Context) string {
|
||||||
|
if v := ctx.Value(contextKeyTableName); v != nil {
|
||||||
|
return v.(string)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithModel adds model to context
|
||||||
|
func WithModel(ctx context.Context, model interface{}) context.Context {
|
||||||
|
return context.WithValue(ctx, contextKeyModel, model)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModel retrieves model from context
|
||||||
|
func GetModel(ctx context.Context) interface{} {
|
||||||
|
return ctx.Value(contextKeyModel)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithModelPtr adds model pointer to context
|
||||||
|
func WithModelPtr(ctx context.Context, modelPtr interface{}) context.Context {
|
||||||
|
return context.WithValue(ctx, contextKeyModelPtr, modelPtr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelPtr retrieves model pointer from context
|
||||||
|
func GetModelPtr(ctx context.Context) interface{} {
|
||||||
|
return ctx.Value(contextKeyModelPtr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithRequestData adds all request-scoped data to context at once
|
||||||
|
func WithRequestData(ctx context.Context, schema, entity, tableName string, model, modelPtr interface{}) context.Context {
|
||||||
|
ctx = WithSchema(ctx, schema)
|
||||||
|
ctx = WithEntity(ctx, entity)
|
||||||
|
ctx = WithTableName(ctx, tableName)
|
||||||
|
ctx = WithModel(ctx, model)
|
||||||
|
ctx = WithModelPtr(ctx, modelPtr)
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
+210
@@ -0,0 +1,210 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CursorDirection defines pagination direction
|
||||||
|
type CursorDirection int
|
||||||
|
|
||||||
|
const (
|
||||||
|
CursorForward CursorDirection = 1
|
||||||
|
CursorBackward CursorDirection = -1
|
||||||
|
)
|
||||||
|
|
||||||
|
// GetCursorFilter generates a SQL `EXISTS` subquery for cursor-based pagination.
|
||||||
|
// It uses the current request's sort and cursor values.
|
||||||
|
//
|
||||||
|
// Parameters:
|
||||||
|
// - tableName: name of the main table (e.g. "posts")
|
||||||
|
// - pkName: primary key column (e.g. "id")
|
||||||
|
// - modelColumns: optional list of valid main-table columns (for validation). Pass nil to skip.
|
||||||
|
// - options: the request options containing sort and cursor information
|
||||||
|
// - expandJoins: optional map[alias]string of JOIN clauses for join-column sort support
|
||||||
|
//
|
||||||
|
// Returns SQL snippet to embed in WHERE clause.
|
||||||
|
func GetCursorFilter(
|
||||||
|
tableName string,
|
||||||
|
pkName string,
|
||||||
|
modelColumns []string,
|
||||||
|
options common.RequestOptions,
|
||||||
|
expandJoins map[string]string,
|
||||||
|
) (string, error) {
|
||||||
|
// Separate schema prefix from bare table name
|
||||||
|
fullTableName := tableName
|
||||||
|
if strings.Contains(tableName, ".") {
|
||||||
|
tableName = strings.SplitN(tableName, ".", 2)[1]
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
// 1. Determine active cursor
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
cursorID, direction := getActiveCursor(options)
|
||||||
|
if cursorID == "" {
|
||||||
|
return "", fmt.Errorf("no cursor provided for table %s", tableName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
// 2. Extract sort columns
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
sortItems := options.Sort
|
||||||
|
if len(sortItems) == 0 {
|
||||||
|
return "", fmt.Errorf("no sort columns defined")
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
// 3. Prepare
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
var whereClauses []string
|
||||||
|
joinSQL := ""
|
||||||
|
reverse := direction < 0
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
// 4. Process each sort column
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
for _, s := range sortItems {
|
||||||
|
col := strings.Trim(strings.TrimSpace(s.Column), "()")
|
||||||
|
if col == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse: "created_at", "user.name", "fn.sortorder", etc.
|
||||||
|
parts := strings.Split(col, ".")
|
||||||
|
field := strings.TrimSpace(parts[len(parts)-1])
|
||||||
|
prefix := strings.Join(parts[:len(parts)-1], ".")
|
||||||
|
|
||||||
|
// Direction from struct
|
||||||
|
desc := strings.EqualFold(s.Direction, "desc")
|
||||||
|
|
||||||
|
if reverse {
|
||||||
|
desc = !desc
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resolve column
|
||||||
|
cursorCol, targetCol, isJoin, err := resolveColumn(
|
||||||
|
field, prefix, tableName, modelColumns,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Skipping invalid sort column %q: %v", col, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle joins
|
||||||
|
if isJoin {
|
||||||
|
if expandJoins != nil {
|
||||||
|
if joinClause, ok := expandJoins[prefix]; ok {
|
||||||
|
jSQL, cRef := rewriteJoin(joinClause, tableName, prefix)
|
||||||
|
joinSQL = jSQL
|
||||||
|
cursorCol = cRef + "." + field
|
||||||
|
targetCol = prefix + "." + field
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if cursorCol == "" {
|
||||||
|
logger.Warn("Skipping cursor sort column %q: join alias %q not in expandJoins", col, prefix)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build inequality
|
||||||
|
op := "<"
|
||||||
|
if desc {
|
||||||
|
op = ">"
|
||||||
|
}
|
||||||
|
whereClauses = append(whereClauses, fmt.Sprintf("%s %s %s", cursorCol, op, targetCol))
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(whereClauses) == 0 {
|
||||||
|
return "", fmt.Errorf("no valid sort columns after filtering")
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
// 5. Build priority OR-AND chain
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
orSQL := buildPriorityChain(whereClauses)
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
// 6. Final EXISTS subquery
|
||||||
|
// --------------------------------------------------------------------- //
|
||||||
|
query := fmt.Sprintf(`EXISTS (
|
||||||
|
SELECT 1
|
||||||
|
FROM %s cursor_select
|
||||||
|
%s
|
||||||
|
WHERE cursor_select.%s = %s
|
||||||
|
AND (%s)
|
||||||
|
)`,
|
||||||
|
fullTableName,
|
||||||
|
joinSQL,
|
||||||
|
pkName,
|
||||||
|
cursorID,
|
||||||
|
orSQL,
|
||||||
|
)
|
||||||
|
|
||||||
|
return query, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ------------------------------------------------------------------------- //
|
||||||
|
// Helper: get active cursor (forward or backward)
|
||||||
|
func getActiveCursor(options common.RequestOptions) (id string, direction CursorDirection) {
|
||||||
|
if options.CursorForward != "" {
|
||||||
|
return options.CursorForward, CursorForward
|
||||||
|
}
|
||||||
|
if options.CursorBackward != "" {
|
||||||
|
return options.CursorBackward, CursorBackward
|
||||||
|
}
|
||||||
|
return "", 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper: resolve column (main table or join)
|
||||||
|
func resolveColumn(
|
||||||
|
field, prefix, tableName string,
|
||||||
|
modelColumns []string,
|
||||||
|
) (cursorCol, targetCol string, isJoin bool, err error) {
|
||||||
|
|
||||||
|
// JSON field
|
||||||
|
if strings.Contains(field, "->") {
|
||||||
|
return "cursor_select." + field, tableName + "." + field, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Main table column
|
||||||
|
if modelColumns != nil {
|
||||||
|
for _, col := range modelColumns {
|
||||||
|
if strings.EqualFold(col, field) {
|
||||||
|
return "cursor_select." + field, tableName + "." + field, false, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// No validation → allow all main-table fields
|
||||||
|
return "cursor_select." + field, tableName + "." + field, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Joined column
|
||||||
|
if prefix != "" && prefix != tableName {
|
||||||
|
return "", "", true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", "", false, fmt.Errorf("invalid column: %s", field)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper: rewrite JOIN clause for cursor subquery
|
||||||
|
func rewriteJoin(joinClause, mainTable, alias string) (joinSQL, cursorAlias string) {
|
||||||
|
joinSQL = strings.ReplaceAll(joinClause, mainTable+".", "cursor_select.")
|
||||||
|
cursorAlias = "cursor_select_" + alias
|
||||||
|
joinSQL = strings.ReplaceAll(joinSQL, " "+alias+" ", " "+cursorAlias+" ")
|
||||||
|
joinSQL = strings.ReplaceAll(joinSQL, " "+alias+".", " "+cursorAlias+".")
|
||||||
|
return joinSQL, cursorAlias
|
||||||
|
}
|
||||||
|
|
||||||
|
// ------------------------------------------------------------------------- //
|
||||||
|
// Helper: build OR-AND priority chain
|
||||||
|
func buildPriorityChain(clauses []string) string {
|
||||||
|
var or []string
|
||||||
|
for i := 0; i < len(clauses); i++ {
|
||||||
|
and := strings.Join(clauses[:i+1], "\n AND ")
|
||||||
|
or = append(or, "("+and+")")
|
||||||
|
}
|
||||||
|
return strings.Join(or, "\n OR ")
|
||||||
|
}
|
||||||
+2231
File diff suppressed because it is too large
Load Diff
+163
@@ -0,0 +1,163 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// HookType defines the type of hook to execute
|
||||||
|
type HookType string
|
||||||
|
|
||||||
|
const (
|
||||||
|
// BeforeHandle fires after model resolution, before operation dispatch.
|
||||||
|
// Use this for auth checks that need model rules and user context simultaneously.
|
||||||
|
BeforeHandle HookType = "before_handle"
|
||||||
|
|
||||||
|
// Read operation hooks
|
||||||
|
BeforeRead HookType = "before_read"
|
||||||
|
AfterRead HookType = "after_read"
|
||||||
|
|
||||||
|
// Create operation hooks
|
||||||
|
BeforeCreate HookType = "before_create"
|
||||||
|
AfterCreate HookType = "after_create"
|
||||||
|
|
||||||
|
// Update operation hooks
|
||||||
|
BeforeUpdate HookType = "before_update"
|
||||||
|
AfterUpdate HookType = "after_update"
|
||||||
|
|
||||||
|
// Delete operation hooks
|
||||||
|
BeforeDelete HookType = "before_delete"
|
||||||
|
AfterDelete HookType = "after_delete"
|
||||||
|
|
||||||
|
// Scan/Execute operation hooks (for query building)
|
||||||
|
BeforeScan HookType = "before_scan"
|
||||||
|
)
|
||||||
|
|
||||||
|
// HookContext contains all the data available to a hook
|
||||||
|
type HookContext struct {
|
||||||
|
Context context.Context
|
||||||
|
Handler *Handler // Reference to the handler for accessing database, registry, etc.
|
||||||
|
Schema string
|
||||||
|
Entity string
|
||||||
|
Model interface{}
|
||||||
|
Options common.RequestOptions
|
||||||
|
Writer common.ResponseWriter
|
||||||
|
Request common.Request
|
||||||
|
|
||||||
|
// Operation being dispatched (e.g. "read", "create", "update", "delete")
|
||||||
|
Operation string
|
||||||
|
|
||||||
|
// Operation-specific fields
|
||||||
|
ID string
|
||||||
|
Data interface{} // For create/update operations
|
||||||
|
Result interface{} // For after hooks
|
||||||
|
Error error // For after hooks
|
||||||
|
|
||||||
|
// Query chain - allows hooks to modify the query before execution
|
||||||
|
Query common.SelectQuery
|
||||||
|
|
||||||
|
// Allow hooks to abort the operation
|
||||||
|
Abort bool // If set to true, the operation will be aborted
|
||||||
|
AbortMessage string // Message to return if aborted
|
||||||
|
AbortCode int // HTTP status code if aborted
|
||||||
|
|
||||||
|
// Tx provides access to the database/transaction for executing additional SQL
|
||||||
|
// This allows hooks to run custom queries in addition to the main Query chain
|
||||||
|
Tx common.Database
|
||||||
|
}
|
||||||
|
|
||||||
|
// HookFunc is the signature for hook functions
|
||||||
|
// It receives a HookContext and can modify it or return an error
|
||||||
|
// If an error is returned, the operation will be aborted
|
||||||
|
type HookFunc func(*HookContext) error
|
||||||
|
|
||||||
|
// HookRegistry manages all registered hooks
|
||||||
|
type HookRegistry struct {
|
||||||
|
hooks map[HookType][]HookFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewHookRegistry creates a new hook registry
|
||||||
|
func NewHookRegistry() *HookRegistry {
|
||||||
|
return &HookRegistry{
|
||||||
|
hooks: make(map[HookType][]HookFunc),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register adds a new hook for the specified hook type
|
||||||
|
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
|
||||||
|
if r.hooks == nil {
|
||||||
|
r.hooks = make(map[HookType][]HookFunc)
|
||||||
|
}
|
||||||
|
r.hooks[hookType] = append(r.hooks[hookType], hook)
|
||||||
|
logger.Info("Registered resolvespec hook for %s (total: %d)", hookType, len(r.hooks[hookType]))
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterMultiple registers a hook for multiple hook types
|
||||||
|
func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
|
||||||
|
for _, hookType := range hookTypes {
|
||||||
|
r.Register(hookType, hook)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute runs all hooks for the specified type in order
|
||||||
|
// If any hook returns an error, execution stops and the error is returned
|
||||||
|
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
|
||||||
|
hooks, exists := r.hooks[hookType]
|
||||||
|
if !exists || len(hooks) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Executing %d resolvespec hook(s) for %s", len(hooks), hookType)
|
||||||
|
|
||||||
|
for i, hook := range hooks {
|
||||||
|
if err := hook(ctx); err != nil {
|
||||||
|
logger.Error("Resolvespec hook %d for %s failed: %v", i+1, hookType, err)
|
||||||
|
return fmt.Errorf("hook execution failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if hook requested abort
|
||||||
|
if ctx.Abort {
|
||||||
|
logger.Warn("Resolvespec hook %d for %s requested abort: %s", i+1, hookType, ctx.AbortMessage)
|
||||||
|
return fmt.Errorf("operation aborted by hook: %s", ctx.AbortMessage)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear removes all hooks for the specified type
|
||||||
|
func (r *HookRegistry) Clear(hookType HookType) {
|
||||||
|
delete(r.hooks, hookType)
|
||||||
|
logger.Info("Cleared all resolvespec hooks for %s", hookType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearAll removes all registered hooks
|
||||||
|
func (r *HookRegistry) ClearAll() {
|
||||||
|
r.hooks = make(map[HookType][]HookFunc)
|
||||||
|
logger.Info("Cleared all resolvespec hooks")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Count returns the number of hooks registered for a specific type
|
||||||
|
func (r *HookRegistry) Count(hookType HookType) int {
|
||||||
|
if hooks, exists := r.hooks[hookType]; exists {
|
||||||
|
return len(hooks)
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// HasHooks returns true if there are any hooks registered for the specified type
|
||||||
|
func (r *HookRegistry) HasHooks(hookType HookType) bool {
|
||||||
|
return r.Count(hookType) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetAllHookTypes returns all hook types that have registered hooks
|
||||||
|
func (r *HookRegistry) GetAllHookTypes() []HookType {
|
||||||
|
types := make([]HookType, 0, len(r.hooks))
|
||||||
|
for hookType := range r.hooks {
|
||||||
|
types = append(types, hookType)
|
||||||
|
}
|
||||||
|
return types
|
||||||
|
}
|
||||||
+27
@@ -0,0 +1,27 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
// Legacy interfaces for backward compatibility
|
||||||
|
type GormTableNameInterface interface {
|
||||||
|
TableName() string
|
||||||
|
}
|
||||||
|
|
||||||
|
type GormTableSchemaInterface interface {
|
||||||
|
TableSchema() string
|
||||||
|
}
|
||||||
|
|
||||||
|
type GormTableCRUDRequest struct {
|
||||||
|
Request *string `json:"_request"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *GormTableCRUDRequest) SetRequest(request string) {
|
||||||
|
r.Request = &request
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r GormTableCRUDRequest) GetRequest() string {
|
||||||
|
return *r.Request
|
||||||
|
}
|
||||||
|
|
||||||
|
// New interfaces that replace the legacy ones above
|
||||||
|
// These are now defined in database.go:
|
||||||
|
// - TableNameProvider (replaces GormTableNameInterface)
|
||||||
|
// - SchemaProvider (replaces GormTableSchemaInterface)
|
||||||
+513
@@ -0,0 +1,513 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bunrouter"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/router"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NewHandlerWithGORM creates a new Handler with GORM adapter
|
||||||
|
func NewHandlerWithGORM(db *gorm.DB) *Handler {
|
||||||
|
gormAdapter := database.NewGormAdapter(db)
|
||||||
|
registry := modelregistry.NewModelRegistry()
|
||||||
|
return NewHandler(gormAdapter, registry)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewHandlerWithBun creates a new Handler with Bun adapter
|
||||||
|
func NewHandlerWithBun(db *bun.DB) *Handler {
|
||||||
|
bunAdapter := database.NewBunAdapter(db)
|
||||||
|
registry := modelregistry.NewModelRegistry()
|
||||||
|
return NewHandler(bunAdapter, registry)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewStandardMuxRouter creates a router with standard Mux HTTP handlers
|
||||||
|
func NewStandardMuxRouter() *router.StandardMuxAdapter {
|
||||||
|
return router.NewStandardMuxAdapter()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewStandardBunRouter creates a router with standard BunRouter handlers
|
||||||
|
func NewStandardBunRouter() *router.StandardBunRouterAdapter {
|
||||||
|
return router.NewStandardBunRouterAdapter()
|
||||||
|
}
|
||||||
|
|
||||||
|
// MiddlewareFunc is a function that wraps an http.Handler with additional functionality
|
||||||
|
type MiddlewareFunc func(http.Handler) http.Handler
|
||||||
|
|
||||||
|
// SetupMuxRoutes sets up routes for the ResolveSpec API with Mux
|
||||||
|
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
||||||
|
// Example: SetupMuxRoutes(router, handler, func(h http.Handler) http.Handler { return security.NewAuthHandler(securityList, h) })
|
||||||
|
func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler, authMiddleware MiddlewareFunc) {
|
||||||
|
// Add global /openapi route
|
||||||
|
openAPIHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
corsConfig := common.DefaultCORSConfig()
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(r)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
|
||||||
|
handler.HandleOpenAPI(respAdapter, reqAdapter)
|
||||||
|
})
|
||||||
|
muxRouter.Handle("/openapi", openAPIHandler).Methods("GET", "OPTIONS")
|
||||||
|
|
||||||
|
// Get all registered models from the registry
|
||||||
|
allModels := handler.registry.GetAllModels()
|
||||||
|
|
||||||
|
// Loop through each registered model and create explicit routes
|
||||||
|
for fullName := range allModels {
|
||||||
|
// Parse the full name (e.g., "public.users" or just "users")
|
||||||
|
schema, entity := parseModelName(fullName)
|
||||||
|
|
||||||
|
// Build the route paths
|
||||||
|
entityPath := buildRoutePath(schema, entity)
|
||||||
|
entityWithIDPath := buildRoutePath(schema, entity) + "/{id}"
|
||||||
|
|
||||||
|
// Create handler functions for this specific entity
|
||||||
|
var postEntityHandler http.Handler = createMuxHandler(handler, schema, entity, "")
|
||||||
|
var postEntityWithIDHandler http.Handler = createMuxHandler(handler, schema, entity, "id")
|
||||||
|
var getEntityHandler http.Handler = createMuxGetHandler(handler, schema, entity, "")
|
||||||
|
optionsEntityHandler := createMuxOptionsHandler(handler, schema, entity, []string{"GET", "POST", "OPTIONS"})
|
||||||
|
optionsEntityWithIDHandler := createMuxOptionsHandler(handler, schema, entity, []string{"POST", "OPTIONS"})
|
||||||
|
|
||||||
|
// Apply authentication middleware if provided
|
||||||
|
if authMiddleware != nil {
|
||||||
|
postEntityHandler = authMiddleware(postEntityHandler)
|
||||||
|
postEntityWithIDHandler = authMiddleware(postEntityWithIDHandler)
|
||||||
|
getEntityHandler = authMiddleware(getEntityHandler)
|
||||||
|
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register routes for this entity
|
||||||
|
muxRouter.Handle(entityPath, postEntityHandler).Methods("POST")
|
||||||
|
muxRouter.Handle(entityWithIDPath, postEntityWithIDHandler).Methods("POST")
|
||||||
|
muxRouter.Handle(entityPath, getEntityHandler).Methods("GET")
|
||||||
|
muxRouter.Handle(entityPath, optionsEntityHandler).Methods("OPTIONS")
|
||||||
|
muxRouter.Handle(entityWithIDPath, optionsEntityWithIDHandler).Methods("OPTIONS")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper function to create Mux handler for a specific entity with CORS support
|
||||||
|
func createMuxHandler(handler *Handler, schema, entity, idParam string) http.HandlerFunc {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Set CORS headers
|
||||||
|
corsConfig := common.DefaultCORSConfig()
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(r)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
|
||||||
|
vars := make(map[string]string)
|
||||||
|
vars["schema"] = schema
|
||||||
|
vars["entity"] = entity
|
||||||
|
if idParam != "" {
|
||||||
|
vars["id"] = mux.Vars(r)[idParam]
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.Handle(respAdapter, reqAdapter, vars)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper function to create Mux GET handler for a specific entity with CORS support
|
||||||
|
func createMuxGetHandler(handler *Handler, schema, entity, idParam string) http.HandlerFunc {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Set CORS headers
|
||||||
|
corsConfig := common.DefaultCORSConfig()
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(r)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
|
||||||
|
vars := make(map[string]string)
|
||||||
|
vars["schema"] = schema
|
||||||
|
vars["entity"] = entity
|
||||||
|
if idParam != "" {
|
||||||
|
vars["id"] = mux.Vars(r)[idParam]
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.HandleGet(respAdapter, reqAdapter, vars)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper function to create Mux OPTIONS handler that returns metadata
|
||||||
|
func createMuxOptionsHandler(handler *Handler, schema, entity string, allowedMethods []string) http.HandlerFunc {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Set CORS headers with the allowed methods for this route
|
||||||
|
corsConfig := common.DefaultCORSConfig()
|
||||||
|
corsConfig.AllowedMethods = allowedMethods
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(r)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
|
||||||
|
// Return metadata in the OPTIONS response body
|
||||||
|
vars := make(map[string]string)
|
||||||
|
vars["schema"] = schema
|
||||||
|
vars["entity"] = entity
|
||||||
|
|
||||||
|
handler.HandleGet(respAdapter, reqAdapter, vars)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseModelName parses a model name like "public.users" into schema and entity
|
||||||
|
// If no schema is present, returns empty string for schema
|
||||||
|
func parseModelName(fullName string) (schema, entity string) {
|
||||||
|
parts := strings.Split(fullName, ".")
|
||||||
|
if len(parts) == 2 {
|
||||||
|
return parts[0], parts[1]
|
||||||
|
}
|
||||||
|
return "", fullName
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildRoutePath builds a route path from schema and entity
|
||||||
|
// If schema is empty, returns just "/entity", otherwise "/{schema}/{entity}"
|
||||||
|
func buildRoutePath(schema, entity string) string {
|
||||||
|
if schema == "" {
|
||||||
|
return "/" + entity
|
||||||
|
}
|
||||||
|
return "/" + schema + "/" + entity
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example usage functions for documentation:
|
||||||
|
|
||||||
|
// ExampleWithGORM shows how to use ResolveSpec with GORM
|
||||||
|
func ExampleWithGORM(db *gorm.DB) {
|
||||||
|
// Create handler using GORM
|
||||||
|
handler := NewHandlerWithGORM(db)
|
||||||
|
|
||||||
|
// Setup router without authentication
|
||||||
|
muxRouter := mux.NewRouter()
|
||||||
|
SetupMuxRoutes(muxRouter, handler, nil)
|
||||||
|
|
||||||
|
// Register models
|
||||||
|
// handler.RegisterModel("public", "users", &User{})
|
||||||
|
|
||||||
|
// To add authentication, pass a middleware function:
|
||||||
|
// import "github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
// secList := security.NewSecurityList(myProvider)
|
||||||
|
// authMiddleware := func(h http.Handler) http.Handler {
|
||||||
|
// return security.NewAuthHandler(secList, h)
|
||||||
|
// }
|
||||||
|
// SetupMuxRoutes(muxRouter, handler, authMiddleware)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleWithBun shows how to switch to Bun ORM
|
||||||
|
func ExampleWithBun(bunDB *bun.DB) {
|
||||||
|
// Create Bun adapter
|
||||||
|
dbAdapter := database.NewBunAdapter(bunDB)
|
||||||
|
|
||||||
|
// Create model registry
|
||||||
|
registry := modelregistry.NewModelRegistry()
|
||||||
|
// registry.RegisterModel("public.users", &User{})
|
||||||
|
|
||||||
|
// Create handler
|
||||||
|
handler := NewHandler(dbAdapter, registry)
|
||||||
|
|
||||||
|
// Setup routes without authentication
|
||||||
|
muxRouter := mux.NewRouter()
|
||||||
|
SetupMuxRoutes(muxRouter, handler, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BunRouterHandler is an interface that both bunrouter.Router and bunrouter.Group implement
|
||||||
|
type BunRouterHandler interface {
|
||||||
|
Handle(method, path string, handler bunrouter.HandlerFunc)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wrapBunRouterHandler wraps a bunrouter handler with auth middleware if provided
|
||||||
|
func wrapBunRouterHandler(handler bunrouter.HandlerFunc, authMiddleware MiddlewareFunc) bunrouter.HandlerFunc {
|
||||||
|
if authMiddleware == nil {
|
||||||
|
return handler
|
||||||
|
}
|
||||||
|
|
||||||
|
return func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
// Create an http.Handler that calls the bunrouter handler
|
||||||
|
httpHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Replace the embedded *http.Request with the middleware-enriched one
|
||||||
|
// so that auth context (user ID, etc.) is visible to the handler.
|
||||||
|
enrichedReq := req
|
||||||
|
enrichedReq.Request = r
|
||||||
|
_ = handler(w, enrichedReq)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Wrap with auth middleware and execute
|
||||||
|
wrappedHandler := authMiddleware(httpHandler)
|
||||||
|
wrappedHandler.ServeHTTP(w, req.Request)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetupBunRouterRoutes sets up bunrouter routes for the ResolveSpec API
|
||||||
|
// Accepts bunrouter.Router or bunrouter.Group
|
||||||
|
// authMiddleware is optional - if provided, routes will be protected with the middleware
|
||||||
|
func SetupBunRouterRoutes(r BunRouterHandler, handler *Handler, authMiddleware MiddlewareFunc) {
|
||||||
|
|
||||||
|
// CORS config
|
||||||
|
corsConfig := common.DefaultCORSConfig()
|
||||||
|
|
||||||
|
// Add global /openapi route
|
||||||
|
r.Handle("GET", "/openapi", func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
handler.HandleOpenAPI(respAdapter, reqAdapter)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
r.Handle("OPTIONS", "/openapi", func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
// Get all registered models from the registry
|
||||||
|
allModels := handler.registry.GetAllModels()
|
||||||
|
|
||||||
|
// Loop through each registered model and create explicit routes
|
||||||
|
for fullName := range allModels {
|
||||||
|
// Parse the full name (e.g., "public.users" or just "users")
|
||||||
|
schema, entity := parseModelName(fullName)
|
||||||
|
|
||||||
|
// Build the route paths
|
||||||
|
entityPath := buildRoutePath(schema, entity)
|
||||||
|
entityWithIDPath := entityPath + "/:id"
|
||||||
|
|
||||||
|
// Create closure variables to capture current schema and entity
|
||||||
|
currentSchema := schema
|
||||||
|
currentEntity := entity
|
||||||
|
|
||||||
|
// POST route without ID
|
||||||
|
postEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
params := map[string]string{
|
||||||
|
"schema": currentSchema,
|
||||||
|
"entity": currentEntity,
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.Handle(respAdapter, reqAdapter, params)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
r.Handle("POST", entityPath, wrapBunRouterHandler(postEntityHandler, authMiddleware))
|
||||||
|
|
||||||
|
// POST route with ID
|
||||||
|
postEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
params := map[string]string{
|
||||||
|
"schema": currentSchema,
|
||||||
|
"entity": currentEntity,
|
||||||
|
"id": req.Param("id"),
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.Handle(respAdapter, reqAdapter, params)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
r.Handle("POST", entityWithIDPath, wrapBunRouterHandler(postEntityWithIDHandler, authMiddleware))
|
||||||
|
|
||||||
|
// GET route without ID
|
||||||
|
getEntityHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
params := map[string]string{
|
||||||
|
"schema": currentSchema,
|
||||||
|
"entity": currentEntity,
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
r.Handle("GET", entityPath, wrapBunRouterHandler(getEntityHandler, authMiddleware))
|
||||||
|
|
||||||
|
// GET route with ID
|
||||||
|
getEntityWithIDHandler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, corsConfig)
|
||||||
|
params := map[string]string{
|
||||||
|
"schema": currentSchema,
|
||||||
|
"entity": currentEntity,
|
||||||
|
"id": req.Param("id"),
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
r.Handle("GET", entityWithIDPath, wrapBunRouterHandler(getEntityWithIDHandler, authMiddleware))
|
||||||
|
|
||||||
|
// OPTIONS route without ID (returns metadata)
|
||||||
|
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||||
|
r.Handle("OPTIONS", entityPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
optionsCorsConfig := corsConfig
|
||||||
|
optionsCorsConfig.AllowedMethods = []string{"GET", "POST", "OPTIONS"}
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||||
|
params := map[string]string{
|
||||||
|
"schema": currentSchema,
|
||||||
|
"entity": currentEntity,
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
// OPTIONS route with ID (returns metadata)
|
||||||
|
// Don't apply auth middleware to OPTIONS - CORS preflight must not require auth
|
||||||
|
r.Handle("OPTIONS", entityWithIDPath, func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||||
|
respAdapter := router.NewHTTPResponseWriter(w)
|
||||||
|
reqAdapter := router.NewHTTPRequest(req.Request)
|
||||||
|
optionsCorsConfig := corsConfig
|
||||||
|
optionsCorsConfig.AllowedMethods = []string{"POST", "OPTIONS"}
|
||||||
|
common.SetCORSHeaders(respAdapter, reqAdapter, optionsCorsConfig)
|
||||||
|
params := map[string]string{
|
||||||
|
"schema": currentSchema,
|
||||||
|
"entity": currentEntity,
|
||||||
|
}
|
||||||
|
|
||||||
|
handler.HandleGet(respAdapter, reqAdapter, params)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleWithBunRouter shows how to use bunrouter from uptrace
|
||||||
|
func ExampleWithBunRouter(bunDB *bun.DB) {
|
||||||
|
// Create handler with Bun adapter
|
||||||
|
handler := NewHandlerWithBun(bunDB)
|
||||||
|
|
||||||
|
// Create bunrouter
|
||||||
|
bunRouter := bunrouter.New()
|
||||||
|
|
||||||
|
// Setup ResolveSpec routes with bunrouter without authentication
|
||||||
|
SetupBunRouterRoutes(bunRouter, handler, nil)
|
||||||
|
|
||||||
|
// Start server
|
||||||
|
// http.ListenAndServe(":8080", bunRouter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleBunRouterWithBunDB shows the full uptrace stack (bunrouter + Bun ORM)
|
||||||
|
func ExampleBunRouterWithBunDB(bunDB *bun.DB) {
|
||||||
|
// Create Bun database adapter
|
||||||
|
dbAdapter := database.NewBunAdapter(bunDB)
|
||||||
|
|
||||||
|
// Create model registry
|
||||||
|
registry := modelregistry.NewModelRegistry()
|
||||||
|
// registry.RegisterModel("public.users", &User{})
|
||||||
|
|
||||||
|
// Create handler with Bun
|
||||||
|
handler := NewHandler(dbAdapter, registry)
|
||||||
|
|
||||||
|
// Create bunrouter
|
||||||
|
bunRouter := bunrouter.New()
|
||||||
|
|
||||||
|
// Setup ResolveSpec routes without authentication
|
||||||
|
SetupBunRouterRoutes(bunRouter, handler, nil)
|
||||||
|
|
||||||
|
// This gives you the full uptrace stack: bunrouter + Bun ORM
|
||||||
|
// http.ListenAndServe(":8080", bunRouter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleBunRouterWithGroup shows how to use SetupBunRouterRoutes with a bunrouter.Group
|
||||||
|
func ExampleBunRouterWithGroup(bunDB *bun.DB) {
|
||||||
|
// Create handler with Bun adapter
|
||||||
|
handler := NewHandlerWithBun(bunDB)
|
||||||
|
|
||||||
|
// Create bunrouter
|
||||||
|
bunRouter := bunrouter.New()
|
||||||
|
|
||||||
|
// Create a route group with a prefix
|
||||||
|
apiGroup := bunRouter.NewGroup("/api")
|
||||||
|
|
||||||
|
// Setup ResolveSpec routes on the group - routes will be under /api
|
||||||
|
SetupBunRouterRoutes(apiGroup, handler, nil)
|
||||||
|
|
||||||
|
// Start server
|
||||||
|
// http.ListenAndServe(":8080", bunRouter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleWithGORMAndAuth shows how to use ResolveSpec with GORM and authentication
|
||||||
|
func ExampleWithGORMAndAuth(db *gorm.DB) {
|
||||||
|
// Create handler using GORM
|
||||||
|
_ = NewHandlerWithGORM(db)
|
||||||
|
|
||||||
|
// Create auth middleware
|
||||||
|
// import "github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
// secList := security.NewSecurityList(myProvider)
|
||||||
|
// authMiddleware := func(h http.Handler) http.Handler {
|
||||||
|
// return security.NewAuthHandler(secList, h)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Setup router with authentication
|
||||||
|
_ = mux.NewRouter()
|
||||||
|
// SetupMuxRoutes(muxRouter, handler, authMiddleware)
|
||||||
|
|
||||||
|
// Register models
|
||||||
|
// handler.RegisterModel("public", "users", &User{})
|
||||||
|
|
||||||
|
// Start server
|
||||||
|
// http.ListenAndServe(":8080", muxRouter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleWithBunAndAuth shows how to use ResolveSpec with Bun and authentication
|
||||||
|
func ExampleWithBunAndAuth(bunDB *bun.DB) {
|
||||||
|
// Create Bun adapter
|
||||||
|
dbAdapter := database.NewBunAdapter(bunDB)
|
||||||
|
|
||||||
|
// Create model registry
|
||||||
|
registry := modelregistry.NewModelRegistry()
|
||||||
|
// registry.RegisterModel("public.users", &User{})
|
||||||
|
|
||||||
|
// Create handler
|
||||||
|
_ = NewHandler(dbAdapter, registry)
|
||||||
|
|
||||||
|
// Create auth middleware
|
||||||
|
// import "github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
// secList := security.NewSecurityList(myProvider)
|
||||||
|
// authMiddleware := func(h http.Handler) http.Handler {
|
||||||
|
// return security.NewAuthHandler(secList, h)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Setup routes with authentication
|
||||||
|
_ = mux.NewRouter()
|
||||||
|
// SetupMuxRoutes(muxRouter, handler, authMiddleware)
|
||||||
|
|
||||||
|
// Start server
|
||||||
|
// http.ListenAndServe(":8080", muxRouter)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExampleBunRouterWithBunDBAndAuth shows the full uptrace stack with authentication
|
||||||
|
func ExampleBunRouterWithBunDBAndAuth(bunDB *bun.DB) {
|
||||||
|
// Create Bun database adapter
|
||||||
|
dbAdapter := database.NewBunAdapter(bunDB)
|
||||||
|
|
||||||
|
// Create model registry
|
||||||
|
registry := modelregistry.NewModelRegistry()
|
||||||
|
// registry.RegisterModel("public.users", &User{})
|
||||||
|
|
||||||
|
// Create handler with Bun
|
||||||
|
_ = NewHandler(dbAdapter, registry)
|
||||||
|
|
||||||
|
// Create auth middleware
|
||||||
|
// import "github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
// secList := security.NewSecurityList(myProvider)
|
||||||
|
// authMiddleware := func(h http.Handler) http.Handler {
|
||||||
|
// return security.NewAuthHandler(secList, h)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Create bunrouter
|
||||||
|
_ = bunrouter.New()
|
||||||
|
|
||||||
|
// Setup ResolveSpec routes with authentication
|
||||||
|
// SetupBunRouterRoutes(bunRouter, handler, authMiddleware)
|
||||||
|
|
||||||
|
// This gives you the full uptrace stack: bunrouter + Bun ORM with authentication
|
||||||
|
// http.ListenAndServe(":8080", bunRouter)
|
||||||
|
}
|
||||||
+109
@@ -0,0 +1,109 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RegisterSecurityHooks registers all security-related hooks with the handler
|
||||||
|
func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) {
|
||||||
|
// Hook 0: BeforeHandle - enforce auth after model resolution
|
||||||
|
handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error {
|
||||||
|
if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil {
|
||||||
|
hookCtx.Abort = true
|
||||||
|
hookCtx.AbortMessage = err.Error()
|
||||||
|
hookCtx.AbortCode = http.StatusUnauthorized
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hook 1: BeforeRead - Load security rules
|
||||||
|
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
||||||
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
return security.LoadSecurityRules(secCtx, securityList)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hook 2: BeforeScan - Apply row-level security filters
|
||||||
|
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
||||||
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
return security.ApplyRowSecurity(secCtx, securityList)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hook 3: AfterRead - Apply column-level security (masking)
|
||||||
|
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||||
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
return security.ApplyColumnSecurity(secCtx, securityList)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hook 4 (Optional): Audit logging
|
||||||
|
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||||
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
return security.LogDataAccess(secCtx)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hook 5: BeforeUpdate - enforce CanUpdate rule from context/registry
|
||||||
|
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||||
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
return security.CheckModelUpdateAllowed(secCtx)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hook 6: BeforeDelete - enforce CanDelete rule from context/registry
|
||||||
|
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
|
||||||
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
return security.CheckModelDeleteAllowed(secCtx)
|
||||||
|
})
|
||||||
|
|
||||||
|
logger.Info("Security hooks registered for resolvespec handler")
|
||||||
|
}
|
||||||
|
|
||||||
|
// securityContext adapts resolvespec.HookContext to security.SecurityContext interface
|
||||||
|
type securityContext struct {
|
||||||
|
ctx *HookContext
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSecurityContext(ctx *HookContext) security.SecurityContext {
|
||||||
|
return &securityContext{ctx: ctx}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetContext() context.Context {
|
||||||
|
return s.ctx.Context
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetUserID() (int, bool) {
|
||||||
|
return security.GetUserID(s.ctx.Context)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetSchema() string {
|
||||||
|
return s.ctx.Schema
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetEntity() string {
|
||||||
|
return s.ctx.Entity
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetModel() interface{} {
|
||||||
|
return s.ctx.Model
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetQuery() interface{} {
|
||||||
|
return s.ctx.Query
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) SetQuery(query interface{}) {
|
||||||
|
if q, ok := query.(common.SelectQuery); ok {
|
||||||
|
s.ctx.Query = q
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) GetResult() interface{} {
|
||||||
|
return s.ctx.Result
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *securityContext) SetResult(result interface{}) {
|
||||||
|
s.ctx.Result = result
|
||||||
|
}
|
||||||
+153
@@ -0,0 +1,153 @@
|
|||||||
|
# Keystore
|
||||||
|
|
||||||
|
Per-user named auth keys with pluggable storage. Each user can hold multiple keys of different types — JWT secrets, header API keys, OAuth2 client credentials, or generic API keys. Keys are identified by a human-readable name ("CI deploy", "mobile app") and can carry scopes and arbitrary metadata.
|
||||||
|
|
||||||
|
## Key types
|
||||||
|
|
||||||
|
| Constant | Value | Use case |
|
||||||
|
|---|---|---|
|
||||||
|
| `KeyTypeJWTSecret` | `jwt_secret` | Per-user JWT signing secret |
|
||||||
|
| `KeyTypeHeaderAPI` | `header_api` | Static API key sent in a request header |
|
||||||
|
| `KeyTypeOAuth2` | `oauth2` | OAuth2 client credentials |
|
||||||
|
| `KeyTypeGenericAPI` | `api` | General-purpose application key |
|
||||||
|
|
||||||
|
## Storage backends
|
||||||
|
|
||||||
|
### ConfigKeyStore
|
||||||
|
|
||||||
|
In-memory store seeded from a static list. Suitable for a small, fixed set of service-account keys loaded from a config file. Keys created at runtime via `CreateKey` are held in memory and lost on restart.
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Pre-load keys from config (KeyHash = SHA-256 hex of the raw key)
|
||||||
|
store := security.NewConfigKeyStore([]security.UserKey{
|
||||||
|
{
|
||||||
|
UserID: 1,
|
||||||
|
KeyType: security.KeyTypeGenericAPI,
|
||||||
|
KeyHash: "e3b0c44298fc1c149afb...", // sha256(rawKey)
|
||||||
|
Name: "CI deploy",
|
||||||
|
Scopes: []string{"deploy"},
|
||||||
|
IsActive: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### DatabaseKeyStore
|
||||||
|
|
||||||
|
Backed by PostgreSQL stored procedures. Supports optional caching (default 2-minute TTL). Apply `keystore_schema.sql` before use.
|
||||||
|
|
||||||
|
```go
|
||||||
|
db, _ := sql.Open("postgres", dsn)
|
||||||
|
|
||||||
|
store := security.NewDatabaseKeyStore(db)
|
||||||
|
|
||||||
|
// With options
|
||||||
|
store = security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
||||||
|
CacheTTL: 5 * time.Minute,
|
||||||
|
SQLNames: &security.KeyStoreSQLNames{
|
||||||
|
ValidateKey: "myapp_keystore_validate", // override one procedure name
|
||||||
|
},
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Managing keys
|
||||||
|
|
||||||
|
```go
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Create — raw key returned once; store it securely
|
||||||
|
resp, err := store.CreateKey(ctx, security.CreateKeyRequest{
|
||||||
|
UserID: 42,
|
||||||
|
KeyType: security.KeyTypeGenericAPI,
|
||||||
|
Name: "mobile app",
|
||||||
|
Scopes: []string{"read", "write"},
|
||||||
|
})
|
||||||
|
fmt.Println(resp.RawKey) // only shown here; hashed internally
|
||||||
|
|
||||||
|
// List
|
||||||
|
keys, err := store.GetUserKeys(ctx, 42, "") // "" = all types
|
||||||
|
keys, err = store.GetUserKeys(ctx, 42, security.KeyTypeGenericAPI)
|
||||||
|
|
||||||
|
// Revoke
|
||||||
|
err = store.DeleteKey(ctx, 42, resp.Key.ID)
|
||||||
|
|
||||||
|
// Validate (used by authenticators internally)
|
||||||
|
key, err := store.ValidateKey(ctx, rawKey, "")
|
||||||
|
```
|
||||||
|
|
||||||
|
## HTTP authentication
|
||||||
|
|
||||||
|
`KeyStoreAuthenticator` wraps any `KeyStore` and implements the `Authenticator` interface. It is drop-in compatible with `DatabaseAuthenticator` and works in `CompositeSecurityProvider`.
|
||||||
|
|
||||||
|
Keys are extracted from the request in this order:
|
||||||
|
|
||||||
|
1. `Authorization: Bearer <key>`
|
||||||
|
2. `Authorization: ApiKey <key>`
|
||||||
|
3. `X-API-Key: <key>`
|
||||||
|
|
||||||
|
```go
|
||||||
|
auth := security.NewKeyStoreAuthenticator(store, "") // "" = accept any key type
|
||||||
|
// Restrict to a specific type:
|
||||||
|
auth = security.NewKeyStoreAuthenticator(store, security.KeyTypeGenericAPI)
|
||||||
|
```
|
||||||
|
|
||||||
|
Plug it into a handler:
|
||||||
|
|
||||||
|
```go
|
||||||
|
handler := resolvespec.NewHandler(db, registry,
|
||||||
|
resolvespec.WithAuthenticator(auth),
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
`Login` and `Logout` return an error — key lifecycle is managed through `KeyStore` directly.
|
||||||
|
|
||||||
|
On successful validation the request context receives a `UserContext` where:
|
||||||
|
|
||||||
|
- `UserID` — from the key
|
||||||
|
- `Roles` — the key's `Scopes`
|
||||||
|
- `Claims["key_type"]` — key type string
|
||||||
|
- `Claims["key_name"]` — key name
|
||||||
|
|
||||||
|
## Database setup
|
||||||
|
|
||||||
|
Apply `keystore_schema.sql` to your PostgreSQL database. It requires the `users` table from the main `database_schema.sql`.
|
||||||
|
|
||||||
|
```sql
|
||||||
|
\i pkg/security/keystore_schema.sql
|
||||||
|
```
|
||||||
|
|
||||||
|
This creates:
|
||||||
|
|
||||||
|
- `user_keys` table with indexes on `user_id`, `key_hash`, and `key_type`
|
||||||
|
- `resolvespec_keystore_get_user_keys(p_user_id, p_key_type)`
|
||||||
|
- `resolvespec_keystore_create_key(p_request jsonb)`
|
||||||
|
- `resolvespec_keystore_delete_key(p_user_id, p_key_id)`
|
||||||
|
- `resolvespec_keystore_validate_key(p_key_hash, p_key_type)`
|
||||||
|
|
||||||
|
### Custom procedure names
|
||||||
|
|
||||||
|
```go
|
||||||
|
store := security.NewDatabaseKeyStore(db, security.DatabaseKeyStoreOptions{
|
||||||
|
SQLNames: &security.KeyStoreSQLNames{
|
||||||
|
GetUserKeys: "myschema_get_keys",
|
||||||
|
CreateKey: "myschema_create_key",
|
||||||
|
DeleteKey: "myschema_delete_key",
|
||||||
|
ValidateKey: "myschema_validate_key",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
// Validate names at startup
|
||||||
|
names := &security.KeyStoreSQLNames{
|
||||||
|
GetUserKeys: "myschema_get_keys",
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
if err := security.ValidateKeyStoreSQLNames(names); err != nil {
|
||||||
|
log.Fatal(err)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Security notes
|
||||||
|
|
||||||
|
- Raw keys are never stored. Only the SHA-256 hex digest is persisted.
|
||||||
|
- The raw key is generated with `crypto/rand` (32 bytes, base64url-encoded) and returned exactly once in `CreateKeyResponse.RawKey`.
|
||||||
|
- Hash comparisons in `ConfigKeyStore` use `crypto/subtle.ConstantTimeCompare` to prevent timing side-channels.
|
||||||
|
- `DeleteKey` performs a soft delete (`is_active = false`). The `DatabaseKeyStore` invalidates the cache entry immediately, but due to the cache TTL a revoked key may authenticate for up to `CacheTTL` (default 2 minutes) in a distributed environment. Set `CacheTTL: 0` to disable caching if immediate revocation is required.
|
||||||
+527
@@ -0,0 +1,527 @@
|
|||||||
|
# OAuth2 Authentication Guide
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
The security package provides OAuth2 authentication support for any OAuth2-compliant provider including Google, GitHub, Microsoft, Facebook, and custom providers.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- **Universal OAuth2 Support**: Works with any OAuth2 provider
|
||||||
|
- **Pre-configured Providers**: Google, GitHub, Microsoft, Facebook
|
||||||
|
- **Multi-Provider Support**: Use all OAuth2 providers simultaneously
|
||||||
|
- **Custom Providers**: Easy configuration for any OAuth2 service
|
||||||
|
- **Session Management**: Database-backed session storage
|
||||||
|
- **Token Refresh**: Automatic token refresh support
|
||||||
|
- **State Validation**: Built-in CSRF protection
|
||||||
|
- **User Auto-Creation**: Automatically creates users on first login
|
||||||
|
- **Unified Authentication**: OAuth2 and traditional auth share same session storage
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
### 1. Database Setup
|
||||||
|
|
||||||
|
```sql
|
||||||
|
-- Run the schema from database_schema.sql
|
||||||
|
CREATE TABLE IF NOT EXISTS users (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
username VARCHAR(255) NOT NULL UNIQUE,
|
||||||
|
email VARCHAR(255) NOT NULL UNIQUE,
|
||||||
|
password VARCHAR(255),
|
||||||
|
user_level INTEGER DEFAULT 0,
|
||||||
|
roles VARCHAR(500),
|
||||||
|
is_active BOOLEAN DEFAULT true,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_login_at TIMESTAMP,
|
||||||
|
remote_id VARCHAR(255),
|
||||||
|
auth_provider VARCHAR(50)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||||
|
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_activity_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
ip_address VARCHAR(45),
|
||||||
|
user_agent TEXT,
|
||||||
|
access_token TEXT,
|
||||||
|
refresh_token TEXT,
|
||||||
|
token_type VARCHAR(50) DEFAULT 'Bearer',
|
||||||
|
auth_provider VARCHAR(50)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- OAuth2 stored procedures (7 functions)
|
||||||
|
-- See database_schema.sql for full implementation
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. Google OAuth2
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
|
||||||
|
// Create authenticator
|
||||||
|
oauth2Auth := security.NewGoogleAuthenticator(
|
||||||
|
"your-google-client-id",
|
||||||
|
"your-google-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Login route - redirects to Google
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := oauth2Auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := oauth2Auth.OAuth2GetAuthURL(state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Callback route - handles Google response
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := oauth2Auth.OAuth2HandleCallback(r.Context(), code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set session cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.Redirect(w, r, "/dashboard", http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. GitHub OAuth2
|
||||||
|
|
||||||
|
```go
|
||||||
|
oauth2Auth := security.NewGitHubAuthenticator(
|
||||||
|
"your-github-client-id",
|
||||||
|
"your-github-client-secret",
|
||||||
|
"http://localhost:8080/auth/github/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Same routes pattern as Google
|
||||||
|
router.HandleFunc("/auth/github/login", ...)
|
||||||
|
router.HandleFunc("/auth/github/callback", ...)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4. Microsoft OAuth2
|
||||||
|
|
||||||
|
```go
|
||||||
|
oauth2Auth := security.NewMicrosoftAuthenticator(
|
||||||
|
"your-microsoft-client-id",
|
||||||
|
"your-microsoft-client-secret",
|
||||||
|
"http://localhost:8080/auth/microsoft/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 5. Facebook OAuth2
|
||||||
|
|
||||||
|
```go
|
||||||
|
oauth2Auth := security.NewFacebookAuthenticator(
|
||||||
|
"your-facebook-client-id",
|
||||||
|
"your-facebook-client-secret",
|
||||||
|
"http://localhost:8080/auth/facebook/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Custom OAuth2 Provider
|
||||||
|
|
||||||
|
```go
|
||||||
|
oauth2Auth := security.NewDatabaseAuthenticator(db).WithOAuth2(security.OAuth2Config{
|
||||||
|
ClientID: "your-client-id",
|
||||||
|
ClientSecret: "your-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://your-provider.com/oauth/authorize",
|
||||||
|
TokenURL: "https://your-provider.com/oauth/token",
|
||||||
|
UserInfoURL: "https://your-provider.com/oauth/userinfo",
|
||||||
|
DB: db,
|
||||||
|
ProviderName: "custom",
|
||||||
|
|
||||||
|
// Optional: Custom user info parser
|
||||||
|
UserInfoParser: func(userInfo map[string]any) (*security.UserContext, error) {
|
||||||
|
return &security.UserContext{
|
||||||
|
UserName: userInfo["username"].(string),
|
||||||
|
Email: userInfo["email"].(string),
|
||||||
|
RemoteID: userInfo["id"].(string),
|
||||||
|
UserLevel: 1,
|
||||||
|
Roles: []string{"user"},
|
||||||
|
Claims: userInfo,
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Protected Routes
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Create security provider
|
||||||
|
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||||
|
rowSec := security.NewDatabaseRowSecurityProvider(db)
|
||||||
|
provider, _ := security.NewCompositeSecurityProvider(oauth2Auth, colSec, rowSec)
|
||||||
|
securityList, _ := security.NewSecurityList(provider)
|
||||||
|
|
||||||
|
// Apply middleware to protected routes
|
||||||
|
protectedRouter := router.PathPrefix("/api").Subrouter()
|
||||||
|
protectedRouter.Use(security.NewAuthMiddleware(securityList))
|
||||||
|
protectedRouter.Use(security.SetSecurityMiddleware(securityList))
|
||||||
|
|
||||||
|
protectedRouter.HandleFunc("/profile", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := security.GetUserContext(r.Context())
|
||||||
|
json.NewEncoder(w).Encode(userCtx)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Token Refresh
|
||||||
|
|
||||||
|
OAuth2 access tokens expire after a period of time. Use the refresh token to obtain a new access token without requiring the user to log in again.
|
||||||
|
|
||||||
|
```go
|
||||||
|
router.HandleFunc("/auth/refresh", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
Provider string `json:"provider"` // "google", "github", etc.
|
||||||
|
}
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Default to google if not specified
|
||||||
|
if req.Provider == "" {
|
||||||
|
req.Provider = "google"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use OAuth2-specific refresh method
|
||||||
|
loginResp, err := oauth2Auth.OAuth2RefreshToken(r.Context(), req.RefreshToken, req.Provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set new session cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
**Important Notes:**
|
||||||
|
- The refresh token is returned in the `LoginResponse.RefreshToken` field after successful OAuth2 callback
|
||||||
|
- Store the refresh token securely on the client side
|
||||||
|
- Each provider must be configured with the appropriate scopes to receive a refresh token (e.g., `access_type=offline` for Google)
|
||||||
|
- The `OAuth2RefreshToken` method requires the provider name to identify which OAuth2 provider to use for refreshing
|
||||||
|
|
||||||
|
## Logout
|
||||||
|
|
||||||
|
```go
|
||||||
|
router.HandleFunc("/auth/logout", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := security.GetUserContext(r.Context())
|
||||||
|
|
||||||
|
oauth2Auth.Logout(r.Context(), security.LogoutRequest{
|
||||||
|
Token: userCtx.SessionID,
|
||||||
|
UserID: userCtx.UserID,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: "",
|
||||||
|
MaxAge: -1,
|
||||||
|
})
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Multi-Provider Setup
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Single DatabaseAuthenticator with ALL OAuth2 providers
|
||||||
|
auth := security.NewDatabaseAuthenticator(db).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ClientID: "google-client-id",
|
||||||
|
ClientSecret: "google-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/google/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://accounts.google.com/o/oauth2/auth",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
||||||
|
ProviderName: "google",
|
||||||
|
}).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ClientID: "github-client-id",
|
||||||
|
ClientSecret: "github-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/github/callback",
|
||||||
|
Scopes: []string{"user:email"},
|
||||||
|
AuthURL: "https://github.com/login/oauth/authorize",
|
||||||
|
TokenURL: "https://github.com/login/oauth/access_token",
|
||||||
|
UserInfoURL: "https://api.github.com/user",
|
||||||
|
ProviderName: "github",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Get list of configured providers
|
||||||
|
providers := auth.OAuth2GetProviders() // ["google", "github"]
|
||||||
|
|
||||||
|
// Google routes
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("google", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, err := auth.OAuth2HandleCallback(r.Context(), "google",
|
||||||
|
r.URL.Query().Get("code"), r.URL.Query().Get("state"))
|
||||||
|
// ... handle response
|
||||||
|
})
|
||||||
|
|
||||||
|
// GitHub routes
|
||||||
|
router.HandleFunc("/auth/github/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("github", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/github/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, err := auth.OAuth2HandleCallback(r.Context(), "github",
|
||||||
|
r.URL.Query().Get("code"), r.URL.Query().Get("state"))
|
||||||
|
// ... handle response
|
||||||
|
})
|
||||||
|
|
||||||
|
// Use same authenticator for protected routes - works for ALL providers
|
||||||
|
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||||
|
securityList, _ := security.NewSecurityList(provider)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration Options
|
||||||
|
|
||||||
|
### OAuth2Config Fields
|
||||||
|
|
||||||
|
| Field | Type | Description |
|
||||||
|
|-------|------|-------------|
|
||||||
|
| ClientID | string | OAuth2 client ID from provider |
|
||||||
|
| ClientSecret | string | OAuth2 client secret |
|
||||||
|
| RedirectURL | string | Callback URL registered with provider |
|
||||||
|
| Scopes | []string | OAuth2 scopes to request |
|
||||||
|
| AuthURL | string | Provider's authorization endpoint |
|
||||||
|
| TokenURL | string | Provider's token endpoint |
|
||||||
|
| UserInfoURL | string | Provider's user info endpoint |
|
||||||
|
| DB | *sql.DB | Database connection for sessions |
|
||||||
|
| UserInfoParser | func | Custom parser for user info (optional) |
|
||||||
|
| StateValidator | func | Custom state validator (optional) |
|
||||||
|
| ProviderName | string | Provider name for logging (optional) |
|
||||||
|
|
||||||
|
## User Info Parsing
|
||||||
|
|
||||||
|
The default parser extracts these standard fields:
|
||||||
|
- `sub` → RemoteID
|
||||||
|
- `email` → Email, UserName
|
||||||
|
- `name` → UserName
|
||||||
|
- `login` → UserName (GitHub)
|
||||||
|
|
||||||
|
Custom parser example:
|
||||||
|
|
||||||
|
```go
|
||||||
|
UserInfoParser: func(userInfo map[string]any) (*security.UserContext, error) {
|
||||||
|
// Extract custom fields
|
||||||
|
ctx := &security.UserContext{
|
||||||
|
UserName: userInfo["preferred_username"].(string),
|
||||||
|
Email: userInfo["email"].(string),
|
||||||
|
RemoteID: userInfo["sub"].(string),
|
||||||
|
UserLevel: 1,
|
||||||
|
Roles: []string{"user"},
|
||||||
|
Claims: userInfo, // Store all claims
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add custom roles based on provider data
|
||||||
|
if groups, ok := userInfo["groups"].([]interface{}); ok {
|
||||||
|
for _, g := range groups {
|
||||||
|
ctx.Roles = append(ctx.Roles, g.(string))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return ctx, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Security Best Practices
|
||||||
|
|
||||||
|
1. **Always use HTTPS in production**
|
||||||
|
```go
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Secure: true, // Only send over HTTPS
|
||||||
|
HttpOnly: true, // Prevent XSS access
|
||||||
|
SameSite: http.SameSiteLaxMode, // CSRF protection
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
2. **Store secrets securely**
|
||||||
|
```go
|
||||||
|
clientID := os.Getenv("GOOGLE_CLIENT_ID")
|
||||||
|
clientSecret := os.Getenv("GOOGLE_CLIENT_SECRET")
|
||||||
|
```
|
||||||
|
|
||||||
|
3. **Validate redirect URLs**
|
||||||
|
- Only register trusted redirect URLs with OAuth2 providers
|
||||||
|
- Never accept redirect URL from request parameters
|
||||||
|
|
||||||
|
5. **Session expiration**
|
||||||
|
- OAuth2 sessions automatically expire based on token expiry
|
||||||
|
- Clean up expired sessions periodically:
|
||||||
|
```sql
|
||||||
|
DELETE FROM user_sessions WHERE expires_at < NOW();
|
||||||
|
```
|
||||||
|
|
||||||
|
4. **State parameter**
|
||||||
|
- Automatically generated with cryptographic randomness
|
||||||
|
- One-time use and expires after 10 minutes
|
||||||
|
- Prevents CSRF attacks
|
||||||
|
|
||||||
|
## Implementation Details
|
||||||
|
|
||||||
|
All database operations use stored procedures for consistency and security:
|
||||||
|
- `resolvespec_oauth_getorcreateuser` - Find or create OAuth2 user
|
||||||
|
- `resolvespec_oauth_createsession` - Create OAuth2 session
|
||||||
|
- `resolvespec_oauth_getsession` - Validate and retrieve session
|
||||||
|
- `resolvespec_oauth_deletesession` - Logout/delete session
|
||||||
|
- `resolvespec_oauth_getrefreshtoken` - Get session by refresh token
|
||||||
|
- `resolvespec_oauth_updaterefreshtoken` - Update tokens after refresh
|
||||||
|
- `resolvespec_oauth_getuser` - Get user data by ID
|
||||||
|
|
||||||
|
## Provider Setup Guides
|
||||||
|
|
||||||
|
### Google
|
||||||
|
|
||||||
|
1. Go to [Google Cloud Console](https://console.cloud.google.com/)
|
||||||
|
2. Create a new project or select existing
|
||||||
|
3. Enable Google+ API
|
||||||
|
4. Create OAuth 2.0 credentials
|
||||||
|
5. Add authorized redirect URI: `http://localhost:8080/auth/google/callback`
|
||||||
|
6. Copy Client ID and Client Secret
|
||||||
|
|
||||||
|
### GitHub
|
||||||
|
|
||||||
|
1. Go to [GitHub Developer Settings](https://github.com/settings/developers)
|
||||||
|
2. Click "New OAuth App"
|
||||||
|
3. Set Homepage URL: `http://localhost:8080`
|
||||||
|
4. Set Authorization callback URL: `http://localhost:8080/auth/github/callback`
|
||||||
|
5. Copy Client ID and Client Secret
|
||||||
|
|
||||||
|
### Microsoft
|
||||||
|
|
||||||
|
1. Go to [Azure Portal](https://portal.azure.com/)
|
||||||
|
2. Register new application in Azure AD
|
||||||
|
3. Add redirect URI: `http://localhost:8080/auth/microsoft/callback`
|
||||||
|
4. Create client secret
|
||||||
|
5. Copy Application (client) ID and secret value
|
||||||
|
|
||||||
|
### Facebook
|
||||||
|
|
||||||
|
1. Go to [Facebook Developers](https://developers.facebook.com/)
|
||||||
|
2. Create new app
|
||||||
|
3. Add Facebook Login product
|
||||||
|
4. Set Valid OAuth Redirect URIs: `http://localhost:8080/auth/facebook/callback`
|
||||||
|
5. Copy App ID and App Secret
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### "redirect_uri_mismatch" error
|
||||||
|
- Ensure the redirect URL in code matches exactly with provider configuration
|
||||||
|
- Include protocol (http/https), domain, port, and path
|
||||||
|
|
||||||
|
### "invalid_client" error
|
||||||
|
- Verify Client ID and Client Secret are correct
|
||||||
|
- Check if credentials are for the correct environment (dev/prod)
|
||||||
|
|
||||||
|
### "invalid_grant" error during token exchange
|
||||||
|
- State parameter validation failed
|
||||||
|
- Token might have expired
|
||||||
|
- Check server time synchronization
|
||||||
|
|
||||||
|
### User not created after successful OAuth2 login
|
||||||
|
- Check database constraints (username/email unique)
|
||||||
|
- Verify UserInfoParser is extracting required fields
|
||||||
|
- Check database logs for constraint violations
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
```go
|
||||||
|
func TestOAuth2Flow(t *testing.T) {
|
||||||
|
// Mock database
|
||||||
|
db, mock, _ := sqlmock.New()
|
||||||
|
|
||||||
|
oauth2Auth := security.NewGoogleAuthenticator(
|
||||||
|
"test-client-id",
|
||||||
|
"test-client-secret",
|
||||||
|
"http://localhost/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Test state generation
|
||||||
|
state, err := oauth2Auth.GenerateState()
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, state)
|
||||||
|
|
||||||
|
// Test auth URL generation
|
||||||
|
authURL := oauth2Auth.GetAuthURL(state)
|
||||||
|
assert.Contains(t, authURL, "accounts.google.com")
|
||||||
|
assert.Contains(t, authURL, state)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## API Reference
|
||||||
|
|
||||||
|
### DatabaseAuthenticator with OAuth2
|
||||||
|
|
||||||
|
| Method | Description |
|
||||||
|
|--------|-------------|
|
||||||
|
| WithOAuth2(cfg) | Adds OAuth2 provider (can be called multiple times, returns *DatabaseAuthenticator) |
|
||||||
|
| OAuth2GetAuthURL(provider, state) | Returns OAuth2 authorization URL for specified provider |
|
||||||
|
| OAuth2GenerateState() | Generates random state for CSRF protection |
|
||||||
|
| OAuth2HandleCallback(ctx, provider, code, state) | Exchanges code for token and creates session |
|
||||||
|
| OAuth2RefreshToken(ctx, refreshToken, provider) | Refreshes expired access token using refresh token |
|
||||||
|
| OAuth2GetProviders() | Returns list of configured OAuth2 provider names |
|
||||||
|
| Login(ctx, req) | Standard username/password login |
|
||||||
|
| Logout(ctx, req) | Invalidates session (works for both OAuth2 and regular sessions) |
|
||||||
|
| Authenticate(r) | Validates session token from request (works for both OAuth2 and regular sessions) |
|
||||||
|
|
||||||
|
### Pre-configured Constructors
|
||||||
|
|
||||||
|
- `NewGoogleAuthenticator(clientID, secret, redirectURL, db)` - Single provider
|
||||||
|
- `NewGitHubAuthenticator(clientID, secret, redirectURL, db)` - Single provider
|
||||||
|
- `NewMicrosoftAuthenticator(clientID, secret, redirectURL, db)` - Single provider
|
||||||
|
- `NewFacebookAuthenticator(clientID, secret, redirectURL, db)` - Single provider
|
||||||
|
- `NewMultiProviderAuthenticator(db, configs)` - Multiple providers at once
|
||||||
|
|
||||||
|
All return `*DatabaseAuthenticator` with OAuth2 pre-configured.
|
||||||
|
|
||||||
|
For multiple providers, use `WithOAuth2()` multiple times or `NewMultiProviderAuthenticator()`.
|
||||||
|
|
||||||
|
## Examples
|
||||||
|
|
||||||
|
Complete working examples available in `oauth2_examples.go`:
|
||||||
|
- Basic Google OAuth2
|
||||||
|
- GitHub OAuth2
|
||||||
|
- Custom provider
|
||||||
|
- Multi-provider setup
|
||||||
|
- Token refresh
|
||||||
|
- Logout flow
|
||||||
|
- Complete integration with security middleware
|
||||||
Generated
Vendored
+281
@@ -0,0 +1,281 @@
|
|||||||
|
# OAuth2 Refresh Token - Quick Reference
|
||||||
|
|
||||||
|
## Quick Setup (3 Steps)
|
||||||
|
|
||||||
|
### 1. Initialize Authenticator
|
||||||
|
```go
|
||||||
|
auth := security.NewGoogleAuthenticator(
|
||||||
|
"client-id",
|
||||||
|
"client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. OAuth2 Login Flow
|
||||||
|
```go
|
||||||
|
// Login - Redirect to Google
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("google", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Callback - Store tokens
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, _ := auth.OAuth2HandleCallback(
|
||||||
|
r.Context(),
|
||||||
|
"google",
|
||||||
|
r.URL.Query().Get("code"),
|
||||||
|
r.URL.Query().Get("state"),
|
||||||
|
)
|
||||||
|
|
||||||
|
// Save refresh_token on client
|
||||||
|
// loginResp.RefreshToken - Store this securely!
|
||||||
|
// loginResp.Token - Session token for API calls
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. Refresh Endpoint
|
||||||
|
```go
|
||||||
|
router.HandleFunc("/auth/refresh", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
}
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Refresh token
|
||||||
|
loginResp, err := auth.OAuth2RefreshToken(r.Context(), req.RefreshToken, "google")
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), 401)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Multi-Provider Example
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Configure multiple providers
|
||||||
|
auth := security.NewDatabaseAuthenticator(db).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ProviderName: "google",
|
||||||
|
ClientID: "google-client-id",
|
||||||
|
ClientSecret: "google-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/google/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://accounts.google.com/o/oauth2/auth",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
||||||
|
}).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ProviderName: "github",
|
||||||
|
ClientID: "github-client-id",
|
||||||
|
ClientSecret: "github-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/github/callback",
|
||||||
|
Scopes: []string{"user:email"},
|
||||||
|
AuthURL: "https://github.com/login/oauth/authorize",
|
||||||
|
TokenURL: "https://github.com/login/oauth/access_token",
|
||||||
|
UserInfoURL: "https://api.github.com/user",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Refresh with provider selection
|
||||||
|
router.HandleFunc("/auth/refresh", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
Provider string `json:"provider"` // "google" or "github"
|
||||||
|
}
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
loginResp, err := auth.OAuth2RefreshToken(r.Context(), req.RefreshToken, req.Provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), 401)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Client-Side JavaScript
|
||||||
|
|
||||||
|
```javascript
|
||||||
|
// Automatic token refresh on 401
|
||||||
|
async function apiCall(url) {
|
||||||
|
let response = await fetch(url, {
|
||||||
|
headers: {
|
||||||
|
'Authorization': 'Bearer ' + localStorage.getItem('access_token')
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// Token expired - refresh it
|
||||||
|
if (response.status === 401) {
|
||||||
|
await refreshToken();
|
||||||
|
|
||||||
|
// Retry request with new token
|
||||||
|
response = await fetch(url, {
|
||||||
|
headers: {
|
||||||
|
'Authorization': 'Bearer ' + localStorage.getItem('access_token')
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
return response.json();
|
||||||
|
}
|
||||||
|
|
||||||
|
async function refreshToken() {
|
||||||
|
const response = await fetch('/auth/refresh', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({
|
||||||
|
refresh_token: localStorage.getItem('refresh_token'),
|
||||||
|
provider: localStorage.getItem('provider')
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.ok) {
|
||||||
|
const data = await response.json();
|
||||||
|
localStorage.setItem('access_token', data.token);
|
||||||
|
localStorage.setItem('refresh_token', data.refresh_token);
|
||||||
|
} else {
|
||||||
|
// Refresh failed - redirect to login
|
||||||
|
window.location.href = '/login';
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## API Methods
|
||||||
|
|
||||||
|
| Method | Parameters | Returns |
|
||||||
|
|--------|-----------|---------|
|
||||||
|
| `OAuth2RefreshToken` | `ctx, refreshToken, provider` | `*LoginResponse, error` |
|
||||||
|
| `OAuth2HandleCallback` | `ctx, provider, code, state` | `*LoginResponse, error` |
|
||||||
|
| `OAuth2GetAuthURL` | `provider, state` | `string, error` |
|
||||||
|
| `OAuth2GenerateState` | none | `string, error` |
|
||||||
|
| `OAuth2GetProviders` | none | `[]string` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## LoginResponse Structure
|
||||||
|
|
||||||
|
```go
|
||||||
|
type LoginResponse struct {
|
||||||
|
Token string // New session token for API calls
|
||||||
|
RefreshToken string // Refresh token (store securely)
|
||||||
|
User *UserContext // User information
|
||||||
|
ExpiresIn int64 // Seconds until token expires
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Database Stored Procedures
|
||||||
|
|
||||||
|
- `resolvespec_oauth_getrefreshtoken(refresh_token)` - Get session by refresh token
|
||||||
|
- `resolvespec_oauth_updaterefreshtoken(update_data)` - Update tokens after refresh
|
||||||
|
- `resolvespec_oauth_getuser(user_id)` - Get user data
|
||||||
|
|
||||||
|
All procedures return: `{p_success bool, p_error text, p_data jsonb}`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Common Errors
|
||||||
|
|
||||||
|
| Error | Cause | Solution |
|
||||||
|
|-------|-------|----------|
|
||||||
|
| `invalid or expired refresh token` | Token revoked/expired | Re-authenticate user |
|
||||||
|
| `OAuth2 provider 'xxx' not found` | Provider not configured | Add with `WithOAuth2()` |
|
||||||
|
| `failed to refresh token with provider` | Provider rejected request | Check credentials, re-auth user |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Security Checklist
|
||||||
|
|
||||||
|
- [ ] Use HTTPS for all OAuth2 endpoints
|
||||||
|
- [ ] Store refresh tokens securely (HttpOnly cookies or encrypted storage)
|
||||||
|
- [ ] Set cookie flags: `HttpOnly`, `Secure`, `SameSite=Strict`
|
||||||
|
- [ ] Implement rate limiting on refresh endpoint
|
||||||
|
- [ ] Log refresh attempts for audit
|
||||||
|
- [ ] Rotate tokens on refresh
|
||||||
|
- [ ] Revoke old sessions after successful refresh
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. Login and get refresh token
|
||||||
|
curl http://localhost:8080/auth/google/login
|
||||||
|
# Follow OAuth2 flow, get refresh_token from callback response
|
||||||
|
|
||||||
|
# 2. Refresh token
|
||||||
|
curl -X POST http://localhost:8080/auth/refresh \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"refresh_token":"ya29.xxx","provider":"google"}'
|
||||||
|
|
||||||
|
# 3. Use new token
|
||||||
|
curl http://localhost:8080/api/protected \
|
||||||
|
-H "Authorization: Bearer sess_abc123..."
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Pre-configured Providers
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Google
|
||||||
|
auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db)
|
||||||
|
|
||||||
|
// GitHub
|
||||||
|
auth := security.NewGitHubAuthenticator(clientID, secret, redirectURL, db)
|
||||||
|
|
||||||
|
// Microsoft
|
||||||
|
auth := security.NewMicrosoftAuthenticator(clientID, secret, redirectURL, db)
|
||||||
|
|
||||||
|
// Facebook
|
||||||
|
auth := security.NewFacebookAuthenticator(clientID, secret, redirectURL, db)
|
||||||
|
|
||||||
|
// All providers at once
|
||||||
|
auth := security.NewMultiProviderAuthenticator(db, map[string]security.OAuth2Config{
|
||||||
|
"google": {...},
|
||||||
|
"github": {...},
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Provider-Specific Notes
|
||||||
|
|
||||||
|
### Google
|
||||||
|
- Add `access_type=offline` to get refresh token
|
||||||
|
- Add `prompt=consent` to force consent screen
|
||||||
|
```go
|
||||||
|
authURL += "&access_type=offline&prompt=consent"
|
||||||
|
```
|
||||||
|
|
||||||
|
### GitHub
|
||||||
|
- Refresh tokens not always provided
|
||||||
|
- May need to request `offline_access` scope
|
||||||
|
|
||||||
|
### Microsoft
|
||||||
|
- Use `offline_access` scope for refresh token
|
||||||
|
|
||||||
|
### Facebook
|
||||||
|
- Tokens expire after 60 days by default
|
||||||
|
- Check app settings for token expiration policy
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Complete Example
|
||||||
|
|
||||||
|
See `/pkg/security/oauth2_examples.go` line 250 for full working example.
|
||||||
|
|
||||||
|
For detailed documentation see `/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md`.
|
||||||
Generated
Vendored
+495
@@ -0,0 +1,495 @@
|
|||||||
|
# OAuth2 Refresh Token Implementation
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
OAuth2 refresh token functionality is **fully implemented** in the ResolveSpec security package. This allows refreshing expired access tokens without requiring users to re-authenticate.
|
||||||
|
|
||||||
|
## Implementation Status: ✅ COMPLETE
|
||||||
|
|
||||||
|
### Components Implemented
|
||||||
|
|
||||||
|
1. **✅ Database Schema** - Tables and stored procedures
|
||||||
|
2. **✅ Go Methods** - OAuth2RefreshToken implementation
|
||||||
|
3. **✅ Thread Safety** - Mutex protection for provider map
|
||||||
|
4. **✅ Examples** - Working code examples
|
||||||
|
5. **✅ Documentation** - Complete API reference
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Database Schema
|
||||||
|
|
||||||
|
### Tables Modified
|
||||||
|
|
||||||
|
```sql
|
||||||
|
-- user_sessions table with OAuth2 token fields
|
||||||
|
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||||
|
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_activity_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
ip_address VARCHAR(45),
|
||||||
|
user_agent TEXT,
|
||||||
|
access_token TEXT, -- OAuth2 access token
|
||||||
|
refresh_token TEXT, -- OAuth2 refresh token
|
||||||
|
token_type VARCHAR(50), -- "Bearer", etc.
|
||||||
|
auth_provider VARCHAR(50) -- "google", "github", etc.
|
||||||
|
);
|
||||||
|
```
|
||||||
|
|
||||||
|
### Stored Procedures
|
||||||
|
|
||||||
|
**`resolvespec_oauth_getrefreshtoken(p_refresh_token)`**
|
||||||
|
- Gets OAuth2 session data by refresh token
|
||||||
|
- Returns: `{user_id, access_token, token_type, expiry}`
|
||||||
|
- Location: `database_schema.sql:714`
|
||||||
|
|
||||||
|
**`resolvespec_oauth_updaterefreshtoken(p_update_data)`**
|
||||||
|
- Updates session with new tokens after refresh
|
||||||
|
- Input: `{user_id, old_refresh_token, new_session_token, new_access_token, new_refresh_token, expires_at}`
|
||||||
|
- Location: `database_schema.sql:752`
|
||||||
|
|
||||||
|
**`resolvespec_oauth_getuser(p_user_id)`**
|
||||||
|
- Gets user data by ID for building UserContext
|
||||||
|
- Location: `database_schema.sql:791`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Go Implementation
|
||||||
|
|
||||||
|
### Method Signature
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (a *DatabaseAuthenticator) OAuth2RefreshToken(
|
||||||
|
ctx context.Context,
|
||||||
|
refreshToken string,
|
||||||
|
providerName string,
|
||||||
|
) (*LoginResponse, error)
|
||||||
|
```
|
||||||
|
|
||||||
|
**Location:** `pkg/security/oauth2_methods.go:375`
|
||||||
|
|
||||||
|
### Implementation Flow
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Validate provider exists
|
||||||
|
├─ getOAuth2Provider(providerName) with RLock
|
||||||
|
└─ Return error if provider not configured
|
||||||
|
|
||||||
|
2. Get session from database
|
||||||
|
├─ Call resolvespec_oauth_getrefreshtoken(refreshToken)
|
||||||
|
└─ Parse session data {user_id, access_token, token_type, expiry}
|
||||||
|
|
||||||
|
3. Refresh token with OAuth2 provider
|
||||||
|
├─ Create oauth2.Token from stored data
|
||||||
|
├─ Use provider.config.TokenSource(ctx, oldToken)
|
||||||
|
└─ Call tokenSource.Token() to get new token
|
||||||
|
|
||||||
|
4. Generate new session token
|
||||||
|
└─ Use OAuth2GenerateState() for secure random token
|
||||||
|
|
||||||
|
5. Update database
|
||||||
|
├─ Call resolvespec_oauth_updaterefreshtoken()
|
||||||
|
└─ Store new session_token, access_token, refresh_token
|
||||||
|
|
||||||
|
6. Get user data
|
||||||
|
├─ Call resolvespec_oauth_getuser(user_id)
|
||||||
|
└─ Build UserContext
|
||||||
|
|
||||||
|
7. Return LoginResponse
|
||||||
|
└─ {Token, RefreshToken, User, ExpiresIn}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Thread Safety
|
||||||
|
|
||||||
|
**Mutex Protection:** All access to `oauth2Providers` map is protected with `sync.RWMutex`
|
||||||
|
|
||||||
|
```go
|
||||||
|
type DatabaseAuthenticator struct {
|
||||||
|
oauth2Providers map[string]*OAuth2Provider
|
||||||
|
oauth2ProvidersMutex sync.RWMutex // Thread-safe access
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read operations use RLock
|
||||||
|
func (a *DatabaseAuthenticator) getOAuth2Provider(name string) {
|
||||||
|
a.oauth2ProvidersMutex.RLock()
|
||||||
|
defer a.oauth2ProvidersMutex.RUnlock()
|
||||||
|
// ... access map
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write operations use Lock
|
||||||
|
func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) {
|
||||||
|
a.oauth2ProvidersMutex.Lock()
|
||||||
|
defer a.oauth2ProvidersMutex.Unlock()
|
||||||
|
// ... modify map
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Usage Examples
|
||||||
|
|
||||||
|
### Single Provider (Google)
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
// Create Google OAuth2 authenticator
|
||||||
|
auth := security.NewGoogleAuthenticator(
|
||||||
|
"your-client-id",
|
||||||
|
"your-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Token refresh endpoint
|
||||||
|
router.HandleFunc("/auth/refresh", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
}
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Refresh token (provider name defaults to "google")
|
||||||
|
loginResp, err := auth.OAuth2RefreshToken(r.Context(), req.RefreshToken, "google")
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set new session cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Multi-Provider Setup
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Single authenticator with multiple OAuth2 providers
|
||||||
|
auth := security.NewDatabaseAuthenticator(db).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ClientID: "google-client-id",
|
||||||
|
ClientSecret: "google-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/google/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://accounts.google.com/o/oauth2/auth",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
||||||
|
ProviderName: "google",
|
||||||
|
}).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ClientID: "github-client-id",
|
||||||
|
ClientSecret: "github-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/github/callback",
|
||||||
|
Scopes: []string{"user:email"},
|
||||||
|
AuthURL: "https://github.com/login/oauth/authorize",
|
||||||
|
TokenURL: "https://github.com/login/oauth/access_token",
|
||||||
|
UserInfoURL: "https://api.github.com/user",
|
||||||
|
ProviderName: "github",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Refresh endpoint with provider selection
|
||||||
|
router.HandleFunc("/auth/refresh", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
Provider string `json:"provider"` // "google" or "github"
|
||||||
|
}
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Refresh with specific provider
|
||||||
|
loginResp, err := auth.OAuth2RefreshToken(r.Context(), req.RefreshToken, req.Provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### Client-Side Usage
|
||||||
|
|
||||||
|
```javascript
|
||||||
|
// JavaScript client example
|
||||||
|
async function refreshAccessToken() {
|
||||||
|
const refreshToken = localStorage.getItem('refresh_token');
|
||||||
|
const provider = localStorage.getItem('auth_provider'); // "google", "github", etc.
|
||||||
|
|
||||||
|
const response = await fetch('/auth/refresh', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({
|
||||||
|
refresh_token: refreshToken,
|
||||||
|
provider: provider
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.ok) {
|
||||||
|
const data = await response.json();
|
||||||
|
|
||||||
|
// Store new tokens
|
||||||
|
localStorage.setItem('access_token', data.token);
|
||||||
|
localStorage.setItem('refresh_token', data.refresh_token);
|
||||||
|
|
||||||
|
console.log('Token refreshed successfully');
|
||||||
|
return data.token;
|
||||||
|
} else {
|
||||||
|
// Refresh failed - redirect to login
|
||||||
|
window.location.href = '/login';
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Automatically refresh token when API returns 401
|
||||||
|
async function apiCall(endpoint) {
|
||||||
|
let response = await fetch(endpoint, {
|
||||||
|
headers: {
|
||||||
|
'Authorization': 'Bearer ' + localStorage.getItem('access_token')
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.status === 401) {
|
||||||
|
// Token expired - try refresh
|
||||||
|
const newToken = await refreshAccessToken();
|
||||||
|
|
||||||
|
// Retry with new token
|
||||||
|
response = await fetch(endpoint, {
|
||||||
|
headers: {
|
||||||
|
'Authorization': 'Bearer ' + newToken
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
return response.json();
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. API Reference
|
||||||
|
|
||||||
|
### DatabaseAuthenticator Methods
|
||||||
|
|
||||||
|
| Method | Signature | Description |
|
||||||
|
|--------|-----------|-------------|
|
||||||
|
| `OAuth2RefreshToken` | `(ctx, refreshToken, provider) (*LoginResponse, error)` | Refreshes expired OAuth2 access token |
|
||||||
|
| `WithOAuth2` | `(cfg OAuth2Config) *DatabaseAuthenticator` | Adds OAuth2 provider (chainable) |
|
||||||
|
| `OAuth2GetAuthURL` | `(provider, state) (string, error)` | Gets authorization URL |
|
||||||
|
| `OAuth2HandleCallback` | `(ctx, provider, code, state) (*LoginResponse, error)` | Handles OAuth2 callback |
|
||||||
|
| `OAuth2GenerateState` | `() (string, error)` | Generates CSRF state token |
|
||||||
|
| `OAuth2GetProviders` | `() []string` | Lists configured providers |
|
||||||
|
|
||||||
|
### LoginResponse Structure
|
||||||
|
|
||||||
|
```go
|
||||||
|
type LoginResponse struct {
|
||||||
|
Token string // New session token
|
||||||
|
RefreshToken string // New refresh token (may be same as input)
|
||||||
|
User *UserContext // User information
|
||||||
|
ExpiresIn int64 // Seconds until expiration
|
||||||
|
}
|
||||||
|
|
||||||
|
type UserContext struct {
|
||||||
|
UserID int // Database user ID
|
||||||
|
UserName string // Username
|
||||||
|
Email string // Email address
|
||||||
|
UserLevel int // Permission level
|
||||||
|
SessionID string // Session token
|
||||||
|
RemoteID string // OAuth2 provider user ID
|
||||||
|
Roles []string // User roles
|
||||||
|
Claims map[string]any // Additional claims
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Important Notes
|
||||||
|
|
||||||
|
### Provider Configuration
|
||||||
|
|
||||||
|
**For Google:** Add `access_type=offline` to get refresh token on first login:
|
||||||
|
|
||||||
|
```go
|
||||||
|
auth := security.NewGoogleAuthenticator(clientID, clientSecret, redirectURL, db)
|
||||||
|
// When generating auth URL, add access_type parameter
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("google", state)
|
||||||
|
authURL += "&access_type=offline&prompt=consent"
|
||||||
|
```
|
||||||
|
|
||||||
|
**For GitHub:** Refresh tokens are not always provided. Check provider documentation.
|
||||||
|
|
||||||
|
### Token Storage
|
||||||
|
|
||||||
|
- Store refresh tokens securely on client (localStorage, secure cookie, etc.)
|
||||||
|
- Never log refresh tokens
|
||||||
|
- Refresh tokens are long-lived (days/months depending on provider)
|
||||||
|
- Access tokens are short-lived (minutes/hours)
|
||||||
|
|
||||||
|
### Error Handling
|
||||||
|
|
||||||
|
Common errors:
|
||||||
|
- `"invalid or expired refresh token"` - Token expired or revoked
|
||||||
|
- `"OAuth2 provider 'xxx' not found"` - Provider not configured
|
||||||
|
- `"failed to refresh token with provider"` - Provider rejected refresh request
|
||||||
|
|
||||||
|
### Security Best Practices
|
||||||
|
|
||||||
|
1. **Always use HTTPS** for token transmission
|
||||||
|
2. **Store refresh tokens securely** on client
|
||||||
|
3. **Set appropriate cookie flags**: `HttpOnly`, `Secure`, `SameSite`
|
||||||
|
4. **Implement token rotation** - issue new refresh token on each refresh
|
||||||
|
5. **Revoke old tokens** after successful refresh
|
||||||
|
6. **Rate limit** refresh endpoints
|
||||||
|
7. **Log refresh attempts** for audit trail
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Testing
|
||||||
|
|
||||||
|
### Manual Test Flow
|
||||||
|
|
||||||
|
1. **Initial Login:**
|
||||||
|
```bash
|
||||||
|
curl http://localhost:8080/auth/google/login
|
||||||
|
# Follow redirect to Google
|
||||||
|
# Returns to callback with LoginResponse containing refresh_token
|
||||||
|
```
|
||||||
|
|
||||||
|
2. **Wait for Token Expiry (or manually expire in DB)**
|
||||||
|
|
||||||
|
3. **Refresh Token:**
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8080/auth/refresh \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"refresh_token": "ya29.a0AfH6SMB...",
|
||||||
|
"provider": "google"
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Response:
|
||||||
|
{
|
||||||
|
"token": "sess_abc123...",
|
||||||
|
"refresh_token": "ya29.a0AfH6SMB...",
|
||||||
|
"user": {
|
||||||
|
"user_id": 1,
|
||||||
|
"user_name": "john_doe",
|
||||||
|
"email": "john@example.com",
|
||||||
|
"session_id": "sess_abc123..."
|
||||||
|
},
|
||||||
|
"expires_in": 3600
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
4. **Use New Token:**
|
||||||
|
```bash
|
||||||
|
curl http://localhost:8080/api/protected \
|
||||||
|
-H "Authorization: Bearer sess_abc123..."
|
||||||
|
```
|
||||||
|
|
||||||
|
### Database Verification
|
||||||
|
|
||||||
|
```sql
|
||||||
|
-- Check session with refresh token
|
||||||
|
SELECT session_token, user_id, expires_at, refresh_token, auth_provider
|
||||||
|
FROM user_sessions
|
||||||
|
WHERE refresh_token = 'ya29.a0AfH6SMB...';
|
||||||
|
|
||||||
|
-- Verify token was updated after refresh
|
||||||
|
SELECT session_token, access_token, refresh_token,
|
||||||
|
expires_at, last_activity_at
|
||||||
|
FROM user_sessions
|
||||||
|
WHERE user_id = 1
|
||||||
|
ORDER BY created_at DESC
|
||||||
|
LIMIT 1;
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. Troubleshooting
|
||||||
|
|
||||||
|
### "Refresh token not found or expired"
|
||||||
|
|
||||||
|
**Cause:** Refresh token doesn't exist in database or session expired
|
||||||
|
|
||||||
|
**Solution:**
|
||||||
|
- Check if initial OAuth2 login stored refresh token
|
||||||
|
- Verify provider returns refresh token (some require `access_type=offline`)
|
||||||
|
- Check session hasn't been deleted from database
|
||||||
|
|
||||||
|
### "Failed to refresh token with provider"
|
||||||
|
|
||||||
|
**Cause:** OAuth2 provider rejected the refresh request
|
||||||
|
|
||||||
|
**Possible reasons:**
|
||||||
|
- Refresh token was revoked by user
|
||||||
|
- OAuth2 app credentials changed
|
||||||
|
- Network connectivity issues
|
||||||
|
- Provider rate limiting
|
||||||
|
|
||||||
|
**Solution:**
|
||||||
|
- Re-authenticate user (full OAuth2 flow)
|
||||||
|
- Check provider dashboard for app status
|
||||||
|
- Verify client credentials are correct
|
||||||
|
|
||||||
|
### "OAuth2 provider 'xxx' not found"
|
||||||
|
|
||||||
|
**Cause:** Provider not registered with `WithOAuth2()`
|
||||||
|
|
||||||
|
**Solution:**
|
||||||
|
```go
|
||||||
|
// Make sure provider is configured
|
||||||
|
auth := security.NewDatabaseAuthenticator(db).
|
||||||
|
WithOAuth2(security.OAuth2Config{
|
||||||
|
ProviderName: "google", // This name must match refresh call
|
||||||
|
// ... other config
|
||||||
|
})
|
||||||
|
|
||||||
|
// Then use same name in refresh
|
||||||
|
auth.OAuth2RefreshToken(ctx, token, "google") // Must match ProviderName
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 8. Complete Working Example
|
||||||
|
|
||||||
|
See `pkg/security/oauth2_examples.go:250` for full working example with token refresh.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Summary
|
||||||
|
|
||||||
|
OAuth2 refresh token functionality is **production-ready** with:
|
||||||
|
|
||||||
|
- ✅ Complete database schema with stored procedures
|
||||||
|
- ✅ Thread-safe Go implementation with mutex protection
|
||||||
|
- ✅ Multi-provider support (Google, GitHub, Microsoft, Facebook, custom)
|
||||||
|
- ✅ Comprehensive error handling
|
||||||
|
- ✅ Working code examples
|
||||||
|
- ✅ Full API documentation
|
||||||
|
- ✅ Security best practices implemented
|
||||||
|
|
||||||
|
**No additional implementation needed - feature is complete and functional.**
|
||||||
+208
@@ -0,0 +1,208 @@
|
|||||||
|
# Passkey Authentication Quick Reference
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
Passkey authentication (WebAuthn/FIDO2) is now integrated into the DatabaseAuthenticator. This provides passwordless authentication using biometrics, security keys, or device credentials.
|
||||||
|
|
||||||
|
## Setup
|
||||||
|
|
||||||
|
### Database Schema
|
||||||
|
Run the passkey SQL schema (in database_schema.sql):
|
||||||
|
- Creates `user_passkey_credentials` table
|
||||||
|
- Adds stored procedures for passkey operations
|
||||||
|
|
||||||
|
### Go Code
|
||||||
|
```go
|
||||||
|
// Create passkey provider
|
||||||
|
passkeyProvider := security.NewDatabasePasskeyProvider(db,
|
||||||
|
security.DatabasePasskeyProviderOptions{
|
||||||
|
RPID: "example.com",
|
||||||
|
RPName: "Example App",
|
||||||
|
RPOrigin: "https://example.com",
|
||||||
|
Timeout: 60000,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Create authenticator with passkey support
|
||||||
|
auth := security.NewDatabaseAuthenticatorWithOptions(db,
|
||||||
|
security.DatabaseAuthenticatorOptions{
|
||||||
|
PasskeyProvider: passkeyProvider,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Or add passkey to existing authenticator
|
||||||
|
auth = security.NewDatabaseAuthenticator(db).WithPasskey(passkeyProvider)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Registration Flow
|
||||||
|
|
||||||
|
### Backend - Step 1: Begin Registration
|
||||||
|
```go
|
||||||
|
options, err := auth.BeginPasskeyRegistration(ctx,
|
||||||
|
security.PasskeyBeginRegistrationRequest{
|
||||||
|
UserID: 1,
|
||||||
|
Username: "alice",
|
||||||
|
DisplayName: "Alice Smith",
|
||||||
|
})
|
||||||
|
// Send options to client as JSON
|
||||||
|
```
|
||||||
|
|
||||||
|
### Frontend - Step 2: Create Credential
|
||||||
|
```javascript
|
||||||
|
// Convert options from server
|
||||||
|
options.challenge = base64ToArrayBuffer(options.challenge);
|
||||||
|
options.user.id = base64ToArrayBuffer(options.user.id);
|
||||||
|
|
||||||
|
// Create credential
|
||||||
|
const credential = await navigator.credentials.create({
|
||||||
|
publicKey: options
|
||||||
|
});
|
||||||
|
|
||||||
|
// Send credential back to server
|
||||||
|
```
|
||||||
|
|
||||||
|
### Backend - Step 3: Complete Registration
|
||||||
|
```go
|
||||||
|
credential, err := auth.CompletePasskeyRegistration(ctx,
|
||||||
|
security.PasskeyRegisterRequest{
|
||||||
|
UserID: 1,
|
||||||
|
Response: clientResponse,
|
||||||
|
ExpectedChallenge: storedChallenge,
|
||||||
|
CredentialName: "My iPhone",
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
## Authentication Flow
|
||||||
|
|
||||||
|
### Backend - Step 1: Begin Authentication
|
||||||
|
```go
|
||||||
|
options, err := auth.BeginPasskeyAuthentication(ctx,
|
||||||
|
security.PasskeyBeginAuthenticationRequest{
|
||||||
|
Username: "alice", // Optional for resident key
|
||||||
|
})
|
||||||
|
// Send options to client as JSON
|
||||||
|
```
|
||||||
|
|
||||||
|
### Frontend - Step 2: Get Credential
|
||||||
|
```javascript
|
||||||
|
// Convert options from server
|
||||||
|
options.challenge = base64ToArrayBuffer(options.challenge);
|
||||||
|
|
||||||
|
// Get credential
|
||||||
|
const credential = await navigator.credentials.get({
|
||||||
|
publicKey: options
|
||||||
|
});
|
||||||
|
|
||||||
|
// Send assertion back to server
|
||||||
|
```
|
||||||
|
|
||||||
|
### Backend - Step 3: Complete Authentication
|
||||||
|
```go
|
||||||
|
loginResponse, err := auth.LoginWithPasskey(ctx,
|
||||||
|
security.PasskeyLoginRequest{
|
||||||
|
Response: clientAssertion,
|
||||||
|
ExpectedChallenge: storedChallenge,
|
||||||
|
Claims: map[string]any{
|
||||||
|
"ip_address": "192.168.1.1",
|
||||||
|
"user_agent": "Mozilla/5.0...",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
// Returns session token and user info
|
||||||
|
```
|
||||||
|
|
||||||
|
## Credential Management
|
||||||
|
|
||||||
|
### List Credentials
|
||||||
|
```go
|
||||||
|
credentials, err := auth.GetPasskeyCredentials(ctx, userID)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Update Credential Name
|
||||||
|
```go
|
||||||
|
err := auth.UpdatePasskeyCredentialName(ctx, userID, credentialID, "New Name")
|
||||||
|
```
|
||||||
|
|
||||||
|
### Delete Credential
|
||||||
|
```go
|
||||||
|
err := auth.DeletePasskeyCredential(ctx, userID, credentialID)
|
||||||
|
```
|
||||||
|
|
||||||
|
## HTTP Endpoints Example
|
||||||
|
|
||||||
|
### POST /api/passkey/register/begin
|
||||||
|
Request: `{user_id, username, display_name}`
|
||||||
|
Response: PasskeyRegistrationOptions
|
||||||
|
|
||||||
|
### POST /api/passkey/register/complete
|
||||||
|
Request: `{user_id, response, credential_name}`
|
||||||
|
Response: PasskeyCredential
|
||||||
|
|
||||||
|
### POST /api/passkey/login/begin
|
||||||
|
Request: `{username}` (optional)
|
||||||
|
Response: PasskeyAuthenticationOptions
|
||||||
|
|
||||||
|
### POST /api/passkey/login/complete
|
||||||
|
Request: `{response}`
|
||||||
|
Response: LoginResponse with session token
|
||||||
|
|
||||||
|
### GET /api/passkey/credentials
|
||||||
|
Response: Array of PasskeyCredential
|
||||||
|
|
||||||
|
### DELETE /api/passkey/credentials/{id}
|
||||||
|
Request: `{credential_id}`
|
||||||
|
Response: 204 No Content
|
||||||
|
|
||||||
|
## Database Stored Procedures
|
||||||
|
|
||||||
|
- `resolvespec_passkey_store_credential` - Store new credential
|
||||||
|
- `resolvespec_passkey_get_credential` - Get credential by ID
|
||||||
|
- `resolvespec_passkey_get_user_credentials` - Get all user credentials
|
||||||
|
- `resolvespec_passkey_update_counter` - Update sign counter (clone detection)
|
||||||
|
- `resolvespec_passkey_delete_credential` - Delete credential
|
||||||
|
- `resolvespec_passkey_update_name` - Update credential name
|
||||||
|
- `resolvespec_passkey_get_credentials_by_username` - Get credentials for login
|
||||||
|
|
||||||
|
## Security Features
|
||||||
|
|
||||||
|
- **Clone Detection**: Sign counter validation detects credential cloning
|
||||||
|
- **Attestation Support**: Stores attestation type (none, indirect, direct)
|
||||||
|
- **Transport Options**: Tracks authenticator transports (usb, nfc, ble, internal)
|
||||||
|
- **Backup State**: Tracks if credential is backed up/synced
|
||||||
|
- **User Verification**: Supports preferred/required user verification
|
||||||
|
|
||||||
|
## Important Notes
|
||||||
|
|
||||||
|
1. **WebAuthn Library**: Current implementation is simplified. For production, use a proper WebAuthn library like `github.com/go-webauthn/webauthn` for full verification.
|
||||||
|
|
||||||
|
2. **Challenge Storage**: Store challenges securely in session/cache. Never expose challenges to client beyond initial request.
|
||||||
|
|
||||||
|
3. **HTTPS Required**: Passkeys only work over HTTPS (except localhost).
|
||||||
|
|
||||||
|
4. **Browser Support**: Check browser compatibility for WebAuthn API.
|
||||||
|
|
||||||
|
5. **Relying Party ID**: Must match your domain exactly.
|
||||||
|
|
||||||
|
## Client-Side Helper Functions
|
||||||
|
|
||||||
|
```javascript
|
||||||
|
function base64ToArrayBuffer(base64) {
|
||||||
|
const binary = atob(base64);
|
||||||
|
const bytes = new Uint8Array(binary.length);
|
||||||
|
for (let i = 0; i < binary.length; i++) {
|
||||||
|
bytes[i] = binary.charCodeAt(i);
|
||||||
|
}
|
||||||
|
return bytes.buffer;
|
||||||
|
}
|
||||||
|
|
||||||
|
function arrayBufferToBase64(buffer) {
|
||||||
|
const bytes = new Uint8Array(buffer);
|
||||||
|
let binary = '';
|
||||||
|
for (let i = 0; i < bytes.length; i++) {
|
||||||
|
binary += String.fromCharCode(bytes[i]);
|
||||||
|
}
|
||||||
|
return btoa(binary);
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
Run tests: `go test -v ./pkg/security -run Passkey`
|
||||||
|
|
||||||
|
All passkey functionality includes comprehensive tests using sqlmock.
|
||||||
+783
@@ -0,0 +1,783 @@
|
|||||||
|
# Security Provider - Quick Reference
|
||||||
|
|
||||||
|
## 3-Step Setup
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Step 1: Create security providers
|
||||||
|
auth := security.NewDatabaseAuthenticator(db) // Session-based (recommended)
|
||||||
|
// OR: auth := security.NewJWTAuthenticator("secret-key", db)
|
||||||
|
// OR: auth := security.NewHeaderAuthenticator()
|
||||||
|
// OR: auth := security.NewGoogleAuthenticator(clientID, secret, redirectURL, db) // OAuth2
|
||||||
|
|
||||||
|
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||||
|
rowSec := security.NewDatabaseRowSecurityProvider(db)
|
||||||
|
|
||||||
|
// Step 2: Combine providers
|
||||||
|
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||||
|
|
||||||
|
// Step 3: Setup and apply middleware
|
||||||
|
securityList, _ := security.SetupSecurityProvider(handler, provider)
|
||||||
|
router.Use(security.NewAuthMiddleware(securityList))
|
||||||
|
router.Use(security.SetSecurityMiddleware(securityList))
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Stored Procedures
|
||||||
|
|
||||||
|
**All database operations use PostgreSQL stored procedures** with `resolvespec_*` naming:
|
||||||
|
|
||||||
|
### Database Authenticators
|
||||||
|
```go
|
||||||
|
// DatabaseAuthenticator uses these stored procedures:
|
||||||
|
resolvespec_login(jsonb) // Login with credentials
|
||||||
|
resolvespec_register(jsonb) // Register new user
|
||||||
|
resolvespec_logout(jsonb) // Invalidate session
|
||||||
|
resolvespec_session(text, text) // Validate session token
|
||||||
|
resolvespec_session_update(text, jsonb) // Update activity timestamp
|
||||||
|
resolvespec_refresh_token(text, jsonb) // Generate new session
|
||||||
|
|
||||||
|
// JWTAuthenticator uses these stored procedures:
|
||||||
|
resolvespec_jwt_login(text, text) // Validate credentials
|
||||||
|
resolvespec_jwt_logout(text, int) // Blacklist token
|
||||||
|
```
|
||||||
|
|
||||||
|
### Security Providers
|
||||||
|
```go
|
||||||
|
// DatabaseColumnSecurityProvider:
|
||||||
|
resolvespec_column_security(int, text, text) // Load column rules
|
||||||
|
|
||||||
|
// DatabaseRowSecurityProvider:
|
||||||
|
resolvespec_row_security(text, text, int) // Load row template
|
||||||
|
```
|
||||||
|
|
||||||
|
All stored procedures return structured results:
|
||||||
|
- Session/Login: `(p_success bool, p_error text, p_data jsonb)`
|
||||||
|
- Security: `(p_success bool, p_error text, p_rules jsonb)`
|
||||||
|
|
||||||
|
See `database_schema.sql` for complete definitions.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Interface Signatures
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Authenticator interface
|
||||||
|
type Authenticator interface {
|
||||||
|
Login(ctx context.Context, req LoginRequest) (*LoginResponse, error)
|
||||||
|
Logout(ctx context.Context, req LogoutRequest) error
|
||||||
|
Authenticate(r *http.Request) (*UserContext, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ColumnSecurityProvider interface
|
||||||
|
type ColumnSecurityProvider interface {
|
||||||
|
GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RowSecurityProvider interface
|
||||||
|
type RowSecurityProvider interface {
|
||||||
|
GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## UserContext Structure
|
||||||
|
|
||||||
|
```go
|
||||||
|
security.UserContext{
|
||||||
|
UserID: 123, // User's unique ID
|
||||||
|
UserName: "john_doe", // Username
|
||||||
|
UserLevel: 5, // User privilege level
|
||||||
|
SessionID: "sess_abc123", // Current session ID
|
||||||
|
RemoteID: "remote_xyz", // Remote system ID
|
||||||
|
Roles: []string{"admin"}, // User roles
|
||||||
|
Email: "john@example.com", // User email
|
||||||
|
Claims: map[string]any{}, // Additional authentication claims
|
||||||
|
Meta: map[string]any{}, // Additional metadata (JSON-serializable)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## ColumnSecurity Structure
|
||||||
|
|
||||||
|
```go
|
||||||
|
security.ColumnSecurity{
|
||||||
|
Path: []string{"column_name"}, // ["ssn"] or ["address", "street"]
|
||||||
|
Accesstype: "mask", // "mask" or "hide"
|
||||||
|
MaskStart: 5, // Mask first N chars
|
||||||
|
MaskEnd: 0, // Mask last N chars
|
||||||
|
MaskChar: "*", // Masking character
|
||||||
|
MaskInvert: false, // true = mask middle
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Common Examples
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Hide entire field
|
||||||
|
{Path: []string{"salary"}, Accesstype: "hide"}
|
||||||
|
|
||||||
|
// Mask SSN (show last 4)
|
||||||
|
{Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5}
|
||||||
|
|
||||||
|
// Mask credit card (show last 4)
|
||||||
|
{Path: []string{"credit_card"}, Accesstype: "mask", MaskStart: 12}
|
||||||
|
|
||||||
|
// Mask email (j***@example.com)
|
||||||
|
{Path: []string{"email"}, Accesstype: "mask", MaskStart: 1, MaskEnd: 0}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## RowSecurity Structure
|
||||||
|
|
||||||
|
```go
|
||||||
|
security.RowSecurity{
|
||||||
|
Schema: "public",
|
||||||
|
Tablename: "orders",
|
||||||
|
UserID: 123,
|
||||||
|
Template: "user_id = {UserID}", // WHERE clause
|
||||||
|
HasBlock: false, // true = block all access
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Template Variables
|
||||||
|
|
||||||
|
- `{UserID}` - Current user ID
|
||||||
|
- `{PrimaryKeyName}` - Primary key column
|
||||||
|
- `{TableName}` - Table name
|
||||||
|
- `{SchemaName}` - Schema name
|
||||||
|
|
||||||
|
### Common Examples
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Users see only their records
|
||||||
|
Template: "user_id = {UserID}"
|
||||||
|
|
||||||
|
// Users see their records OR public ones
|
||||||
|
Template: "user_id = {UserID} OR is_public = true"
|
||||||
|
|
||||||
|
// Tenant isolation
|
||||||
|
Template: "tenant_id = 5 AND user_id = {UserID}"
|
||||||
|
|
||||||
|
// Complex with subquery
|
||||||
|
Template: "dept_id IN (SELECT dept_id FROM user_depts WHERE user_id = {UserID})"
|
||||||
|
|
||||||
|
// Block all access
|
||||||
|
HasBlock: true
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Example Implementations
|
||||||
|
|
||||||
|
### Database Session Authenticator (Recommended)
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Create authenticator
|
||||||
|
auth := security.NewDatabaseAuthenticator(db)
|
||||||
|
|
||||||
|
// Requires these tables:
|
||||||
|
// - users (id, username, email, password, user_level, roles, is_active)
|
||||||
|
// - user_sessions (session_token, user_id, expires_at, created_at, last_activity_at)
|
||||||
|
// See database_schema.sql for full schema
|
||||||
|
|
||||||
|
// Features:
|
||||||
|
// - Login with username/password
|
||||||
|
// - Session management in database
|
||||||
|
// - Token refresh support (implements Refreshable)
|
||||||
|
// - Automatic session expiration
|
||||||
|
// - Tracks IP address and user agent
|
||||||
|
// - Works with Authorization header or cookie
|
||||||
|
```
|
||||||
|
|
||||||
|
### Simple Header Authenticator
|
||||||
|
|
||||||
|
```go
|
||||||
|
type HeaderAuthenticator struct{}
|
||||||
|
|
||||||
|
func NewHeaderAuthenticator() *HeaderAuthenticator {
|
||||||
|
return &HeaderAuthenticator{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *HeaderAuthenticator) Login(ctx context.Context, req security.LoginRequest) (*security.LoginResponse, error) {
|
||||||
|
return nil, fmt.Errorf("not supported")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *HeaderAuthenticator) Logout(ctx context.Context, req security.LogoutRequest) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *HeaderAuthenticator) Authenticate(r *http.Request) (*security.UserContext, error) {
|
||||||
|
userIDStr := r.Header.Get("X-User-ID")
|
||||||
|
if userIDStr == "" {
|
||||||
|
return nil, fmt.Errorf("X-User-ID required")
|
||||||
|
}
|
||||||
|
userID, _ := strconv.Atoi(userIDStr)
|
||||||
|
return &security.UserContext{
|
||||||
|
UserID: userID,
|
||||||
|
UserName: r.Header.Get("X-User-Name"),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### JWT Authenticator
|
||||||
|
|
||||||
|
```go
|
||||||
|
type JWTAuthenticator struct {
|
||||||
|
secretKey []byte
|
||||||
|
db *gorm.DB
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewJWTAuthenticator(secret string, db *gorm.DB) *JWTAuthenticator {
|
||||||
|
return &JWTAuthenticator{secretKey: []byte(secret), db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *JWTAuthenticator) Login(ctx context.Context, req security.LoginRequest) (*security.LoginResponse, error) {
|
||||||
|
// Validate credentials against database
|
||||||
|
var user User
|
||||||
|
err := a.db.WithContext(ctx).Where("username = ?", req.Username).First(&user).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid credentials")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate JWT token
|
||||||
|
token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||||
|
"user_id": user.ID,
|
||||||
|
"exp": time.Now().Add(24 * time.Hour).Unix(),
|
||||||
|
})
|
||||||
|
tokenString, _ := token.SignedString(a.secretKey)
|
||||||
|
|
||||||
|
return &security.LoginResponse{
|
||||||
|
Token: tokenString,
|
||||||
|
User: &security.UserContext{UserID: user.ID},
|
||||||
|
ExpiresIn: 86400,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *JWTAuthenticator) Logout(ctx context.Context, req security.LogoutRequest) error {
|
||||||
|
// Invalidate session via stored procedure
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *JWTAuthenticator) Authenticate(r *http.Request) (*security.UserContext, error) {
|
||||||
|
tokenString := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")
|
||||||
|
token, err := jwt.Parse(tokenString, func(t *jwt.Token) (any, error) {
|
||||||
|
return a.secretKey, nil
|
||||||
|
})
|
||||||
|
if err != nil || !token.Valid {
|
||||||
|
return nil, fmt.Errorf("invalid token")
|
||||||
|
}
|
||||||
|
claims := token.Claims.(jwt.MapClaims)
|
||||||
|
return &security.UserContext{
|
||||||
|
UserID: int(claims["user_id"].(float64)),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Static Column Security
|
||||||
|
|
||||||
|
```go
|
||||||
|
type ConfigColumnSecurityProvider struct {
|
||||||
|
rules map[string][]security.ColumnSecurity
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewConfigColumnSecurityProvider(rules map[string][]security.ColumnSecurity) *ConfigColumnSecurityProvider {
|
||||||
|
return &ConfigColumnSecurityProvider{rules: rules}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *ConfigColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) {
|
||||||
|
key := fmt.Sprintf("%s.%s", schema, table)
|
||||||
|
return p.rules[key], nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Database Column Security
|
||||||
|
|
||||||
|
```go
|
||||||
|
type DatabaseColumnSecurityProvider struct {
|
||||||
|
db *gorm.DB
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewDatabaseColumnSecurityProvider(db *gorm.DB) *DatabaseColumnSecurityProvider {
|
||||||
|
return &DatabaseColumnSecurityProvider{db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) {
|
||||||
|
var records []struct {
|
||||||
|
Control string
|
||||||
|
Accesstype string
|
||||||
|
JSONValue string
|
||||||
|
}
|
||||||
|
|
||||||
|
query := `
|
||||||
|
SELECT control, accesstype, jsonvalue
|
||||||
|
FROM core.secaccess
|
||||||
|
WHERE rid_hub IN (
|
||||||
|
SELECT rid_hub_parent FROM core.hub_link
|
||||||
|
WHERE rid_hub_child = ? AND parent_hubtype = 'secgroup'
|
||||||
|
)
|
||||||
|
AND control ILIKE ?
|
||||||
|
`
|
||||||
|
|
||||||
|
err := p.db.WithContext(ctx).Raw(query, userID, fmt.Sprintf("%s.%s%%", schema, table)).Scan(&records).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var rules []security.ColumnSecurity
|
||||||
|
for _, rec := range records {
|
||||||
|
parts := strings.Split(rec.Control, ".")
|
||||||
|
if len(parts) < 3 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rules = append(rules, security.ColumnSecurity{
|
||||||
|
Schema: schema,
|
||||||
|
Tablename: table,
|
||||||
|
Path: parts[2:],
|
||||||
|
Accesstype: rec.Accesstype,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return rules, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Static Row Security
|
||||||
|
|
||||||
|
```go
|
||||||
|
type ConfigRowSecurityProvider struct {
|
||||||
|
templates map[string]string
|
||||||
|
blocked map[string]bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewConfigRowSecurityProvider(templates map[string]string, blocked map[string]bool) *ConfigRowSecurityProvider {
|
||||||
|
return &ConfigRowSecurityProvider{templates: templates, blocked: blocked}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *ConfigRowSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (security.RowSecurity, error) {
|
||||||
|
key := fmt.Sprintf("%s.%s", schema, table)
|
||||||
|
|
||||||
|
if p.blocked[key] {
|
||||||
|
return security.RowSecurity{HasBlock: true}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return security.RowSecurity{
|
||||||
|
Schema: schema,
|
||||||
|
Tablename: table,
|
||||||
|
UserID: userID,
|
||||||
|
Template: p.templates[key],
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Test Authenticator
|
||||||
|
auth := security.NewHeaderAuthenticator()
|
||||||
|
req := httptest.NewRequest("GET", "/", nil)
|
||||||
|
req.Header.Set("X-User-ID", "123")
|
||||||
|
userCtx, err := auth.Authenticate(req)
|
||||||
|
assert.Equal(t, 123, userCtx.UserID)
|
||||||
|
|
||||||
|
// Test ColumnSecurityProvider
|
||||||
|
colSec := security.NewConfigColumnSecurityProvider(rules)
|
||||||
|
cols, err := colSec.GetColumnSecurity(context.Background(), 123, "public", "employees")
|
||||||
|
assert.Equal(t, "mask", cols[0].Accesstype)
|
||||||
|
|
||||||
|
// Test RowSecurityProvider
|
||||||
|
rowSec := security.NewConfigRowSecurityProvider(templates, blocked)
|
||||||
|
row, err := rowSec.GetRowSecurity(context.Background(), 123, "public", "orders")
|
||||||
|
assert.Equal(t, "user_id = {UserID}", row.Template)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Request Flow
|
||||||
|
|
||||||
|
```
|
||||||
|
HTTP Request
|
||||||
|
↓
|
||||||
|
NewOptionalAuthMiddleware → calls provider.Authenticate()
|
||||||
|
↓ (adds UserContext or guest context; never 401)
|
||||||
|
SetSecurityMiddleware → adds SecurityList to context
|
||||||
|
↓
|
||||||
|
Handler.Handle() → resolves model
|
||||||
|
↓
|
||||||
|
BeforeHandle Hook → CheckModelAuthAllowed(secCtx, operation)
|
||||||
|
├─ SecurityDisabled → allow
|
||||||
|
├─ CanPublicRead/Create/Update/Delete → allow unauthenticated
|
||||||
|
└─ UserID == 0 → abort 401
|
||||||
|
↓
|
||||||
|
BeforeRead Hook → calls provider.GetColumnSecurity() + GetRowSecurity()
|
||||||
|
↓
|
||||||
|
BeforeScan Hook → applies row security (WHERE clause)
|
||||||
|
↓
|
||||||
|
Database Query
|
||||||
|
↓
|
||||||
|
AfterRead Hook → applies column security (masking)
|
||||||
|
↓
|
||||||
|
HTTP Response
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Common Patterns
|
||||||
|
|
||||||
|
### Role-Based Security
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (p *MyColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) {
|
||||||
|
userCtx, _ := security.GetUserContext(ctx)
|
||||||
|
|
||||||
|
if contains(userCtx.Roles, "admin") {
|
||||||
|
return []security.ColumnSecurity{}, nil // No restrictions
|
||||||
|
}
|
||||||
|
|
||||||
|
return loadRestrictions(userID, schema, table), nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Tenant Isolation
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (p *MyRowSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (security.RowSecurity, error) {
|
||||||
|
tenantID := getUserTenant(userID)
|
||||||
|
return security.RowSecurity{
|
||||||
|
Template: fmt.Sprintf("tenant_id = %d", tenantID),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Caching with Decorator
|
||||||
|
|
||||||
|
```go
|
||||||
|
type CachedColumnSecurityProvider struct {
|
||||||
|
inner security.ColumnSecurityProvider
|
||||||
|
cache *cache.Cache
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *CachedColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) {
|
||||||
|
key := fmt.Sprintf("%d:%s.%s", userID, schema, table)
|
||||||
|
|
||||||
|
if cached, found := p.cache.Get(key); found {
|
||||||
|
return cached.([]security.ColumnSecurity), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
rules, err := p.inner.GetColumnSecurity(ctx, userID, schema, table)
|
||||||
|
if err == nil {
|
||||||
|
p.cache.Set(key, rules, cache.DefaultExpiration)
|
||||||
|
}
|
||||||
|
return rules, err
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Error Handling
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Panic if provider is nil
|
||||||
|
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||||
|
// panics if any parameter is nil
|
||||||
|
|
||||||
|
// Auth middleware returns 401 if Authenticate fails
|
||||||
|
func (a *MyAuthenticator) Authenticate(r *http.Request) (*security.UserContext, error) {
|
||||||
|
if invalid {
|
||||||
|
return nil, fmt.Errorf("invalid credentials") // Returns HTTP 401
|
||||||
|
}
|
||||||
|
return &security.UserContext{UserID: userID}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Security loading can fail gracefully
|
||||||
|
func (p *MyProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) {
|
||||||
|
rules, err := db.Load(...)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("Failed to load security: %v", err)
|
||||||
|
return []security.ColumnSecurity{}, nil // No rules = no restrictions
|
||||||
|
}
|
||||||
|
return rules, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Login/Logout/Register Endpoints
|
||||||
|
|
||||||
|
```go
|
||||||
|
func SetupAuthRoutes(router *mux.Router, securityList *security.SecurityList) {
|
||||||
|
// Register
|
||||||
|
router.HandleFunc("/auth/register", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req security.RegisterRequest
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Check if provider supports registration
|
||||||
|
registrable, ok := securityList.Provider().(security.Registrable)
|
||||||
|
if !ok {
|
||||||
|
http.Error(w, "Registration not supported", http.StatusNotImplemented)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := registrable.Register(r.Context(), req)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}).Methods("POST")
|
||||||
|
|
||||||
|
// Login
|
||||||
|
router.HandleFunc("/auth/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req security.LoginRequest
|
||||||
|
json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
resp, err := securityList.Provider().Login(r.Context(), req)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}).Methods("POST")
|
||||||
|
|
||||||
|
// Logout
|
||||||
|
router.HandleFunc("/auth/logout", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
token := r.Header.Get("Authorization")
|
||||||
|
userID, _ := security.GetUserID(r.Context())
|
||||||
|
|
||||||
|
err := securityList.Provider().Logout(r.Context(), security.LogoutRequest{
|
||||||
|
Token: token,
|
||||||
|
UserID: userID,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}).Methods("POST")
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Debugging
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Enable debug logging
|
||||||
|
import "github.com/bitechdev/GoCore/pkg/cfg"
|
||||||
|
cfg.SetLogLevel("DEBUG")
|
||||||
|
|
||||||
|
// Log in provider methods
|
||||||
|
func (a *MyAuthenticator) Authenticate(r *http.Request) (*security.UserContext, error) {
|
||||||
|
token := r.Header.Get("Authorization")
|
||||||
|
log.Printf("Auth: token=%s", token)
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if methods are called
|
||||||
|
func (p *MyColumnSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) {
|
||||||
|
log.Printf("Loading column security: user=%d, schema=%s, table=%s", userID, schema, table)
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Complete Minimal Example
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/restheadspec"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Simple all-in-one provider
|
||||||
|
type SimpleProvider struct{}
|
||||||
|
|
||||||
|
func (p *SimpleProvider) Login(ctx context.Context, req security.LoginRequest) (*security.LoginResponse, error) {
|
||||||
|
return nil, fmt.Errorf("not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *SimpleProvider) Logout(ctx context.Context, req security.LogoutRequest) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *SimpleProvider) Authenticate(r *http.Request) (*security.UserContext, error) {
|
||||||
|
id, _ := strconv.Atoi(r.Header.Get("X-User-ID"))
|
||||||
|
return &security.UserContext{UserID: id}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *SimpleProvider) GetColumnSecurity(ctx context.Context, u int, s, t string) ([]security.ColumnSecurity, error) {
|
||||||
|
return []security.ColumnSecurity{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *SimpleProvider) GetRowSecurity(ctx context.Context, u int, s, t string) (security.RowSecurity, error) {
|
||||||
|
return security.RowSecurity{Template: fmt.Sprintf("user_id = %d", u)}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
handler := restheadspec.NewHandlerWithGORM(db)
|
||||||
|
|
||||||
|
// Setup security
|
||||||
|
provider := &SimpleProvider{}
|
||||||
|
securityList := security.SetupSecurityProvider(handler, provider)
|
||||||
|
|
||||||
|
// Apply middleware
|
||||||
|
router := mux.NewRouter()
|
||||||
|
restheadspec.SetupMuxRoutes(router, handler)
|
||||||
|
router.Use(security.NewAuthMiddleware(securityList))
|
||||||
|
router.Use(security.SetSecurityMiddleware(securityList))
|
||||||
|
|
||||||
|
http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Authentication Modes
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Required authentication (default)
|
||||||
|
// Authentication must succeed or returns 401
|
||||||
|
router.Use(security.NewAuthMiddleware(securityList))
|
||||||
|
|
||||||
|
// Skip authentication for specific routes
|
||||||
|
// Always sets guest user context
|
||||||
|
func PublicRoute(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx := security.SkipAuth(r.Context())
|
||||||
|
r = r.WithContext(ctx)
|
||||||
|
// Guest context will be set
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optional authentication for specific routes
|
||||||
|
// Tries to authenticate, falls back to guest if it fails
|
||||||
|
func HomeRoute(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx := security.OptionalAuth(r.Context())
|
||||||
|
r = r.WithContext(ctx)
|
||||||
|
|
||||||
|
userCtx, _ := security.GetUserContext(r.Context())
|
||||||
|
if userCtx.UserID == 0 {
|
||||||
|
// Guest user
|
||||||
|
} else {
|
||||||
|
// Authenticated user
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Comparison:**
|
||||||
|
- **Required**: Auth must succeed or return 401 (default)
|
||||||
|
- **SkipAuth**: Never tries to authenticate, always guest
|
||||||
|
- **OptionalAuth**: Tries to authenticate, guest on failure
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Standalone Handlers
|
||||||
|
|
||||||
|
```go
|
||||||
|
// NewAuthHandler - Required authentication (returns 401 on failure)
|
||||||
|
authHandler := security.NewAuthHandler(securityList, myHandler)
|
||||||
|
http.Handle("/api/protected", authHandler)
|
||||||
|
|
||||||
|
// NewOptionalAuthHandler - Optional authentication (guest on failure)
|
||||||
|
optionalHandler := security.NewOptionalAuthHandler(securityList, myHandler)
|
||||||
|
http.Handle("/home", optionalHandler)
|
||||||
|
|
||||||
|
// NewOptionalAuthMiddleware - For spec routes; auth enforcement deferred to BeforeHandle
|
||||||
|
apiRouter.Use(security.NewOptionalAuthMiddleware(securityList))
|
||||||
|
apiRouter.Use(security.SetSecurityMiddleware(securityList))
|
||||||
|
restheadspec.RegisterSecurityHooks(handler, securityList) // includes BeforeHandle
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Model-Level Access Control
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Register model with rules (pkg/modelregistry)
|
||||||
|
modelregistry.RegisterModelWithRules("public.products", &Product{}, modelregistry.ModelRules{
|
||||||
|
SecurityDisabled: false, // skip all auth when true
|
||||||
|
CanPublicRead: true, // unauthenticated reads allowed
|
||||||
|
CanPublicCreate: false, // requires auth
|
||||||
|
CanPublicUpdate: false, // requires auth
|
||||||
|
CanPublicDelete: false, // requires auth
|
||||||
|
CanUpdate: true, // authenticated can update
|
||||||
|
CanDelete: false, // authenticated cannot delete (enforced in BeforeDelete)
|
||||||
|
})
|
||||||
|
|
||||||
|
// CheckModelAuthAllowed used automatically in BeforeHandle hook
|
||||||
|
// No code needed — call RegisterSecurityHooks and it's applied
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Context Helpers
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Get full user context
|
||||||
|
userCtx, ok := security.GetUserContext(ctx)
|
||||||
|
|
||||||
|
// Get individual fields
|
||||||
|
userID, ok := security.GetUserID(ctx)
|
||||||
|
userName, ok := security.GetUserName(ctx)
|
||||||
|
userLevel, ok := security.GetUserLevel(ctx)
|
||||||
|
sessionID, ok := security.GetSessionID(ctx)
|
||||||
|
remoteID, ok := security.GetRemoteID(ctx)
|
||||||
|
roles, ok := security.GetUserRoles(ctx)
|
||||||
|
email, ok := security.GetUserEmail(ctx)
|
||||||
|
meta, ok := security.GetUserMeta(ctx)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Resources
|
||||||
|
|
||||||
|
| File | Description |
|
||||||
|
|------|-------------|
|
||||||
|
| `INTERFACE_GUIDE.md` | **Start here** - Complete implementation guide |
|
||||||
|
| `OAUTH2.md` | **OAuth2 Guide** - Google, GitHub, Microsoft, Facebook, custom providers |
|
||||||
|
| `examples.go` | Working provider implementations to copy |
|
||||||
|
| `setup_example.go` | 6 complete integration examples |
|
||||||
|
| `README.md` | Architecture overview and migration guide |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Cheat Sheet
|
||||||
|
|
||||||
|
```go
|
||||||
|
// ===== REQUIRED SETUP =====
|
||||||
|
auth := security.NewJWTAuthenticator("secret", db)
|
||||||
|
colSec := security.NewDatabaseColumnSecurityProvider(db)
|
||||||
|
rowSec := security.NewDatabaseRowSecurityProvider(db)
|
||||||
|
provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||||
|
securityList := security.SetupSecurityProvider(handler, provider)
|
||||||
|
|
||||||
|
// ===== INTERFACE METHODS =====
|
||||||
|
Authenticate(r *http.Request) (*UserContext, error)
|
||||||
|
Login(ctx context.Context, req LoginRequest) (*LoginResponse, error)
|
||||||
|
Logout(ctx context.Context, req LogoutRequest) error
|
||||||
|
GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error)
|
||||||
|
GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error)
|
||||||
|
|
||||||
|
// ===== QUICK EXAMPLES =====
|
||||||
|
// Header auth
|
||||||
|
&UserContext{UserID: 123, UserName: "john"}
|
||||||
|
|
||||||
|
// Mask SSN
|
||||||
|
{Path: []string{"ssn"}, Accesstype: "mask", MaskStart: 5}
|
||||||
|
|
||||||
|
// User isolation
|
||||||
|
{Template: "user_id = {UserID}"}
|
||||||
|
```
|
||||||
+1417
File diff suppressed because it is too large
Load Diff
+440
@@ -0,0 +1,440 @@
|
|||||||
|
# Security Features: Blacklist & Rate Limit Inspection
|
||||||
|
|
||||||
|
## IP Blacklist
|
||||||
|
|
||||||
|
The IP blacklist middleware allows you to block specific IP addresses or CIDR ranges from accessing your application.
|
||||||
|
|
||||||
|
### Basic Usage
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/middleware"
|
||||||
|
|
||||||
|
// Create blacklist (UseProxy=true if behind a proxy)
|
||||||
|
blacklist := middleware.NewIPBlacklist(middleware.BlacklistConfig{
|
||||||
|
UseProxy: true, // Checks X-Forwarded-For and X-Real-IP headers
|
||||||
|
})
|
||||||
|
|
||||||
|
// Block individual IP
|
||||||
|
blacklist.BlockIP("192.168.1.100", "Suspicious activity detected")
|
||||||
|
|
||||||
|
// Block entire CIDR range
|
||||||
|
blacklist.BlockCIDR("10.0.0.0/8", "Private network blocked")
|
||||||
|
|
||||||
|
// Apply middleware
|
||||||
|
http.Handle("/api/", blacklist.Middleware(yourHandler))
|
||||||
|
```
|
||||||
|
|
||||||
|
### Managing Blacklist
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Unblock an IP
|
||||||
|
blacklist.UnblockIP("192.168.1.100")
|
||||||
|
|
||||||
|
// Unblock a CIDR range
|
||||||
|
blacklist.UnblockCIDR("10.0.0.0/8")
|
||||||
|
|
||||||
|
// Get all blacklisted IPs and CIDRs
|
||||||
|
ips, cidrs := blacklist.GetBlacklist()
|
||||||
|
fmt.Printf("Blocked IPs: %v\n", ips)
|
||||||
|
fmt.Printf("Blocked CIDRs: %v\n", cidrs)
|
||||||
|
|
||||||
|
// Check if specific IP is blocked
|
||||||
|
blocked, reason := blacklist.IsBlocked("192.168.1.100")
|
||||||
|
if blocked {
|
||||||
|
fmt.Printf("IP blocked: %s\n", reason)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Blacklist Statistics Endpoint
|
||||||
|
|
||||||
|
Expose blacklist statistics via HTTP:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Add stats endpoint
|
||||||
|
http.Handle("/admin/blacklist-stats", blacklist.StatsHandler())
|
||||||
|
```
|
||||||
|
|
||||||
|
**Example Response:**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"blocked_ips": ["192.168.1.100", "192.168.1.101"],
|
||||||
|
"blocked_cidrs": ["10.0.0.0/8"],
|
||||||
|
"total_ips": 2,
|
||||||
|
"total_cidrs": 1
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Integration Example
|
||||||
|
|
||||||
|
```go
|
||||||
|
func main() {
|
||||||
|
// Create blacklist
|
||||||
|
blacklist := middleware.NewIPBlacklist(middleware.BlacklistConfig{
|
||||||
|
UseProxy: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Block known malicious IPs
|
||||||
|
blacklist.BlockIP("203.0.113.1", "Known scanner")
|
||||||
|
blacklist.BlockCIDR("198.51.100.0/24", "Spam network")
|
||||||
|
|
||||||
|
// Create your router
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
|
||||||
|
// Protected routes
|
||||||
|
mux.Handle("/api/", blacklist.Middleware(apiHandler))
|
||||||
|
|
||||||
|
// Admin endpoint to manage blacklist
|
||||||
|
mux.HandleFunc("/admin/block-ip", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ip := r.URL.Query().Get("ip")
|
||||||
|
reason := r.URL.Query().Get("reason")
|
||||||
|
|
||||||
|
if err := blacklist.BlockIP(ip, reason); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
fmt.Fprintf(w, "Blocked %s: %s", ip, reason)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Stats endpoint
|
||||||
|
mux.Handle("/admin/blacklist-stats", blacklist.StatsHandler())
|
||||||
|
|
||||||
|
http.ListenAndServe(":8080", mux)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Rate Limit Inspection
|
||||||
|
|
||||||
|
Monitor and inspect rate limit status per IP address in real-time.
|
||||||
|
|
||||||
|
### Basic Usage
|
||||||
|
|
||||||
|
```go
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/middleware"
|
||||||
|
|
||||||
|
// Create rate limiter (10 req/sec, burst of 20)
|
||||||
|
rateLimiter := middleware.NewRateLimiter(10, 20)
|
||||||
|
|
||||||
|
// Apply middleware
|
||||||
|
http.Handle("/api/", rateLimiter.Middleware(yourHandler))
|
||||||
|
```
|
||||||
|
|
||||||
|
### Programmatic Inspection
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Get all tracked IPs
|
||||||
|
trackedIPs := rateLimiter.GetTrackedIPs()
|
||||||
|
fmt.Printf("Currently tracking %d IPs\n", len(trackedIPs))
|
||||||
|
|
||||||
|
// Get rate limit info for specific IP
|
||||||
|
info := rateLimiter.GetRateLimitInfo("192.168.1.1")
|
||||||
|
fmt.Printf("IP: %s\n", info.IP)
|
||||||
|
fmt.Printf("Tokens Remaining: %.2f\n", info.TokensRemaining)
|
||||||
|
fmt.Printf("Limit: %.2f req/sec\n", info.Limit)
|
||||||
|
fmt.Printf("Burst: %d\n", info.Burst)
|
||||||
|
|
||||||
|
// Get info for all tracked IPs
|
||||||
|
allInfo := rateLimiter.GetAllRateLimitInfo()
|
||||||
|
for _, info := range allInfo {
|
||||||
|
fmt.Printf("%s: %.2f tokens remaining\n", info.IP, info.TokensRemaining)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Rate Limit Stats Endpoint
|
||||||
|
|
||||||
|
Expose rate limit statistics via HTTP:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Add stats endpoint
|
||||||
|
http.Handle("/admin/rate-limit-stats", rateLimiter.StatsHandler())
|
||||||
|
```
|
||||||
|
|
||||||
|
**Example Response (all IPs):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"total_tracked_ips": 3,
|
||||||
|
"rate_limit_config": {
|
||||||
|
"requests_per_second": 10,
|
||||||
|
"burst": 20
|
||||||
|
},
|
||||||
|
"tracked_ips": [
|
||||||
|
{
|
||||||
|
"ip": "192.168.1.1",
|
||||||
|
"tokens_remaining": 15.5,
|
||||||
|
"limit": 10,
|
||||||
|
"burst": 20
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ip": "192.168.1.2",
|
||||||
|
"tokens_remaining": 18.2,
|
||||||
|
"limit": 10,
|
||||||
|
"burst": 20
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Example Response (specific IP):**
|
||||||
|
```bash
|
||||||
|
GET /admin/rate-limit-stats?ip=192.168.1.1
|
||||||
|
```
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"ip": "192.168.1.1",
|
||||||
|
"tokens_remaining": 15.5,
|
||||||
|
"limit": 10,
|
||||||
|
"burst": 20
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Complete Integration Example
|
||||||
|
|
||||||
|
```go
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/middleware"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
// Create rate limiter
|
||||||
|
rateLimiter := middleware.NewRateLimiter(10, 20)
|
||||||
|
|
||||||
|
// Create blacklist
|
||||||
|
blacklist := middleware.NewIPBlacklist(middleware.BlacklistConfig{
|
||||||
|
UseProxy: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
|
||||||
|
// API handler with both middlewares (blacklist first, then rate limit)
|
||||||
|
apiHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{
|
||||||
|
"message": "Success",
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
// Apply middleware chain: blacklist -> rate limit -> handler
|
||||||
|
mux.Handle("/api/", blacklist.Middleware(rateLimiter.Middleware(apiHandler)))
|
||||||
|
|
||||||
|
// Admin endpoints
|
||||||
|
mux.Handle("/admin/rate-limit-stats", rateLimiter.StatsHandler())
|
||||||
|
mux.Handle("/admin/blacklist-stats", blacklist.StatsHandler())
|
||||||
|
|
||||||
|
// Custom monitoring endpoint
|
||||||
|
mux.HandleFunc("/admin/monitor", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Get rate limit stats
|
||||||
|
rateLimitInfo := rateLimiter.GetAllRateLimitInfo()
|
||||||
|
|
||||||
|
// Get blacklist stats
|
||||||
|
blockedIPs, blockedCIDRs := blacklist.GetBlacklist()
|
||||||
|
|
||||||
|
response := map[string]interface{}{
|
||||||
|
"rate_limits": rateLimitInfo,
|
||||||
|
"blacklist": map[string]interface{}{
|
||||||
|
"ips": blockedIPs,
|
||||||
|
"cidrs": blockedCIDRs,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(response)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Dynamic blacklist management
|
||||||
|
mux.HandleFunc("/admin/block", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ip := r.URL.Query().Get("ip")
|
||||||
|
reason := r.URL.Query().Get("reason")
|
||||||
|
|
||||||
|
if ip == "" {
|
||||||
|
http.Error(w, "IP required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := blacklist.BlockIP(ip, reason); err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(w, "Blocked %s: %s", ip, reason)
|
||||||
|
})
|
||||||
|
|
||||||
|
mux.HandleFunc("/admin/unblock", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ip := r.URL.Query().Get("ip")
|
||||||
|
if ip == "" {
|
||||||
|
http.Error(w, "IP required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
blacklist.UnblockIP(ip)
|
||||||
|
fmt.Fprintf(w, "Unblocked %s", ip)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Auto-block IPs that exceed rate limit
|
||||||
|
mux.HandleFunc("/admin/auto-block-heavy-users", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
blocked := 0
|
||||||
|
|
||||||
|
for _, info := range rateLimiter.GetAllRateLimitInfo() {
|
||||||
|
// If tokens are very low, IP is making many requests
|
||||||
|
if info.TokensRemaining < 1.0 {
|
||||||
|
blacklist.BlockIP(info.IP, "Exceeded rate limit")
|
||||||
|
blocked++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(w, "Blocked %d IPs exceeding rate limits", blocked)
|
||||||
|
})
|
||||||
|
|
||||||
|
fmt.Println("Server starting on :8080")
|
||||||
|
fmt.Println("Rate limit stats: http://localhost:8080/admin/rate-limit-stats")
|
||||||
|
fmt.Println("Blacklist stats: http://localhost:8080/admin/blacklist-stats")
|
||||||
|
http.ListenAndServe(":8080", mux)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Monitoring Dashboard Example
|
||||||
|
|
||||||
|
Create a simple monitoring page:
|
||||||
|
|
||||||
|
```go
|
||||||
|
mux.HandleFunc("/admin/dashboard", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
html := `
|
||||||
|
<html>
|
||||||
|
<head>
|
||||||
|
<title>Security Dashboard</title>
|
||||||
|
<script>
|
||||||
|
async function loadStats() {
|
||||||
|
const rateLimitRes = await fetch('/admin/rate-limit-stats');
|
||||||
|
const rateLimitData = await rateLimitRes.json();
|
||||||
|
|
||||||
|
const blacklistRes = await fetch('/admin/blacklist-stats');
|
||||||
|
const blacklistData = await blacklistRes.json();
|
||||||
|
|
||||||
|
document.getElementById('rate-limit').innerHTML =
|
||||||
|
JSON.stringify(rateLimitData, null, 2);
|
||||||
|
document.getElementById('blacklist').innerHTML =
|
||||||
|
JSON.stringify(blacklistData, null, 2);
|
||||||
|
}
|
||||||
|
|
||||||
|
setInterval(loadStats, 5000); // Refresh every 5 seconds
|
||||||
|
loadStats();
|
||||||
|
</script>
|
||||||
|
</head>
|
||||||
|
<body>
|
||||||
|
<h1>Security Dashboard</h1>
|
||||||
|
|
||||||
|
<h2>Rate Limits</h2>
|
||||||
|
<pre id="rate-limit">Loading...</pre>
|
||||||
|
|
||||||
|
<h2>Blacklist</h2>
|
||||||
|
<pre id="blacklist">Loading...</pre>
|
||||||
|
</body>
|
||||||
|
</html>
|
||||||
|
`
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "text/html")
|
||||||
|
w.Write([]byte(html))
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Best Practices
|
||||||
|
|
||||||
|
### 1. Proxy Configuration
|
||||||
|
Always set `UseProxy: true` when running behind a reverse proxy (nginx, Cloudflare, etc.):
|
||||||
|
```go
|
||||||
|
blacklist := middleware.NewIPBlacklist(middleware.BlacklistConfig{
|
||||||
|
UseProxy: true, // Checks X-Forwarded-For headers
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. Middleware Order
|
||||||
|
Apply blacklist before rate limiting to save resources:
|
||||||
|
```go
|
||||||
|
// Correct order: blacklist -> rate limit -> handler
|
||||||
|
handler := blacklist.Middleware(
|
||||||
|
rateLimiter.Middleware(yourHandler)
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. Secure Admin Endpoints
|
||||||
|
Protect admin endpoints with authentication:
|
||||||
|
```go
|
||||||
|
mux.Handle("/admin/", authMiddleware(adminHandler))
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4. Monitoring
|
||||||
|
Set up alerts when:
|
||||||
|
- Many IPs are being rate limited
|
||||||
|
- Blacklist grows too large
|
||||||
|
- Specific IPs are repeatedly blocked
|
||||||
|
|
||||||
|
### 5. Dynamic Response
|
||||||
|
Automatically block IPs that consistently exceed rate limits:
|
||||||
|
```go
|
||||||
|
// Check every minute
|
||||||
|
ticker := time.NewTicker(1 * time.Minute)
|
||||||
|
go func() {
|
||||||
|
for range ticker.C {
|
||||||
|
for _, info := range rateLimiter.GetAllRateLimitInfo() {
|
||||||
|
if info.TokensRemaining < 0.5 {
|
||||||
|
blacklist.BlockIP(info.IP, "Automated block: rate limit exceeded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
```
|
||||||
|
|
||||||
|
### 6. CIDR for Network Blocks
|
||||||
|
Use CIDR ranges to block entire networks efficiently:
|
||||||
|
```go
|
||||||
|
// Block entire subnets
|
||||||
|
blacklist.BlockCIDR("10.0.0.0/8", "Private network")
|
||||||
|
blacklist.BlockCIDR("192.168.0.0/16", "Local network")
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## API Reference
|
||||||
|
|
||||||
|
### IPBlacklist
|
||||||
|
|
||||||
|
#### Methods
|
||||||
|
- `BlockIP(ip, reason string) error` - Block a single IP address
|
||||||
|
- `BlockCIDR(cidr, reason string) error` - Block a CIDR range
|
||||||
|
- `UnblockIP(ip string)` - Remove IP from blacklist
|
||||||
|
- `UnblockCIDR(cidr string)` - Remove CIDR from blacklist
|
||||||
|
- `IsBlocked(ip string) (blocked bool, reason string)` - Check if IP is blocked
|
||||||
|
- `GetBlacklist() (ips, cidrs []string)` - Get all blocked IPs and CIDRs
|
||||||
|
- `Middleware(next http.Handler) http.Handler` - HTTP middleware
|
||||||
|
- `StatsHandler() http.Handler` - HTTP handler for statistics
|
||||||
|
|
||||||
|
### RateLimiter
|
||||||
|
|
||||||
|
#### Methods
|
||||||
|
- `GetTrackedIPs() []string` - Get all tracked IP addresses
|
||||||
|
- `GetRateLimitInfo(ip string) *RateLimitInfo` - Get info for specific IP
|
||||||
|
- `GetAllRateLimitInfo() []*RateLimitInfo` - Get info for all tracked IPs
|
||||||
|
- `Middleware(next http.Handler) http.Handler` - HTTP middleware
|
||||||
|
- `StatsHandler() http.Handler` - HTTP handler for statistics
|
||||||
|
|
||||||
|
#### RateLimitInfo Structure
|
||||||
|
```go
|
||||||
|
type RateLimitInfo struct {
|
||||||
|
IP string `json:"ip"`
|
||||||
|
TokensRemaining float64 `json:"tokens_remaining"`
|
||||||
|
Limit float64 `json:"limit"`
|
||||||
|
Burst int `json:"burst"`
|
||||||
|
}
|
||||||
|
```
|
||||||
+57
@@ -0,0 +1,57 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ChainAuthenticator tries each authenticator in order, returning the first success.
|
||||||
|
// Login and Logout are delegated to the primary authenticator.
|
||||||
|
type ChainAuthenticator struct {
|
||||||
|
authenticators []Authenticator
|
||||||
|
authenticateCallback func(r *http.Request) (*UserContext, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewChainAuthenticator creates a ChainAuthenticator from the given authenticators.
|
||||||
|
// At least one authenticator is required; the first is treated as primary for Login/Logout.
|
||||||
|
func NewChainAuthenticator(primary Authenticator, rest ...Authenticator) *ChainAuthenticator {
|
||||||
|
return &ChainAuthenticator{
|
||||||
|
authenticators: append([]Authenticator{primary}, rest...),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ChainAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
var lastErr error
|
||||||
|
for _, a := range c.authenticators {
|
||||||
|
if uc, err := a.Authenticate(r); err == nil {
|
||||||
|
return uc, nil
|
||||||
|
} else {
|
||||||
|
lastErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if c.authenticateCallback != nil {
|
||||||
|
return c.authenticateCallback(r)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("all authenticators failed; last error: %w", lastErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ChainAuthenticator) SetAuthenticateCallback(fn func(r *http.Request) (*UserContext, error)) {
|
||||||
|
c.authenticateCallback = fn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ChainAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||||
|
return c.authenticators[0].Login(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ChainAuthenticator) LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error) {
|
||||||
|
return c.authenticators[0].LoginWithCookie(ctx, req, w)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ChainAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
|
||||||
|
return c.authenticators[0].Logout(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *ChainAuthenticator) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
|
||||||
|
return c.authenticators[0].LogoutWithCookie(ctx, req, w)
|
||||||
|
}
|
||||||
+120
@@ -0,0 +1,120 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CompositeSecurityProvider combines multiple security providers
|
||||||
|
// Allows separating authentication, column security, and row security concerns
|
||||||
|
type CompositeSecurityProvider struct {
|
||||||
|
auth Authenticator
|
||||||
|
colSec ColumnSecurityProvider
|
||||||
|
rowSec RowSecurityProvider
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCompositeSecurityProvider creates a composite provider
|
||||||
|
// All parameters are required
|
||||||
|
func NewCompositeSecurityProvider(
|
||||||
|
auth Authenticator,
|
||||||
|
colSec ColumnSecurityProvider,
|
||||||
|
rowSec RowSecurityProvider,
|
||||||
|
) (*CompositeSecurityProvider, error) {
|
||||||
|
if auth == nil {
|
||||||
|
return nil, fmt.Errorf("authenticator cannot be nil")
|
||||||
|
}
|
||||||
|
if colSec == nil {
|
||||||
|
return nil, fmt.Errorf("column security provider cannot be nil")
|
||||||
|
}
|
||||||
|
if rowSec == nil {
|
||||||
|
return nil, fmt.Errorf("row security provider cannot be nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
return &CompositeSecurityProvider{
|
||||||
|
auth: auth,
|
||||||
|
colSec: colSec,
|
||||||
|
rowSec: rowSec,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Login delegates to the authenticator
|
||||||
|
func (c *CompositeSecurityProvider) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||||
|
return c.auth.Login(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoginWithCookie delegates to the authenticator
|
||||||
|
func (c *CompositeSecurityProvider) LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error) {
|
||||||
|
return c.auth.LoginWithCookie(ctx, req, w)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Logout delegates to the authenticator
|
||||||
|
func (c *CompositeSecurityProvider) Logout(ctx context.Context, req LogoutRequest) error {
|
||||||
|
return c.auth.Logout(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogoutWithCookie delegates to the authenticator
|
||||||
|
func (c *CompositeSecurityProvider) LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error {
|
||||||
|
return c.auth.LogoutWithCookie(ctx, req, w)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticate delegates to the authenticator
|
||||||
|
func (c *CompositeSecurityProvider) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
return c.auth.Authenticate(r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetAuthenticateCallback delegates to the authenticator
|
||||||
|
func (c *CompositeSecurityProvider) SetAuthenticateCallback(fn func(r *http.Request) (*UserContext, error)) {
|
||||||
|
c.auth.SetAuthenticateCallback(fn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetColumnSecurity delegates to the column security provider
|
||||||
|
func (c *CompositeSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
|
||||||
|
return c.colSec.GetColumnSecurity(ctx, userID, schema, table)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRowSecurity delegates to the row security provider
|
||||||
|
func (c *CompositeSecurityProvider) GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error) {
|
||||||
|
return c.rowSec.GetRowSecurity(ctx, userID, schema, table)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optional interface implementations (if wrapped providers support them)
|
||||||
|
|
||||||
|
// RefreshToken implements Refreshable if the authenticator supports it
|
||||||
|
func (c *CompositeSecurityProvider) RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error) {
|
||||||
|
if refreshable, ok := c.auth.(Refreshable); ok {
|
||||||
|
return refreshable.RefreshToken(ctx, refreshToken)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("authenticator does not support token refresh")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateToken implements Validatable if the authenticator supports it
|
||||||
|
func (c *CompositeSecurityProvider) ValidateToken(ctx context.Context, token string) (bool, error) {
|
||||||
|
if validatable, ok := c.auth.(Validatable); ok {
|
||||||
|
return validatable.ValidateToken(ctx, token)
|
||||||
|
}
|
||||||
|
return false, fmt.Errorf("authenticator does not support token validation")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearCache implements Cacheable if any provider supports it
|
||||||
|
func (c *CompositeSecurityProvider) ClearCache(ctx context.Context, userID int, schema, table string) error {
|
||||||
|
var errs []error
|
||||||
|
|
||||||
|
if cacheable, ok := c.colSec.(Cacheable); ok {
|
||||||
|
if err := cacheable.ClearCache(ctx, userID, schema, table); err != nil {
|
||||||
|
errs = append(errs, fmt.Errorf("column security cache clear failed: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cacheable, ok := c.rowSec.(Cacheable); ok {
|
||||||
|
if err := cacheable.ClearCache(ctx, userID, schema, table); err != nil {
|
||||||
|
errs = append(errs, fmt.Errorf("row security cache clear failed: %w", err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(errs) > 0 {
|
||||||
|
return fmt.Errorf("cache clear errors: %v", errs)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+1763
File diff suppressed because it is too large
Load Diff
+385
@@ -0,0 +1,385 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
// Optional: Uncomment if you want to use JWT authentication
|
||||||
|
// "github.com/golang-jwt/jwt/v5"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Example 1: Simple Header-Based Authenticator
|
||||||
|
// =============================================
|
||||||
|
|
||||||
|
type HeaderAuthenticatorExample struct {
|
||||||
|
// Optional: Add any dependencies here (e.g., database, cache)
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewHeaderAuthenticatorExample() *HeaderAuthenticatorExample {
|
||||||
|
return &HeaderAuthenticatorExample{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *HeaderAuthenticatorExample) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||||
|
// For header-based auth, login might not be used
|
||||||
|
// Could validate credentials against a database here
|
||||||
|
return nil, fmt.Errorf("header authentication does not support login")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *HeaderAuthenticatorExample) Logout(ctx context.Context, req LogoutRequest) error {
|
||||||
|
// For header-based auth, logout is a no-op
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *HeaderAuthenticatorExample) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
userIDStr := r.Header.Get("X-User-ID")
|
||||||
|
if userIDStr == "" {
|
||||||
|
return nil, fmt.Errorf("X-User-ID header required")
|
||||||
|
}
|
||||||
|
|
||||||
|
userID, err := strconv.Atoi(userIDStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid user ID: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &UserContext{
|
||||||
|
UserID: userID,
|
||||||
|
UserName: r.Header.Get("X-User-Name"),
|
||||||
|
UserLevel: parseIntHeader(r, "X-User-Level", 0),
|
||||||
|
SessionID: r.Header.Get("X-Session-ID"),
|
||||||
|
RemoteID: r.Header.Get("X-Remote-ID"),
|
||||||
|
Email: r.Header.Get("X-User-Email"),
|
||||||
|
Roles: parseRoles(r.Header.Get("X-User-Roles")),
|
||||||
|
Claims: make(map[string]any),
|
||||||
|
Meta: make(map[string]any),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 2: JWT Token Authenticator
|
||||||
|
// ====================================
|
||||||
|
// NOTE: To use this, uncomment the jwt import and install: go get github.com/golang-jwt/jwt/v5
|
||||||
|
|
||||||
|
type JWTAuthenticatorExample struct {
|
||||||
|
secretKey []byte
|
||||||
|
db *gorm.DB
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewJWTAuthenticatorExample(secretKey string, db *gorm.DB) *JWTAuthenticatorExample {
|
||||||
|
return &JWTAuthenticatorExample{
|
||||||
|
secretKey: []byte(secretKey),
|
||||||
|
db: db,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *JWTAuthenticatorExample) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||||
|
// Validate credentials against database
|
||||||
|
var user struct {
|
||||||
|
ID int
|
||||||
|
Username string
|
||||||
|
Email string
|
||||||
|
Password string // Should be hashed
|
||||||
|
UserLevel int
|
||||||
|
Roles string
|
||||||
|
}
|
||||||
|
|
||||||
|
err := a.db.WithContext(ctx).
|
||||||
|
Table("users").
|
||||||
|
Where("username = ?", req.Username).
|
||||||
|
First(&user).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid credentials")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: Verify password hash
|
||||||
|
// if !verifyPassword(user.Password, req.Password) {
|
||||||
|
// return nil, fmt.Errorf("invalid credentials")
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Create JWT token
|
||||||
|
expiresAt := time.Now().Add(24 * time.Hour)
|
||||||
|
|
||||||
|
// Uncomment when using JWT:
|
||||||
|
// token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
|
||||||
|
// "user_id": user.ID,
|
||||||
|
// "username": user.Username,
|
||||||
|
// "email": user.Email,
|
||||||
|
// "user_level": user.UserLevel,
|
||||||
|
// "roles": user.Roles,
|
||||||
|
// "exp": expiresAt.Unix(),
|
||||||
|
// })
|
||||||
|
// tokenString, err := token.SignedString(a.secretKey)
|
||||||
|
// if err != nil {
|
||||||
|
// return nil, fmt.Errorf("failed to generate token: %w", err)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Placeholder token for example (replace with actual JWT)
|
||||||
|
tokenString := fmt.Sprintf("token_%d_%d", user.ID, expiresAt.Unix())
|
||||||
|
|
||||||
|
return &LoginResponse{
|
||||||
|
Token: tokenString,
|
||||||
|
User: &UserContext{
|
||||||
|
UserID: user.ID,
|
||||||
|
UserName: user.Username,
|
||||||
|
Email: user.Email,
|
||||||
|
UserLevel: user.UserLevel,
|
||||||
|
Roles: parseRoles(user.Roles),
|
||||||
|
Claims: req.Claims,
|
||||||
|
Meta: req.Meta,
|
||||||
|
},
|
||||||
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *JWTAuthenticatorExample) Logout(ctx context.Context, req LogoutRequest) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *JWTAuthenticatorExample) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
authHeader := r.Header.Get("Authorization")
|
||||||
|
if authHeader == "" {
|
||||||
|
return nil, fmt.Errorf("authorization header required")
|
||||||
|
}
|
||||||
|
|
||||||
|
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
||||||
|
if tokenString == authHeader {
|
||||||
|
return nil, fmt.Errorf("bearer token required")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Uncomment when using JWT:
|
||||||
|
// token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
|
||||||
|
// if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||||
|
// return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
||||||
|
// }
|
||||||
|
// return a.secretKey, nil
|
||||||
|
// })
|
||||||
|
//
|
||||||
|
// if err != nil || !token.Valid {
|
||||||
|
// return nil, fmt.Errorf("invalid token: %w", err)
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// claims, ok := token.Claims.(jwt.MapClaims)
|
||||||
|
// if !ok {
|
||||||
|
// return nil, fmt.Errorf("invalid token claims")
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// return &UserContext{
|
||||||
|
// UserID: int(claims["user_id"].(float64)),
|
||||||
|
// UserName: getString(claims, "username"),
|
||||||
|
// Email: getString(claims, "email"),
|
||||||
|
// UserLevel: getInt(claims, "user_level"),
|
||||||
|
// Roles: parseRoles(getString(claims, "roles")),
|
||||||
|
// Claims: claims,
|
||||||
|
// }, nil
|
||||||
|
|
||||||
|
// Placeholder implementation (replace with actual JWT parsing)
|
||||||
|
return nil, fmt.Errorf("JWT parsing not implemented - uncomment JWT code above")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example 3: Database Session Authenticator
|
||||||
|
// ==========================================
|
||||||
|
|
||||||
|
type DatabaseAuthenticatorExample struct {
|
||||||
|
db *gorm.DB
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewDatabaseAuthenticatorExample(db *gorm.DB) *DatabaseAuthenticatorExample {
|
||||||
|
return &DatabaseAuthenticatorExample{db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *DatabaseAuthenticatorExample) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||||
|
// Query user from database
|
||||||
|
var user struct {
|
||||||
|
ID int
|
||||||
|
Username string
|
||||||
|
Email string
|
||||||
|
Password string // Should be hashed with bcrypt
|
||||||
|
UserLevel int
|
||||||
|
Roles string
|
||||||
|
IsActive bool
|
||||||
|
}
|
||||||
|
|
||||||
|
err := a.db.WithContext(ctx).
|
||||||
|
Table("users").
|
||||||
|
Where("username = ? AND is_active = true", req.Username).
|
||||||
|
First(&user).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid credentials")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: Verify password with bcrypt
|
||||||
|
// if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(req.Password)); err != nil {
|
||||||
|
// return nil, fmt.Errorf("invalid credentials")
|
||||||
|
// }
|
||||||
|
|
||||||
|
// Generate session token
|
||||||
|
sessionToken := fmt.Sprintf("sess_%s_%d", generateRandomString(32), time.Now().Unix())
|
||||||
|
expiresAt := time.Now().Add(24 * time.Hour)
|
||||||
|
|
||||||
|
// Create session in database
|
||||||
|
err = a.db.WithContext(ctx).Table("user_sessions").Create(map[string]any{
|
||||||
|
"session_token": sessionToken,
|
||||||
|
"user_id": user.ID,
|
||||||
|
"expires_at": expiresAt,
|
||||||
|
"created_at": time.Now(),
|
||||||
|
"ip_address": req.Claims["ip_address"],
|
||||||
|
"user_agent": req.Claims["user_agent"],
|
||||||
|
}).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &LoginResponse{
|
||||||
|
Token: sessionToken,
|
||||||
|
User: &UserContext{
|
||||||
|
UserID: user.ID,
|
||||||
|
UserName: user.Username,
|
||||||
|
Email: user.Email,
|
||||||
|
UserLevel: user.UserLevel,
|
||||||
|
Roles: parseRoles(user.Roles),
|
||||||
|
SessionID: sessionToken,
|
||||||
|
Claims: req.Claims,
|
||||||
|
Meta: req.Meta,
|
||||||
|
},
|
||||||
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *DatabaseAuthenticatorExample) Logout(ctx context.Context, req LogoutRequest) error {
|
||||||
|
// Delete session from database
|
||||||
|
err := a.db.WithContext(ctx).
|
||||||
|
Table("user_sessions").
|
||||||
|
Where("session_token = ? AND user_id = ?", req.Token, req.UserID).
|
||||||
|
Delete(nil).Error
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to delete session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *DatabaseAuthenticatorExample) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
// Extract session token from header or cookie
|
||||||
|
sessionToken := r.Header.Get("Authorization")
|
||||||
|
if sessionToken == "" {
|
||||||
|
// Try cookie
|
||||||
|
cookie, err := r.Cookie("session_token")
|
||||||
|
if err == nil {
|
||||||
|
sessionToken = cookie.Value
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Remove "Bearer " prefix if present
|
||||||
|
sessionToken = strings.TrimPrefix(sessionToken, "Bearer ")
|
||||||
|
}
|
||||||
|
|
||||||
|
if sessionToken == "" {
|
||||||
|
return nil, fmt.Errorf("session token required")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Query session and user from database
|
||||||
|
var session struct {
|
||||||
|
SessionToken string
|
||||||
|
UserID int
|
||||||
|
ExpiresAt time.Time
|
||||||
|
Username string
|
||||||
|
Email string
|
||||||
|
UserLevel int
|
||||||
|
Roles string
|
||||||
|
}
|
||||||
|
|
||||||
|
query := `
|
||||||
|
SELECT
|
||||||
|
s.session_token,
|
||||||
|
s.user_id,
|
||||||
|
s.expires_at,
|
||||||
|
u.username,
|
||||||
|
u.email,
|
||||||
|
u.user_level,
|
||||||
|
u.roles
|
||||||
|
FROM user_sessions s
|
||||||
|
JOIN users u ON s.user_id = u.id
|
||||||
|
WHERE s.session_token = ?
|
||||||
|
AND s.expires_at > ?
|
||||||
|
AND u.is_active = true
|
||||||
|
`
|
||||||
|
|
||||||
|
err := a.db.Raw(query, sessionToken, time.Now()).Scan(&session).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid or expired session")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update last activity timestamp
|
||||||
|
go a.updateSessionActivity(sessionToken)
|
||||||
|
|
||||||
|
return &UserContext{
|
||||||
|
UserID: session.UserID,
|
||||||
|
UserName: session.Username,
|
||||||
|
Email: session.Email,
|
||||||
|
UserLevel: session.UserLevel,
|
||||||
|
SessionID: sessionToken,
|
||||||
|
Roles: parseRoles(session.Roles),
|
||||||
|
Claims: make(map[string]any),
|
||||||
|
Meta: make(map[string]any),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// updateSessionActivity updates the last activity timestamp for the session
|
||||||
|
func (a *DatabaseAuthenticatorExample) updateSessionActivity(sessionToken string) {
|
||||||
|
a.db.Table("user_sessions").
|
||||||
|
Where("session_token = ?", sessionToken).
|
||||||
|
Update("last_activity_at", time.Now())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optional: Implement Refreshable interface
|
||||||
|
func (a *DatabaseAuthenticatorExample) RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error) {
|
||||||
|
// Query the refresh token
|
||||||
|
var session struct {
|
||||||
|
UserID int
|
||||||
|
Username string
|
||||||
|
Email string
|
||||||
|
}
|
||||||
|
|
||||||
|
err := a.db.WithContext(ctx).Raw(`
|
||||||
|
SELECT u.id as user_id, u.username, u.email
|
||||||
|
FROM user_sessions s
|
||||||
|
JOIN users u ON s.user_id = u.id
|
||||||
|
WHERE s.session_token = ? AND s.expires_at > ?
|
||||||
|
`, refreshToken, time.Now()).Scan(&session).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid refresh token")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate new session token
|
||||||
|
newSessionToken := fmt.Sprintf("sess_%s_%d", generateRandomString(32), time.Now().Unix())
|
||||||
|
expiresAt := time.Now().Add(24 * time.Hour)
|
||||||
|
|
||||||
|
// Create new session
|
||||||
|
err = a.db.WithContext(ctx).Table("user_sessions").Create(map[string]any{
|
||||||
|
"session_token": newSessionToken,
|
||||||
|
"user_id": session.UserID,
|
||||||
|
"expires_at": expiresAt,
|
||||||
|
"created_at": time.Now(),
|
||||||
|
}).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create new session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete old session
|
||||||
|
a.db.WithContext(ctx).Table("user_sessions").Where("session_token = ?", refreshToken).Delete(nil)
|
||||||
|
|
||||||
|
return &LoginResponse{
|
||||||
|
Token: newSessionToken,
|
||||||
|
User: &UserContext{
|
||||||
|
UserID: session.UserID,
|
||||||
|
UserName: session.Username,
|
||||||
|
Email: session.Email,
|
||||||
|
SessionID: newSessionToken,
|
||||||
|
Claims: make(map[string]any),
|
||||||
|
Meta: make(map[string]any),
|
||||||
|
},
|
||||||
|
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
+160
@@ -0,0 +1,160 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
// This file contains usage examples for integrating security with funcspec handlers
|
||||||
|
// These are example snippets - not executable code
|
||||||
|
|
||||||
|
/*
|
||||||
|
Example 1: Wrap handlers with authentication (required)
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/funcspec"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Setup
|
||||||
|
db := ... // your database connection
|
||||||
|
securityList := ... // your security list
|
||||||
|
handler := funcspec.NewHandler(db)
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Wrap handler with required authentication (returns 401 if not authenticated)
|
||||||
|
ordersHandler := security.WithAuth(
|
||||||
|
handler.SqlQueryList("SELECT * FROM orders WHERE user_id = [rid_user]", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/orders", ordersHandler).Methods("GET")
|
||||||
|
|
||||||
|
Example 2: Wrap handlers with optional authentication
|
||||||
|
|
||||||
|
// Wrap handler with optional authentication (falls back to guest if not authenticated)
|
||||||
|
productsHandler := security.WithOptionalAuth(
|
||||||
|
handler.SqlQueryList("SELECT * FROM products WHERE deleted = false", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/products", productsHandler).Methods("GET")
|
||||||
|
|
||||||
|
// The handler will show all products for guests, but could show personalized pricing
|
||||||
|
// or recommendations for authenticated users based on [rid_user]
|
||||||
|
|
||||||
|
Example 3: Wrap handlers with both authentication and security context
|
||||||
|
|
||||||
|
// Use the convenience function for both auth and security context
|
||||||
|
usersHandler := security.WithAuthAndSecurity(
|
||||||
|
handler.SqlQueryList("SELECT * FROM users WHERE active = true", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/users", usersHandler).Methods("GET")
|
||||||
|
|
||||||
|
// Or use WithOptionalAuthAndSecurity for optional auth
|
||||||
|
postsHandler := security.WithOptionalAuthAndSecurity(
|
||||||
|
handler.SqlQueryList("SELECT * FROM posts WHERE published = true", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/posts", postsHandler).Methods("GET")
|
||||||
|
|
||||||
|
Example 4: Wrap a single funcspec handler with security context only
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/funcspec"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Setup
|
||||||
|
db := ... // your database connection
|
||||||
|
securityList := ... // your security list
|
||||||
|
handler := funcspec.NewHandler(db)
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Wrap a specific handler with security context
|
||||||
|
usersHandler := security.WithSecurityContext(
|
||||||
|
handler.SqlQueryList("SELECT * FROM users WHERE active = true", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/users", usersHandler).Methods("GET")
|
||||||
|
|
||||||
|
Example 5: Wrap multiple handlers for different paths
|
||||||
|
|
||||||
|
// Products list endpoint
|
||||||
|
productsHandler := security.WithSecurityContext(
|
||||||
|
handler.SqlQueryList("SELECT * FROM products WHERE deleted = false", false, true, true),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/products", productsHandler).Methods("GET")
|
||||||
|
|
||||||
|
// Single product endpoint
|
||||||
|
productHandler := security.WithSecurityContext(
|
||||||
|
handler.SqlQuery("SELECT * FROM products WHERE id = [id]", true),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/products/{id}", productHandler).Methods("GET")
|
||||||
|
|
||||||
|
// Orders endpoint with user filtering
|
||||||
|
ordersHandler := security.WithSecurityContext(
|
||||||
|
handler.SqlQueryList("SELECT * FROM orders WHERE user_id = [rid_user]", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/orders", ordersHandler).Methods("GET")
|
||||||
|
|
||||||
|
Example 6: Helper function to wrap multiple handlers
|
||||||
|
|
||||||
|
// Create a helper function for your application
|
||||||
|
func secureHandler(h funcspec.HTTPFuncType, sl *SecurityList) funcspec.HTTPFuncType {
|
||||||
|
return security.WithSecurityContext(h, sl)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use it to wrap handlers
|
||||||
|
router.HandleFunc("/api/users", secureHandler(
|
||||||
|
handler.SqlQueryList("SELECT * FROM users", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)).Methods("GET")
|
||||||
|
|
||||||
|
router.HandleFunc("/api/roles", secureHandler(
|
||||||
|
handler.SqlQueryList("SELECT * FROM roles", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)).Methods("GET")
|
||||||
|
|
||||||
|
Example 7: Access SecurityList and user context in hooks
|
||||||
|
|
||||||
|
// In your funcspec hook, you can now access the SecurityList and user context
|
||||||
|
handler.Hooks().Register(funcspec.BeforeQueryList, func(ctx *funcspec.HookContext) error {
|
||||||
|
// Get SecurityList from context
|
||||||
|
if secList, ok := security.GetSecurityList(ctx.Context); ok {
|
||||||
|
// Use secList to apply security rules
|
||||||
|
// e.g., apply row-level security, column masking, etc.
|
||||||
|
_ = secList
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get user context
|
||||||
|
if userCtx, ok := security.GetUserContext(ctx.Context); ok {
|
||||||
|
// Access user information
|
||||||
|
logger.Info("User %s (ID: %d) accessing resource", userCtx.UserName, userCtx.UserID)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
|
||||||
|
Example 8: Mixing authentication and security patterns
|
||||||
|
|
||||||
|
// Public endpoint - no auth required, but has security context
|
||||||
|
publicHandler := security.WithSecurityContext(
|
||||||
|
handler.SqlQueryList("SELECT * FROM public_data", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/public", publicHandler).Methods("GET")
|
||||||
|
|
||||||
|
// Optional auth - personalized for logged-in users, works for guests
|
||||||
|
personalizedHandler := security.WithOptionalAuth(
|
||||||
|
handler.SqlQueryList("SELECT * FROM products WHERE category = [category]", false, true, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/products/category/{category}", personalizedHandler).Methods("GET")
|
||||||
|
|
||||||
|
// Required auth - must be logged in
|
||||||
|
privateHandler := security.WithAuthAndSecurity(
|
||||||
|
handler.SqlQueryList("SELECT * FROM private_data WHERE user_id = [rid_user]", false, false, false),
|
||||||
|
securityList,
|
||||||
|
)
|
||||||
|
router.HandleFunc("/api/private", privateHandler).Methods("GET")
|
||||||
|
*/
|
||||||
+385
@@ -0,0 +1,385 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SecurityContext is a generic interface that any spec can implement to integrate with security features
|
||||||
|
// This interface abstracts the common security context needs across different specs
|
||||||
|
type SecurityContext interface {
|
||||||
|
GetContext() context.Context
|
||||||
|
GetUserID() (int, bool)
|
||||||
|
GetSchema() string
|
||||||
|
GetEntity() string
|
||||||
|
GetModel() interface{}
|
||||||
|
GetQuery() interface{}
|
||||||
|
SetQuery(interface{})
|
||||||
|
GetResult() interface{}
|
||||||
|
SetResult(interface{})
|
||||||
|
}
|
||||||
|
|
||||||
|
// loadSecurityRules loads security configuration for the user and entity (generic version)
|
||||||
|
func loadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
// Extract user ID from context
|
||||||
|
userID, ok := secCtx.GetUserID()
|
||||||
|
if !ok {
|
||||||
|
logger.Warn("No user ID in context for security check")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := secCtx.GetSchema()
|
||||||
|
tablename := secCtx.GetEntity()
|
||||||
|
|
||||||
|
logger.Debug("Loading security rules for user=%d, schema=%s, table=%s", userID, schema, tablename)
|
||||||
|
|
||||||
|
// Load column security rules using the provider
|
||||||
|
err := securityList.LoadColumnSecurity(secCtx.GetContext(), userID, schema, tablename, false)
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Failed to load column security: %v", err)
|
||||||
|
// Don't fail the request if no security rules exist
|
||||||
|
// return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load row security rules using the provider
|
||||||
|
_, err = securityList.LoadRowSecurity(secCtx.GetContext(), userID, schema, tablename, false)
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Failed to load row security: %v", err)
|
||||||
|
// Don't fail the request if no security rules exist
|
||||||
|
// return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// applyRowSecurity applies row-level security filters to the query (generic version)
|
||||||
|
func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
userID, ok := secCtx.GetUserID()
|
||||||
|
if !ok {
|
||||||
|
return nil // No user context, skip
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := secCtx.GetSchema()
|
||||||
|
tablename := secCtx.GetEntity()
|
||||||
|
|
||||||
|
// Get row security template
|
||||||
|
rowSec, err := securityList.GetRowSecurityTemplate(userID, schema, tablename)
|
||||||
|
if err != nil {
|
||||||
|
// No row security defined, allow query to proceed
|
||||||
|
logger.Debug("No row security for %s.%s@%d: %v", schema, tablename, userID, err)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if user has a blocking rule
|
||||||
|
if rowSec.HasBlock {
|
||||||
|
logger.Warn("User %d blocked from accessing %s.%s", userID, schema, tablename)
|
||||||
|
return fmt.Errorf("access denied to %s", tablename)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If there's a security template, apply it as a WHERE clause
|
||||||
|
if rowSec.Template != "" {
|
||||||
|
model := secCtx.GetModel()
|
||||||
|
if model == nil {
|
||||||
|
logger.Debug("No model available for row security on %s.%s", schema, tablename)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get primary key name from model
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
if modelType.Kind() == reflect.Pointer {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find primary key field
|
||||||
|
pkName := "id" // default
|
||||||
|
for i := 0; i < modelType.NumField(); i++ {
|
||||||
|
field := modelType.Field(i)
|
||||||
|
if tag := field.Tag.Get("bun"); tag != "" {
|
||||||
|
// Check for primary key tag
|
||||||
|
if contains(tag, "pk") || contains(tag, "primary_key") {
|
||||||
|
if sqlName := extractSQLName(tag); sqlName != "" {
|
||||||
|
pkName = sqlName
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate the WHERE clause from template
|
||||||
|
whereClause := rowSec.GetTemplate(pkName, modelType)
|
||||||
|
|
||||||
|
logger.Info("Applying row security filter for user %d on %s.%s: %s",
|
||||||
|
userID, schema, tablename, whereClause)
|
||||||
|
|
||||||
|
// Apply the WHERE clause to the query
|
||||||
|
query := secCtx.GetQuery()
|
||||||
|
if selectQuery, ok := query.(interface {
|
||||||
|
Where(string, ...interface{}) interface{}
|
||||||
|
}); ok {
|
||||||
|
secCtx.SetQuery(selectQuery.Where(whereClause))
|
||||||
|
} else {
|
||||||
|
logger.Debug("Query doesn't support Where method, skipping row security")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// applyColumnSecurity applies column-level security (masking/hiding) to results (generic version)
|
||||||
|
func applyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
userID, ok := secCtx.GetUserID()
|
||||||
|
if !ok {
|
||||||
|
return nil // No user context, skip
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := secCtx.GetSchema()
|
||||||
|
tablename := secCtx.GetEntity()
|
||||||
|
|
||||||
|
// Get result data
|
||||||
|
result := secCtx.GetResult()
|
||||||
|
if result == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Debug("Applying column security for user=%d, schema=%s, table=%s", userID, schema, tablename)
|
||||||
|
|
||||||
|
model := secCtx.GetModel()
|
||||||
|
if model == nil {
|
||||||
|
logger.Debug("No model available for column security on %s.%s", schema, tablename)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get model type
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
if modelType.Kind() == reflect.Pointer {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply column security masking
|
||||||
|
resultValue := reflect.ValueOf(result)
|
||||||
|
if resultValue.Kind() == reflect.Pointer {
|
||||||
|
resultValue = resultValue.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
maskedResult, err := securityList.ApplyColumnSecurity(resultValue, modelType, userID, schema, tablename)
|
||||||
|
if err != nil {
|
||||||
|
logger.Warn("Column security error: %v", err)
|
||||||
|
// Don't fail the request, just log the issue
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update the result with masked data
|
||||||
|
if maskedResult.IsValid() && maskedResult.CanInterface() {
|
||||||
|
secCtx.SetResult(maskedResult.Interface())
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// logDataAccess logs all data access for audit purposes (generic version)
|
||||||
|
func logDataAccess(secCtx SecurityContext) error {
|
||||||
|
userID, _ := secCtx.GetUserID()
|
||||||
|
|
||||||
|
logger.Info("AUDIT: User %d accessed %s.%s",
|
||||||
|
userID,
|
||||||
|
secCtx.GetSchema(),
|
||||||
|
secCtx.GetEntity(),
|
||||||
|
)
|
||||||
|
|
||||||
|
// TODO: Write to audit log table or external audit service
|
||||||
|
// auditLog := AuditLog{
|
||||||
|
// UserID: userID,
|
||||||
|
// Schema: secCtx.GetSchema(),
|
||||||
|
// Entity: secCtx.GetEntity(),
|
||||||
|
// Action: "READ",
|
||||||
|
// Timestamp: time.Now(),
|
||||||
|
// }
|
||||||
|
// db.Create(&auditLog)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogDataAccess is a public wrapper for logDataAccess that accepts a SecurityContext
|
||||||
|
// This allows other packages to use the audit logging functionality
|
||||||
|
func LogDataAccess(secCtx SecurityContext) error {
|
||||||
|
return logDataAccess(secCtx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoadSecurityRules is a public wrapper for loadSecurityRules that accepts a SecurityContext
|
||||||
|
// This allows other packages to load security rules using the generic interface
|
||||||
|
func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
return loadSecurityRules(secCtx, securityList)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
|
||||||
|
// This allows other packages to apply row-level security using the generic interface
|
||||||
|
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
return applyRowSecurity(secCtx, securityList)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplyColumnSecurity is a public wrapper for applyColumnSecurity that accepts a SecurityContext
|
||||||
|
// This allows other packages to apply column-level security using the generic interface
|
||||||
|
func ApplyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
return applyColumnSecurity(secCtx, securityList)
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkModelUpdateAllowed returns an error if CanUpdate is false for the model.
|
||||||
|
// Rules are read from context (set by NewModelAuthMiddleware) with a fallback to the model registry.
|
||||||
|
func checkModelUpdateAllowed(secCtx SecurityContext) error {
|
||||||
|
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
|
||||||
|
if !ok {
|
||||||
|
schema := secCtx.GetSchema()
|
||||||
|
entity := secCtx.GetEntity()
|
||||||
|
var err error
|
||||||
|
if schema != "" {
|
||||||
|
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
|
||||||
|
}
|
||||||
|
if err != nil || schema == "" {
|
||||||
|
rules, err = modelregistry.GetModelRulesByName(entity)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil // model not registered, allow by default
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !rules.CanUpdate {
|
||||||
|
return fmt.Errorf("update not allowed for %s", secCtx.GetEntity())
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkModelDeleteAllowed returns an error if CanDelete is false for the model.
|
||||||
|
// Rules are read from context (set by NewModelAuthMiddleware) with a fallback to the model registry.
|
||||||
|
func checkModelDeleteAllowed(secCtx SecurityContext) error {
|
||||||
|
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
|
||||||
|
if !ok {
|
||||||
|
schema := secCtx.GetSchema()
|
||||||
|
entity := secCtx.GetEntity()
|
||||||
|
var err error
|
||||||
|
if schema != "" {
|
||||||
|
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
|
||||||
|
}
|
||||||
|
if err != nil || schema == "" {
|
||||||
|
rules, err = modelregistry.GetModelRulesByName(entity)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil // model not registered, allow by default
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !rules.CanDelete {
|
||||||
|
return fmt.Errorf("delete not allowed for %s", secCtx.GetEntity())
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CheckModelAuthAllowed checks whether the requested operation is permitted based on
|
||||||
|
// model rules and the current user's authentication state. It is intended for use in
|
||||||
|
// a BeforeHandle hook, fired after model resolution.
|
||||||
|
//
|
||||||
|
// Logic:
|
||||||
|
// 1. Load model rules from context (set by NewModelAuthMiddleware) or fall back to registry.
|
||||||
|
// 2. SecurityDisabled → allow.
|
||||||
|
// 3. operation == "read" && CanPublicRead → allow.
|
||||||
|
// 4. operation == "create" && CanPublicCreate → allow.
|
||||||
|
// 5. operation == "update" && CanPublicUpdate → allow.
|
||||||
|
// 6. operation == "delete" && CanPublicDelete → allow.
|
||||||
|
// 7. Guest (UserID == 0) → return "authentication required".
|
||||||
|
// 8. Authenticated user → allow (operation-specific checks remain in BeforeUpdate/BeforeDelete).
|
||||||
|
func CheckModelAuthAllowed(secCtx SecurityContext, operation string) error {
|
||||||
|
rules, ok := GetModelRulesFromContext(secCtx.GetContext())
|
||||||
|
if !ok {
|
||||||
|
schema := secCtx.GetSchema()
|
||||||
|
entity := secCtx.GetEntity()
|
||||||
|
var err error
|
||||||
|
if schema != "" {
|
||||||
|
rules, err = modelregistry.GetModelRulesByName(fmt.Sprintf("%s.%s", schema, entity))
|
||||||
|
}
|
||||||
|
if err != nil || schema == "" {
|
||||||
|
rules, err = modelregistry.GetModelRulesByName(entity)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
// Model not registered - fall through to auth check
|
||||||
|
userID, _ := secCtx.GetUserID()
|
||||||
|
if userID == 0 {
|
||||||
|
return fmt.Errorf("authentication required")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if rules.SecurityDisabled {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if operation == "read" && rules.CanPublicRead {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if operation == "create" && rules.CanPublicCreate {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if operation == "update" && rules.CanPublicUpdate {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if operation == "delete" && rules.CanPublicDelete {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
userID, _ := secCtx.GetUserID()
|
||||||
|
if userID == 0 {
|
||||||
|
return fmt.Errorf("authentication required")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed.
|
||||||
|
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
|
||||||
|
return checkModelUpdateAllowed(secCtx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CheckModelDeleteAllowed is the public wrapper for checkModelDeleteAllowed.
|
||||||
|
func CheckModelDeleteAllowed(secCtx SecurityContext) error {
|
||||||
|
return checkModelDeleteAllowed(secCtx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper functions
|
||||||
|
|
||||||
|
func contains(s, substr string) bool {
|
||||||
|
return len(s) >= len(substr) && s[:len(substr)] == substr ||
|
||||||
|
len(s) > len(substr) && s[len(s)-len(substr):] == substr
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractSQLName(tag string) string {
|
||||||
|
// Simple parser for "column:name" or just "name"
|
||||||
|
// This is a simplified version
|
||||||
|
parts := splitTag(tag, ',')
|
||||||
|
for _, part := range parts {
|
||||||
|
if part != "" && !contains(part, ":") {
|
||||||
|
return part
|
||||||
|
}
|
||||||
|
if contains(part, "column:") {
|
||||||
|
return part[7:] // Skip "column:"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func splitTag(tag string, sep rune) []string {
|
||||||
|
var parts []string
|
||||||
|
var current string
|
||||||
|
for _, ch := range tag {
|
||||||
|
if ch == sep {
|
||||||
|
if current != "" {
|
||||||
|
parts = append(parts, current)
|
||||||
|
current = ""
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
current += string(ch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if current != "" {
|
||||||
|
parts = append(parts, current)
|
||||||
|
}
|
||||||
|
return parts
|
||||||
|
}
|
||||||
+162
@@ -0,0 +1,162 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
// UserContext holds authenticated user information
|
||||||
|
type UserContext struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
UserName string `json:"user_name"`
|
||||||
|
UserLevel int `json:"user_level"`
|
||||||
|
SessionID string `json:"session_id"`
|
||||||
|
SessionRID int64 `json:"session_rid"`
|
||||||
|
RemoteID string `json:"remote_id"`
|
||||||
|
Roles []string `json:"roles"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
Claims map[string]any `json:"claims"`
|
||||||
|
Meta map[string]any `json:"meta"` // Additional metadata that can hold any JSON-serializable values
|
||||||
|
TwoFactorEnabled bool `json:"two_factor_enabled"` // Indicates if 2FA is enabled for this user
|
||||||
|
ProgramUserID int `json:"program_user_id"`
|
||||||
|
ProgramUserTable string `json:"program_user_table"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoginRequest contains credentials for login
|
||||||
|
type LoginRequest struct {
|
||||||
|
Username string `json:"username"`
|
||||||
|
Password string `json:"password"`
|
||||||
|
TwoFactorCode string `json:"two_factor_code,omitempty"` // TOTP or backup code
|
||||||
|
Claims map[string]any `json:"claims"` // Additional login data
|
||||||
|
Meta map[string]any `json:"meta"` // Additional metadata to be set on user context
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterRequest contains information for new user registration
|
||||||
|
type RegisterRequest struct {
|
||||||
|
Username string `json:"username"`
|
||||||
|
Password string `json:"password"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
UserLevel int `json:"user_level"`
|
||||||
|
Roles []string `json:"roles"`
|
||||||
|
Claims map[string]any `json:"claims"` // Additional registration data
|
||||||
|
Meta map[string]any `json:"meta"` // Additional metadata
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoginResponse contains the result of a login attempt
|
||||||
|
type LoginResponse struct {
|
||||||
|
Token string `json:"token"`
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
User *UserContext `json:"user"`
|
||||||
|
ExpiresIn int64 `json:"expires_in"` // Token expiration in seconds
|
||||||
|
Requires2FA bool `json:"requires_2fa"` // True if 2FA code is required
|
||||||
|
TwoFactorSetupData *TwoFactorSecret `json:"two_factor_setup,omitempty"` // Present when setting up 2FA
|
||||||
|
Meta map[string]any `json:"meta"` // Additional metadata to be set on user context
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogoutRequest contains information for logout
|
||||||
|
type LogoutRequest struct {
|
||||||
|
Token string `json:"token"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasswordResetRequest initiates a password reset for a user
|
||||||
|
type PasswordResetRequest struct {
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
Username string `json:"username,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasswordResetResponse is returned when a reset is initiated
|
||||||
|
type PasswordResetResponse struct {
|
||||||
|
// Token is the reset token to be delivered out-of-band (e.g. email).
|
||||||
|
// The stored procedure may return it for delivery or leave it empty
|
||||||
|
// if the delivery is handled entirely in the database.
|
||||||
|
Token string `json:"token"`
|
||||||
|
ExpiresIn int64 `json:"expires_in"` // seconds
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasswordResetCompleteRequest completes a password reset using the token
|
||||||
|
type PasswordResetCompleteRequest struct {
|
||||||
|
Token string `json:"token"`
|
||||||
|
NewPassword string `json:"new_password"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticator handles user authentication operations
|
||||||
|
type Authenticator interface {
|
||||||
|
// Login authenticates credentials and returns a token
|
||||||
|
Login(ctx context.Context, req LoginRequest) (*LoginResponse, error)
|
||||||
|
|
||||||
|
// LoginWithCookie authenticates credentials and, when cookie sessions are enabled,
|
||||||
|
// writes the session cookie to w. Implementations that do not support cookies
|
||||||
|
// should delegate to Login and ignore w.
|
||||||
|
LoginWithCookie(ctx context.Context, req LoginRequest, w http.ResponseWriter) (*LoginResponse, error)
|
||||||
|
|
||||||
|
// Logout invalidates a user's session/token
|
||||||
|
Logout(ctx context.Context, req LogoutRequest) error
|
||||||
|
|
||||||
|
// LogoutWithCookie invalidates a user's session/token and, when cookie sessions are
|
||||||
|
// enabled, clears the session cookie on w. Implementations that do not support cookies
|
||||||
|
// should delegate to Logout and ignore w.
|
||||||
|
LogoutWithCookie(ctx context.Context, req LogoutRequest, w http.ResponseWriter) error
|
||||||
|
|
||||||
|
// Authenticate extracts and validates user from HTTP request
|
||||||
|
// Returns UserContext or error if authentication fails
|
||||||
|
Authenticate(r *http.Request) (*UserContext, error)
|
||||||
|
|
||||||
|
// SetAuthenticateCallback registers a fallback called when primary authentication fails.
|
||||||
|
// If the callback returns a non-nil UserContext, that result is used instead of the error.
|
||||||
|
SetAuthenticateCallback(fn func(r *http.Request) (*UserContext, error))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Registrable allows providers to support user registration
|
||||||
|
type Registrable interface {
|
||||||
|
// Register creates a new user account
|
||||||
|
Register(ctx context.Context, req RegisterRequest) (*LoginResponse, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ColumnSecurityProvider handles column-level security (masking/hiding)
|
||||||
|
type ColumnSecurityProvider interface {
|
||||||
|
// GetColumnSecurity loads column security rules for a user and entity
|
||||||
|
GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RowSecurityProvider handles row-level security (filtering)
|
||||||
|
type RowSecurityProvider interface {
|
||||||
|
// GetRowSecurity loads row security rules for a user and entity
|
||||||
|
GetRowSecurity(ctx context.Context, userID int, schema, table string) (RowSecurity, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SecurityProvider is the main interface combining all security concerns
|
||||||
|
type SecurityProvider interface {
|
||||||
|
Authenticator
|
||||||
|
ColumnSecurityProvider
|
||||||
|
RowSecurityProvider
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optional interfaces for advanced functionality
|
||||||
|
|
||||||
|
// Refreshable allows providers to support token refresh
|
||||||
|
type Refreshable interface {
|
||||||
|
// RefreshToken exchanges a refresh token for a new access token
|
||||||
|
RefreshToken(ctx context.Context, refreshToken string) (*LoginResponse, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validatable allows providers to validate tokens without full authentication
|
||||||
|
type Validatable interface {
|
||||||
|
// ValidateToken checks if a token is valid without extracting full user context
|
||||||
|
ValidateToken(ctx context.Context, token string) (bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cacheable allows providers to support caching of security rules
|
||||||
|
type Cacheable interface {
|
||||||
|
// ClearCache clears cached security rules for a user/entity
|
||||||
|
ClearCache(ctx context.Context, userID int, schema, table string) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasswordResettable allows providers to support self-service password reset
|
||||||
|
type PasswordResettable interface {
|
||||||
|
// RequestPasswordReset creates a reset token for the given email/username
|
||||||
|
RequestPasswordReset(ctx context.Context, req PasswordResetRequest) (*PasswordResetResponse, error)
|
||||||
|
|
||||||
|
// CompletePasswordReset validates the token and sets the new password
|
||||||
|
CompletePasswordReset(ctx context.Context, req PasswordResetCompleteRequest) error
|
||||||
|
}
|
||||||
+81
@@ -0,0 +1,81 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// hashSHA256Hex returns the lowercase hex SHA-256 digest of the given string.
|
||||||
|
// Used by all keystore implementations to hash raw keys before storage or lookup.
|
||||||
|
func hashSHA256Hex(raw string) string {
|
||||||
|
sum := sha256.Sum256([]byte(raw))
|
||||||
|
return hex.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyType identifies the category of an auth key.
|
||||||
|
type KeyType string
|
||||||
|
|
||||||
|
const (
|
||||||
|
// KeyTypeJWTSecret is a per-user JWT signing secret for token generation.
|
||||||
|
KeyTypeJWTSecret KeyType = "jwt_secret"
|
||||||
|
// KeyTypeHeaderAPI is a static API key sent via a request header.
|
||||||
|
KeyTypeHeaderAPI KeyType = "header_api"
|
||||||
|
// KeyTypeOAuth2 holds OAuth2 client credentials (client_id / client_secret).
|
||||||
|
KeyTypeOAuth2 KeyType = "oauth2"
|
||||||
|
// KeyTypeGenericAPI is a generic application API key.
|
||||||
|
KeyTypeGenericAPI KeyType = "api"
|
||||||
|
)
|
||||||
|
|
||||||
|
// UserKey represents a single named auth key belonging to a user.
|
||||||
|
// KeyHash stores the SHA-256 hex digest of the raw key; the raw key is never persisted.
|
||||||
|
type UserKey struct {
|
||||||
|
ID int64 `json:"id"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
KeyType KeyType `json:"key_type"`
|
||||||
|
KeyHash string `json:"key_hash"` // SHA-256 hex; never the raw key
|
||||||
|
Name string `json:"name"`
|
||||||
|
Scopes []string `json:"scopes,omitempty"`
|
||||||
|
Meta map[string]any `json:"meta,omitempty"`
|
||||||
|
ExpiresAt *time.Time `json:"expires_at,omitempty"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
LastUsedAt *time.Time `json:"last_used_at,omitempty"`
|
||||||
|
IsActive bool `json:"is_active"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateKeyRequest specifies the parameters for a new key.
|
||||||
|
type CreateKeyRequest struct {
|
||||||
|
UserID int
|
||||||
|
KeyType KeyType
|
||||||
|
Name string
|
||||||
|
Scopes []string
|
||||||
|
Meta map[string]any
|
||||||
|
ExpiresAt *time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateKeyResponse is returned exactly once when a key is created.
|
||||||
|
// The caller is responsible for persisting RawKey; it is not stored anywhere.
|
||||||
|
type CreateKeyResponse struct {
|
||||||
|
Key UserKey
|
||||||
|
RawKey string // crypto/rand 32 bytes, base64url-encoded
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyStore manages per-user auth keys with pluggable storage backends.
|
||||||
|
// Implementations: ConfigKeyStore (static list) and DatabaseKeyStore (stored procedures).
|
||||||
|
type KeyStore interface {
|
||||||
|
// CreateKey generates a new key, stores its hash, and returns the raw key once.
|
||||||
|
CreateKey(ctx context.Context, req CreateKeyRequest) (*CreateKeyResponse, error)
|
||||||
|
|
||||||
|
// GetUserKeys returns all active, non-expired keys for a user.
|
||||||
|
// Pass an empty KeyType to return all types.
|
||||||
|
GetUserKeys(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error)
|
||||||
|
|
||||||
|
// DeleteKey soft-deletes a key by ID after verifying ownership.
|
||||||
|
DeleteKey(ctx context.Context, userID int, keyID int64) error
|
||||||
|
|
||||||
|
// ValidateKey checks a raw key, returns the matching UserKey on success.
|
||||||
|
// The implementation hashes the raw key before any lookup.
|
||||||
|
// Pass an empty KeyType to accept any type.
|
||||||
|
ValidateKey(ctx context.Context, rawKey string, keyType KeyType) (*UserKey, error)
|
||||||
|
}
|
||||||
+119
@@ -0,0 +1,119 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// KeyStoreAuthenticator implements the Authenticator interface using a KeyStore.
|
||||||
|
// It is suitable for long-lived application credentials (API keys, JWT secrets, etc.)
|
||||||
|
// rather than interactive sessions. Login and Logout are not supported — key lifecycle
|
||||||
|
// is managed directly through the KeyStore.
|
||||||
|
//
|
||||||
|
// Key extraction order:
|
||||||
|
// 1. Authorization: Bearer <key>
|
||||||
|
// 2. Authorization: ApiKey <key>
|
||||||
|
// 3. X-API-Key header
|
||||||
|
type KeyStoreAuthenticator struct {
|
||||||
|
keyStore KeyStore
|
||||||
|
keyType KeyType // empty = accept any type
|
||||||
|
authenticateCallback func(r *http.Request) (*UserContext, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewKeyStoreAuthenticator creates a KeyStoreAuthenticator.
|
||||||
|
// Pass an empty keyType to accept keys of any type.
|
||||||
|
func NewKeyStoreAuthenticator(ks KeyStore, keyType KeyType) *KeyStoreAuthenticator {
|
||||||
|
return &KeyStoreAuthenticator{keyStore: ks, keyType: keyType}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Login is not supported for keystore authentication.
|
||||||
|
func (a *KeyStoreAuthenticator) Login(_ context.Context, _ LoginRequest) (*LoginResponse, error) {
|
||||||
|
return nil, fmt.Errorf("keystore authenticator does not support login")
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoginWithCookie is not supported for keystore authentication.
|
||||||
|
func (a *KeyStoreAuthenticator) LoginWithCookie(_ context.Context, _ LoginRequest, _ http.ResponseWriter) (*LoginResponse, error) {
|
||||||
|
return nil, fmt.Errorf("keystore authenticator does not support login")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Logout is not supported for keystore authentication.
|
||||||
|
func (a *KeyStoreAuthenticator) Logout(_ context.Context, _ LogoutRequest) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogoutWithCookie is not supported for keystore authentication.
|
||||||
|
func (a *KeyStoreAuthenticator) LogoutWithCookie(_ context.Context, _ LogoutRequest, _ http.ResponseWriter) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetAuthenticateCallback registers a fallback called when key authentication fails.
|
||||||
|
func (a *KeyStoreAuthenticator) SetAuthenticateCallback(fn func(r *http.Request) (*UserContext, error)) {
|
||||||
|
a.authenticateCallback = fn
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticate extracts an API key from the request and validates it against the KeyStore.
|
||||||
|
// Returns a UserContext built from the matching UserKey on success.
|
||||||
|
func (a *KeyStoreAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
rawKey := extractAPIKey(r)
|
||||||
|
if rawKey == "" {
|
||||||
|
if a.authenticateCallback != nil {
|
||||||
|
return a.authenticateCallback(r)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("API key required (Authorization: Bearer/ApiKey <key> or X-API-Key header)")
|
||||||
|
}
|
||||||
|
|
||||||
|
userKey, err := a.keyStore.ValidateKey(r.Context(), rawKey, a.keyType)
|
||||||
|
if err != nil {
|
||||||
|
if a.authenticateCallback != nil {
|
||||||
|
return a.authenticateCallback(r)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("invalid API key: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return userKeyToUserContext(userKey), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractAPIKey extracts a raw key from the request using the following precedence:
|
||||||
|
// 1. Authorization: Bearer <key>
|
||||||
|
// 2. Authorization: ApiKey <key>
|
||||||
|
// 3. X-API-Key header
|
||||||
|
func extractAPIKey(r *http.Request) string {
|
||||||
|
if auth := r.Header.Get("Authorization"); auth != "" {
|
||||||
|
if after, ok := strings.CutPrefix(auth, "Bearer "); ok {
|
||||||
|
return strings.TrimSpace(after)
|
||||||
|
}
|
||||||
|
if after, ok := strings.CutPrefix(auth, "ApiKey "); ok {
|
||||||
|
return strings.TrimSpace(after)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(r.Header.Get("X-API-Key"))
|
||||||
|
}
|
||||||
|
|
||||||
|
// userKeyToUserContext converts a UserKey into a UserContext.
|
||||||
|
// Scopes are mapped to Roles. Key type and name are stored in Claims.
|
||||||
|
func userKeyToUserContext(k *UserKey) *UserContext {
|
||||||
|
claims := map[string]any{
|
||||||
|
"key_type": string(k.KeyType),
|
||||||
|
"key_name": k.Name,
|
||||||
|
}
|
||||||
|
|
||||||
|
meta := k.Meta
|
||||||
|
if meta == nil {
|
||||||
|
meta = map[string]any{}
|
||||||
|
}
|
||||||
|
|
||||||
|
roles := k.Scopes
|
||||||
|
if roles == nil {
|
||||||
|
roles = []string{}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &UserContext{
|
||||||
|
UserID: k.UserID,
|
||||||
|
SessionID: fmt.Sprintf("key:%d", k.ID),
|
||||||
|
Roles: roles,
|
||||||
|
Claims: claims,
|
||||||
|
Meta: meta,
|
||||||
|
}
|
||||||
|
}
|
||||||
+149
@@ -0,0 +1,149 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/subtle"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/hex"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ConfigKeyStore is an in-memory keystore backed by a static slice of UserKey values.
|
||||||
|
// It is designed for config-file driven setups (e.g. service accounts defined in YAML)
|
||||||
|
// with a small, bounded number of keys. For large or dynamic key sets use DatabaseKeyStore.
|
||||||
|
//
|
||||||
|
// Pre-existing entries must have KeyHash set to the SHA-256 hex of the intended raw key.
|
||||||
|
// Keys created at runtime via CreateKey are held in memory only and lost on restart.
|
||||||
|
type ConfigKeyStore struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
keys []UserKey
|
||||||
|
next int64 // monotonic ID counter for runtime-created keys (atomic)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewConfigKeyStore creates a ConfigKeyStore seeded with the provided keys.
|
||||||
|
// Pass nil or an empty slice to start with no pre-loaded keys.
|
||||||
|
// Zero-value entries (CreatedAt is zero) are treated as active and assigned the current time.
|
||||||
|
func NewConfigKeyStore(keys []UserKey) *ConfigKeyStore {
|
||||||
|
var maxID int64
|
||||||
|
copied := make([]UserKey, len(keys))
|
||||||
|
copy(copied, keys)
|
||||||
|
for i := range copied {
|
||||||
|
if copied[i].CreatedAt.IsZero() {
|
||||||
|
copied[i].IsActive = true
|
||||||
|
copied[i].CreatedAt = time.Now()
|
||||||
|
}
|
||||||
|
if copied[i].ID > maxID {
|
||||||
|
maxID = copied[i].ID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return &ConfigKeyStore{keys: copied, next: maxID}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateKey generates a new raw key, stores its SHA-256 hash, and returns the raw key once.
|
||||||
|
func (s *ConfigKeyStore) CreateKey(_ context.Context, req CreateKeyRequest) (*CreateKeyResponse, error) {
|
||||||
|
rawBytes := make([]byte, 32)
|
||||||
|
if _, err := rand.Read(rawBytes); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate key material: %w", err)
|
||||||
|
}
|
||||||
|
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
||||||
|
hash := hashSHA256Hex(rawKey)
|
||||||
|
|
||||||
|
id := atomic.AddInt64(&s.next, 1)
|
||||||
|
key := UserKey{
|
||||||
|
ID: id,
|
||||||
|
UserID: req.UserID,
|
||||||
|
KeyType: req.KeyType,
|
||||||
|
KeyHash: hash,
|
||||||
|
Name: req.Name,
|
||||||
|
Scopes: req.Scopes,
|
||||||
|
Meta: req.Meta,
|
||||||
|
ExpiresAt: req.ExpiresAt,
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
IsActive: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
s.keys = append(s.keys, key)
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
return &CreateKeyResponse{Key: key, RawKey: rawKey}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserKeys returns all active, non-expired keys for the given user.
|
||||||
|
// Pass an empty KeyType to return all types.
|
||||||
|
func (s *ConfigKeyStore) GetUserKeys(_ context.Context, userID int, keyType KeyType) ([]UserKey, error) {
|
||||||
|
now := time.Now()
|
||||||
|
s.mu.RLock()
|
||||||
|
defer s.mu.RUnlock()
|
||||||
|
|
||||||
|
var result []UserKey
|
||||||
|
for i := range s.keys {
|
||||||
|
k := &s.keys[i]
|
||||||
|
if k.UserID != userID || !k.IsActive {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if k.ExpiresAt != nil && k.ExpiresAt.Before(now) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if keyType != "" && k.KeyType != keyType {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
result = append(result, *k)
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteKey soft-deletes a key by setting IsActive to false after ownership verification.
|
||||||
|
func (s *ConfigKeyStore) DeleteKey(_ context.Context, userID int, keyID int64) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
|
for i := range s.keys {
|
||||||
|
if s.keys[i].ID == keyID {
|
||||||
|
if s.keys[i].UserID != userID {
|
||||||
|
return fmt.Errorf("key not found or permission denied")
|
||||||
|
}
|
||||||
|
s.keys[i].IsActive = false
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("key not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateKey hashes the raw key and finds a matching, active, non-expired entry.
|
||||||
|
// Uses constant-time comparison to prevent timing side-channels.
|
||||||
|
// Pass an empty KeyType to accept any type.
|
||||||
|
func (s *ConfigKeyStore) ValidateKey(_ context.Context, rawKey string, keyType KeyType) (*UserKey, error) {
|
||||||
|
hash := hashSHA256Hex(rawKey)
|
||||||
|
hashBytes, _ := hex.DecodeString(hash)
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
// Write lock: ValidateKey updates LastUsedAt on the matched entry.
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
|
||||||
|
for i := range s.keys {
|
||||||
|
k := &s.keys[i]
|
||||||
|
if !k.IsActive {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if k.ExpiresAt != nil && k.ExpiresAt.Before(now) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if keyType != "" && k.KeyType != keyType {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
stored, _ := hex.DecodeString(k.KeyHash)
|
||||||
|
if subtle.ConstantTimeCompare(hashBytes, stored) != 1 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
k.LastUsedAt = &now
|
||||||
|
result := *k
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("invalid or expired key")
|
||||||
|
}
|
||||||
+256
@@ -0,0 +1,256 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
|
||||||
|
type DatabaseKeyStoreOptions struct {
|
||||||
|
// Cache is an optional cache instance. If nil, uses the default cache.
|
||||||
|
Cache *cache.Cache
|
||||||
|
// CacheTTL is the duration to cache ValidateKey results.
|
||||||
|
// Default: 2 minutes.
|
||||||
|
CacheTTL time.Duration
|
||||||
|
// SQLNames provides custom procedure names. If nil, uses DefaultKeyStoreSQLNames().
|
||||||
|
SQLNames *KeyStoreSQLNames
|
||||||
|
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
||||||
|
// If nil, reconnection is disabled.
|
||||||
|
DBFactory func() (*sql.DB, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DatabaseKeyStore is a KeyStore backed by PostgreSQL stored procedures.
|
||||||
|
// All DB operations go through configurable procedure names; the raw key is
|
||||||
|
// never passed to the database.
|
||||||
|
//
|
||||||
|
// See keystore_schema.sql for the required table and procedure definitions.
|
||||||
|
//
|
||||||
|
// Note: DeleteKey invalidates the cache entry for the deleted key. Due to the
|
||||||
|
// cache TTL, a deleted key may continue to authenticate for up to CacheTTL
|
||||||
|
// (default 2 minutes) if the cache entry cannot be invalidated.
|
||||||
|
type DatabaseKeyStore struct {
|
||||||
|
db *sql.DB
|
||||||
|
dbMu sync.RWMutex
|
||||||
|
dbFactory func() (*sql.DB, error)
|
||||||
|
sqlNames *KeyStoreSQLNames
|
||||||
|
cache *cache.Cache
|
||||||
|
cacheTTL time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDatabaseKeyStore creates a DatabaseKeyStore with optional configuration.
|
||||||
|
func NewDatabaseKeyStore(db *sql.DB, opts ...DatabaseKeyStoreOptions) *DatabaseKeyStore {
|
||||||
|
o := DatabaseKeyStoreOptions{}
|
||||||
|
if len(opts) > 0 {
|
||||||
|
o = opts[0]
|
||||||
|
}
|
||||||
|
if o.CacheTTL == 0 {
|
||||||
|
o.CacheTTL = 2 * time.Minute
|
||||||
|
}
|
||||||
|
c := o.Cache
|
||||||
|
if c == nil {
|
||||||
|
c = cache.GetDefaultCache()
|
||||||
|
}
|
||||||
|
names := MergeKeyStoreSQLNames(DefaultKeyStoreSQLNames(), o.SQLNames)
|
||||||
|
return &DatabaseKeyStore{
|
||||||
|
db: db,
|
||||||
|
dbFactory: o.DBFactory,
|
||||||
|
sqlNames: names,
|
||||||
|
cache: c,
|
||||||
|
cacheTTL: o.CacheTTL,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ks *DatabaseKeyStore) getDB() *sql.DB {
|
||||||
|
ks.dbMu.RLock()
|
||||||
|
defer ks.dbMu.RUnlock()
|
||||||
|
return ks.db
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ks *DatabaseKeyStore) reconnectDB() error {
|
||||||
|
if ks.dbFactory == nil {
|
||||||
|
return fmt.Errorf("no db factory configured for reconnect")
|
||||||
|
}
|
||||||
|
newDB, err := ks.dbFactory()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ks.dbMu.Lock()
|
||||||
|
ks.db = newDB
|
||||||
|
ks.dbMu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateKey generates a raw key, stores its SHA-256 hash via the create procedure,
|
||||||
|
// and returns the raw key once.
|
||||||
|
func (ks *DatabaseKeyStore) CreateKey(ctx context.Context, req CreateKeyRequest) (*CreateKeyResponse, error) {
|
||||||
|
rawBytes := make([]byte, 32)
|
||||||
|
if _, err := rand.Read(rawBytes); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate key material: %w", err)
|
||||||
|
}
|
||||||
|
rawKey := base64.RawURLEncoding.EncodeToString(rawBytes)
|
||||||
|
hash := hashSHA256Hex(rawKey)
|
||||||
|
|
||||||
|
type createRequest struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
KeyType KeyType `json:"key_type"`
|
||||||
|
KeyHash string `json:"key_hash"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Scopes []string `json:"scopes,omitempty"`
|
||||||
|
Meta map[string]any `json:"meta,omitempty"`
|
||||||
|
ExpiresAt *time.Time `json:"expires_at,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
reqJSON, err := json.Marshal(createRequest{
|
||||||
|
UserID: req.UserID,
|
||||||
|
KeyType: req.KeyType,
|
||||||
|
KeyHash: hash,
|
||||||
|
Name: req.Name,
|
||||||
|
Scopes: req.Scopes,
|
||||||
|
Meta: req.Meta,
|
||||||
|
ExpiresAt: req.ExpiresAt,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal create key request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var keyJSON sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1::jsonb)`, ks.sqlNames.CreateKey)
|
||||||
|
if err = ks.getDB().QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &keyJSON); err != nil {
|
||||||
|
return nil, fmt.Errorf("create key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, errors.New(nullStringOr(errorMsg, "create key failed"))
|
||||||
|
}
|
||||||
|
|
||||||
|
var key UserKey
|
||||||
|
if err = json.Unmarshal([]byte(keyJSON.String), &key); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse created key: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &CreateKeyResponse{Key: key, RawKey: rawKey}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserKeys returns all active, non-expired keys for the given user.
|
||||||
|
// Pass an empty KeyType to return all types.
|
||||||
|
func (ks *DatabaseKeyStore) GetUserKeys(ctx context.Context, userID int, keyType KeyType) ([]UserKey, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var keysJSON sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_keys::text FROM %s($1, $2)`, ks.sqlNames.GetUserKeys)
|
||||||
|
if err := ks.getDB().QueryRowContext(ctx, query, userID, string(keyType)).Scan(&success, &errorMsg, &keysJSON); err != nil {
|
||||||
|
return nil, fmt.Errorf("get user keys procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, errors.New(nullStringOr(errorMsg, "get user keys failed"))
|
||||||
|
}
|
||||||
|
|
||||||
|
var keys []UserKey
|
||||||
|
if keysJSON.Valid && keysJSON.String != "" && keysJSON.String != "[]" {
|
||||||
|
if err := json.Unmarshal([]byte(keysJSON.String), &keys); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user keys: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if keys == nil {
|
||||||
|
keys = []UserKey{}
|
||||||
|
}
|
||||||
|
return keys, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteKey soft-deletes a key after verifying ownership and invalidates its cache entry.
|
||||||
|
// The delete procedure returns the key_hash so no separate lookup is needed.
|
||||||
|
// Note: cache invalidation is best-effort; a cached entry may persist for up to CacheTTL.
|
||||||
|
func (ks *DatabaseKeyStore) DeleteKey(ctx context.Context, userID int, keyID int64) error {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var keyHash sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_key_hash FROM %s($1, $2)`, ks.sqlNames.DeleteKey)
|
||||||
|
if err := ks.getDB().QueryRowContext(ctx, query, userID, keyID).Scan(&success, &errorMsg, &keyHash); err != nil {
|
||||||
|
return fmt.Errorf("delete key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return errors.New(nullStringOr(errorMsg, "delete key failed"))
|
||||||
|
}
|
||||||
|
|
||||||
|
if keyHash.Valid && keyHash.String != "" && ks.cache != nil {
|
||||||
|
_ = ks.cache.Delete(ctx, keystoreCacheKey(keyHash.String))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateKey hashes the raw key and calls the validate procedure.
|
||||||
|
// Results are cached for CacheTTL to reduce DB load on hot paths.
|
||||||
|
func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyType KeyType) (*UserKey, error) {
|
||||||
|
hash := hashSHA256Hex(rawKey)
|
||||||
|
cacheKey := keystoreCacheKey(hash)
|
||||||
|
|
||||||
|
if ks.cache != nil {
|
||||||
|
var cached UserKey
|
||||||
|
if err := ks.cache.Get(ctx, cacheKey, &cached); err == nil {
|
||||||
|
if cached.IsActive {
|
||||||
|
return &cached, nil
|
||||||
|
}
|
||||||
|
return nil, errors.New("invalid or expired key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var keyJSON sql.NullString
|
||||||
|
|
||||||
|
runQuery := func() error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1, $2)`, ks.sqlNames.ValidateKey)
|
||||||
|
return ks.getDB().QueryRowContext(ctx, query, hash, string(keyType)).Scan(&success, &errorMsg, &keyJSON)
|
||||||
|
}
|
||||||
|
if err := runQuery(); err != nil {
|
||||||
|
if isDBClosed(err) {
|
||||||
|
if reconnErr := ks.reconnectDB(); reconnErr == nil {
|
||||||
|
err = runQuery()
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("validate key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
return nil, fmt.Errorf("validate key procedure failed: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
return nil, errors.New(nullStringOr(errorMsg, "invalid or expired key"))
|
||||||
|
}
|
||||||
|
|
||||||
|
var key UserKey
|
||||||
|
if err := json.Unmarshal([]byte(keyJSON.String), &key); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse validated key: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if ks.cache != nil {
|
||||||
|
_ = ks.cache.Set(ctx, cacheKey, key, ks.cacheTTL)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func keystoreCacheKey(hash string) string {
|
||||||
|
return "keystore:validate:" + hash
|
||||||
|
}
|
||||||
|
|
||||||
|
// nullStringOr returns s.String if valid, otherwise the fallback.
|
||||||
|
func nullStringOr(s sql.NullString, fallback string) string {
|
||||||
|
if s.Valid && s.String != "" {
|
||||||
|
return s.String
|
||||||
|
}
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
+187
@@ -0,0 +1,187 @@
|
|||||||
|
-- Keystore schema for per-user auth keys
|
||||||
|
-- Apply alongside database_schema.sql (requires the users table)
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS user_keys (
|
||||||
|
id BIGSERIAL PRIMARY KEY,
|
||||||
|
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
key_type VARCHAR(50) NOT NULL,
|
||||||
|
key_hash VARCHAR(64) NOT NULL UNIQUE, -- SHA-256 hex digest (64 chars)
|
||||||
|
name VARCHAR(255) NOT NULL DEFAULT '',
|
||||||
|
scopes TEXT, -- JSON array, e.g. '["read","write"]'
|
||||||
|
meta JSONB,
|
||||||
|
expires_at TIMESTAMP,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_used_at TIMESTAMP,
|
||||||
|
is_active BOOLEAN DEFAULT true
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_user_keys_key_hash ON user_keys(key_hash);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type);
|
||||||
|
|
||||||
|
-- resolvespec_keystore_get_user_keys
|
||||||
|
-- Returns all active, non-expired keys for a user.
|
||||||
|
-- Pass empty p_key_type to return all key types.
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_keystore_get_user_keys(
|
||||||
|
p_user_id INTEGER,
|
||||||
|
p_key_type TEXT DEFAULT ''
|
||||||
|
)
|
||||||
|
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_keys JSONB)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_keys JSONB;
|
||||||
|
BEGIN
|
||||||
|
SELECT COALESCE(
|
||||||
|
jsonb_agg(
|
||||||
|
jsonb_build_object(
|
||||||
|
'id', k.id,
|
||||||
|
'user_id', k.user_id,
|
||||||
|
'key_type', k.key_type,
|
||||||
|
'name', k.name,
|
||||||
|
'scopes', CASE WHEN k.scopes IS NOT NULL THEN k.scopes::jsonb ELSE '[]'::jsonb END,
|
||||||
|
'meta', COALESCE(k.meta, '{}'::jsonb),
|
||||||
|
'expires_at', k.expires_at,
|
||||||
|
'created_at', k.created_at,
|
||||||
|
'last_used_at', k.last_used_at,
|
||||||
|
'is_active', k.is_active
|
||||||
|
)
|
||||||
|
),
|
||||||
|
'[]'::jsonb
|
||||||
|
)
|
||||||
|
INTO v_keys
|
||||||
|
FROM user_keys k
|
||||||
|
WHERE k.user_id = p_user_id
|
||||||
|
AND k.is_active = true
|
||||||
|
AND (k.expires_at IS NULL OR k.expires_at > NOW())
|
||||||
|
AND (p_key_type = '' OR k.key_type = p_key_type);
|
||||||
|
|
||||||
|
RETURN QUERY SELECT true, NULL::TEXT, v_keys;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
-- resolvespec_keystore_create_key
|
||||||
|
-- Inserts a new key row. key_hash is provided by the caller (Go hashes the raw key).
|
||||||
|
-- Returns the created key record (without key_hash).
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_keystore_create_key(
|
||||||
|
p_request JSONB
|
||||||
|
)
|
||||||
|
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_key JSONB)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_id BIGINT;
|
||||||
|
v_created_at TIMESTAMP;
|
||||||
|
v_key JSONB;
|
||||||
|
BEGIN
|
||||||
|
INSERT INTO user_keys (user_id, key_type, key_hash, name, scopes, meta, expires_at)
|
||||||
|
VALUES (
|
||||||
|
(p_request->>'user_id')::INTEGER,
|
||||||
|
p_request->>'key_type',
|
||||||
|
p_request->>'key_hash',
|
||||||
|
COALESCE(p_request->>'name', ''),
|
||||||
|
p_request->>'scopes',
|
||||||
|
p_request->'meta',
|
||||||
|
CASE WHEN p_request->>'expires_at' IS NOT NULL
|
||||||
|
THEN (p_request->>'expires_at')::TIMESTAMP
|
||||||
|
ELSE NULL
|
||||||
|
END
|
||||||
|
)
|
||||||
|
RETURNING id, created_at INTO v_id, v_created_at;
|
||||||
|
|
||||||
|
v_key := jsonb_build_object(
|
||||||
|
'id', v_id,
|
||||||
|
'user_id', (p_request->>'user_id')::INTEGER,
|
||||||
|
'key_type', p_request->>'key_type',
|
||||||
|
'name', COALESCE(p_request->>'name', ''),
|
||||||
|
'scopes', CASE WHEN p_request->>'scopes' IS NOT NULL
|
||||||
|
THEN (p_request->>'scopes')::jsonb
|
||||||
|
ELSE '[]'::jsonb END,
|
||||||
|
'meta', COALESCE(p_request->'meta', '{}'::jsonb),
|
||||||
|
'expires_at', p_request->>'expires_at',
|
||||||
|
'created_at', v_created_at,
|
||||||
|
'is_active', true
|
||||||
|
);
|
||||||
|
|
||||||
|
RETURN QUERY SELECT true, NULL::TEXT, v_key;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
-- resolvespec_keystore_delete_key
|
||||||
|
-- Soft-deletes a key (is_active = false) after verifying ownership.
|
||||||
|
-- Returns p_key_hash so the caller can invalidate cache entries without a separate query.
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_keystore_delete_key(
|
||||||
|
p_user_id INTEGER,
|
||||||
|
p_key_id BIGINT
|
||||||
|
)
|
||||||
|
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_key_hash TEXT)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_hash TEXT;
|
||||||
|
BEGIN
|
||||||
|
UPDATE user_keys
|
||||||
|
SET is_active = false
|
||||||
|
WHERE id = p_key_id AND user_id = p_user_id AND is_active = true
|
||||||
|
RETURNING key_hash INTO v_hash;
|
||||||
|
|
||||||
|
IF NOT FOUND THEN
|
||||||
|
RETURN QUERY SELECT false, 'key not found or already deleted'::TEXT, NULL::TEXT;
|
||||||
|
RETURN;
|
||||||
|
END IF;
|
||||||
|
|
||||||
|
RETURN QUERY SELECT true, NULL::TEXT, v_hash;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, NULL::TEXT;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
|
|
||||||
|
-- resolvespec_keystore_validate_key
|
||||||
|
-- Looks up a key by its SHA-256 hash, checks active status and expiry,
|
||||||
|
-- updates last_used_at, and returns the key record.
|
||||||
|
-- p_key_type can be empty to accept any key type.
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_keystore_validate_key(
|
||||||
|
p_key_hash TEXT,
|
||||||
|
p_key_type TEXT DEFAULT ''
|
||||||
|
)
|
||||||
|
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_key JSONB)
|
||||||
|
LANGUAGE plpgsql AS $$
|
||||||
|
DECLARE
|
||||||
|
v_key_rec user_keys%ROWTYPE;
|
||||||
|
v_key JSONB;
|
||||||
|
BEGIN
|
||||||
|
SELECT * INTO v_key_rec
|
||||||
|
FROM user_keys
|
||||||
|
WHERE key_hash = p_key_hash
|
||||||
|
AND is_active = true
|
||||||
|
AND (expires_at IS NULL OR expires_at > NOW())
|
||||||
|
AND (p_key_type = '' OR key_type = p_key_type);
|
||||||
|
|
||||||
|
IF NOT FOUND THEN
|
||||||
|
RETURN QUERY SELECT false, 'invalid or expired key'::TEXT, NULL::JSONB;
|
||||||
|
RETURN;
|
||||||
|
END IF;
|
||||||
|
|
||||||
|
UPDATE user_keys SET last_used_at = NOW() WHERE id = v_key_rec.id;
|
||||||
|
|
||||||
|
v_key := jsonb_build_object(
|
||||||
|
'id', v_key_rec.id,
|
||||||
|
'user_id', v_key_rec.user_id,
|
||||||
|
'key_type', v_key_rec.key_type,
|
||||||
|
'name', v_key_rec.name,
|
||||||
|
'scopes', CASE WHEN v_key_rec.scopes IS NOT NULL
|
||||||
|
THEN v_key_rec.scopes::jsonb
|
||||||
|
ELSE '[]'::jsonb END,
|
||||||
|
'meta', COALESCE(v_key_rec.meta, '{}'::jsonb),
|
||||||
|
'expires_at', v_key_rec.expires_at,
|
||||||
|
'created_at', v_key_rec.created_at,
|
||||||
|
'last_used_at', NOW(),
|
||||||
|
'is_active', v_key_rec.is_active
|
||||||
|
);
|
||||||
|
|
||||||
|
RETURN QUERY SELECT true, NULL::TEXT, v_key;
|
||||||
|
EXCEPTION WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||||
|
END;
|
||||||
|
$$;
|
||||||
+61
@@ -0,0 +1,61 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import "fmt"
|
||||||
|
|
||||||
|
// KeyStoreSQLNames holds the configurable stored procedure names used by DatabaseKeyStore.
|
||||||
|
// Use DefaultKeyStoreSQLNames() for defaults and MergeKeyStoreSQLNames() for partial overrides.
|
||||||
|
type KeyStoreSQLNames struct {
|
||||||
|
GetUserKeys string // default: "resolvespec_keystore_get_user_keys"
|
||||||
|
CreateKey string // default: "resolvespec_keystore_create_key"
|
||||||
|
DeleteKey string // default: "resolvespec_keystore_delete_key"
|
||||||
|
ValidateKey string // default: "resolvespec_keystore_validate_key"
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultKeyStoreSQLNames returns a KeyStoreSQLNames with all default resolvespec_keystore_* values.
|
||||||
|
func DefaultKeyStoreSQLNames() *KeyStoreSQLNames {
|
||||||
|
return &KeyStoreSQLNames{
|
||||||
|
GetUserKeys: "resolvespec_keystore_get_user_keys",
|
||||||
|
CreateKey: "resolvespec_keystore_create_key",
|
||||||
|
DeleteKey: "resolvespec_keystore_delete_key",
|
||||||
|
ValidateKey: "resolvespec_keystore_validate_key",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MergeKeyStoreSQLNames returns a copy of base with any non-empty fields from override applied.
|
||||||
|
// If override is nil, a copy of base is returned.
|
||||||
|
func MergeKeyStoreSQLNames(base, override *KeyStoreSQLNames) *KeyStoreSQLNames {
|
||||||
|
if override == nil {
|
||||||
|
copied := *base
|
||||||
|
return &copied
|
||||||
|
}
|
||||||
|
merged := *base
|
||||||
|
if override.GetUserKeys != "" {
|
||||||
|
merged.GetUserKeys = override.GetUserKeys
|
||||||
|
}
|
||||||
|
if override.CreateKey != "" {
|
||||||
|
merged.CreateKey = override.CreateKey
|
||||||
|
}
|
||||||
|
if override.DeleteKey != "" {
|
||||||
|
merged.DeleteKey = override.DeleteKey
|
||||||
|
}
|
||||||
|
if override.ValidateKey != "" {
|
||||||
|
merged.ValidateKey = override.ValidateKey
|
||||||
|
}
|
||||||
|
return &merged
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateKeyStoreSQLNames checks that all non-empty procedure names are valid SQL identifiers.
|
||||||
|
func ValidateKeyStoreSQLNames(names *KeyStoreSQLNames) error {
|
||||||
|
fields := map[string]string{
|
||||||
|
"GetUserKeys": names.GetUserKeys,
|
||||||
|
"CreateKey": names.CreateKey,
|
||||||
|
"DeleteKey": names.DeleteKey,
|
||||||
|
"ValidateKey": names.ValidateKey,
|
||||||
|
}
|
||||||
|
for field, val := range fields {
|
||||||
|
if val != "" && !validSQLIdentifier.MatchString(val) {
|
||||||
|
return fmt.Errorf("KeyStoreSQLNames.%s contains invalid characters: %q", field, val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+737
@@ -0,0 +1,737 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
)
|
||||||
|
|
||||||
|
// contextKey is a custom type for context keys to avoid collisions
|
||||||
|
type contextKey string
|
||||||
|
|
||||||
|
const (
|
||||||
|
// Context keys for user information
|
||||||
|
UserIDKey contextKey = "user_id"
|
||||||
|
UserNameKey contextKey = "user_name"
|
||||||
|
UserLevelKey contextKey = "user_level"
|
||||||
|
SessionIDKey contextKey = "session_id"
|
||||||
|
SessionRIDKey contextKey = "session_rid"
|
||||||
|
RemoteIDKey contextKey = "remote_id"
|
||||||
|
UserRolesKey contextKey = "user_roles"
|
||||||
|
UserEmailKey contextKey = "user_email"
|
||||||
|
UserContextKey contextKey = "user_context"
|
||||||
|
UserMetaKey contextKey = "user_meta"
|
||||||
|
SkipAuthKey contextKey = "skip_auth"
|
||||||
|
OptionalAuthKey contextKey = "optional_auth"
|
||||||
|
ModelRulesKey contextKey = "model_rules"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SkipAuth returns a context with skip auth flag set to true
|
||||||
|
// Use this to mark routes that should bypass authentication middleware
|
||||||
|
func SkipAuth(ctx context.Context) context.Context {
|
||||||
|
return context.WithValue(ctx, SkipAuthKey, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// OptionalAuth returns a context with optional auth flag set to true
|
||||||
|
// Use this to mark routes that should try to authenticate, but fall back to guest if authentication fails
|
||||||
|
func OptionalAuth(ctx context.Context) context.Context {
|
||||||
|
return context.WithValue(ctx, OptionalAuthKey, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// createGuestContext creates a guest user context for unauthenticated requests
|
||||||
|
func createGuestContext(r *http.Request) *UserContext {
|
||||||
|
return &UserContext{
|
||||||
|
UserID: 0,
|
||||||
|
UserName: "guest",
|
||||||
|
UserLevel: 0,
|
||||||
|
SessionID: "",
|
||||||
|
RemoteID: r.RemoteAddr,
|
||||||
|
Roles: []string{"guest"},
|
||||||
|
Email: "",
|
||||||
|
Claims: map[string]any{},
|
||||||
|
Meta: map[string]any{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// setUserContext adds a user context to the request context
|
||||||
|
func setUserContext(r *http.Request, userCtx *UserContext) *http.Request {
|
||||||
|
ctx := r.Context()
|
||||||
|
ctx = context.WithValue(ctx, UserContextKey, userCtx)
|
||||||
|
ctx = context.WithValue(ctx, UserIDKey, userCtx.UserID)
|
||||||
|
ctx = context.WithValue(ctx, UserNameKey, userCtx.UserName)
|
||||||
|
ctx = context.WithValue(ctx, UserLevelKey, userCtx.UserLevel)
|
||||||
|
ctx = context.WithValue(ctx, SessionIDKey, userCtx.SessionID)
|
||||||
|
ctx = context.WithValue(ctx, SessionRIDKey, userCtx.SessionRID)
|
||||||
|
ctx = context.WithValue(ctx, RemoteIDKey, userCtx.RemoteID)
|
||||||
|
ctx = context.WithValue(ctx, UserRolesKey, userCtx.Roles)
|
||||||
|
|
||||||
|
if userCtx.Email != "" {
|
||||||
|
ctx = context.WithValue(ctx, UserEmailKey, userCtx.Email)
|
||||||
|
}
|
||||||
|
if len(userCtx.Meta) > 0 {
|
||||||
|
ctx = context.WithValue(ctx, UserMetaKey, userCtx.Meta)
|
||||||
|
}
|
||||||
|
|
||||||
|
return r.WithContext(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// authenticateRequest performs authentication and adds user context to the request
|
||||||
|
// This is the shared authentication logic used by both handler and middleware
|
||||||
|
func authenticateRequest(w http.ResponseWriter, r *http.Request, provider SecurityProvider) (*http.Request, bool) {
|
||||||
|
// Call the provider's Authenticate method
|
||||||
|
userCtx, err := provider.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Authentication failed: "+err.Error(), http.StatusUnauthorized)
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return setUserContext(r, userCtx), true
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAuthHandler creates an authentication handler that can be used standalone
|
||||||
|
// This handler performs authentication and returns 401 if authentication fails
|
||||||
|
// Use this when you need authentication logic without middleware wrapping
|
||||||
|
func NewAuthHandler(securityList *SecurityList, next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Get the security provider
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticate the request
|
||||||
|
authenticatedReq, ok := authenticateRequest(w, r, provider)
|
||||||
|
if !ok {
|
||||||
|
return // authenticateRequest already wrote the error response
|
||||||
|
}
|
||||||
|
|
||||||
|
// Continue with authenticated context
|
||||||
|
next.ServeHTTP(w, authenticatedReq)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOptionalAuthHandler creates an optional authentication handler that can be used standalone
|
||||||
|
// This handler tries to authenticate but falls back to guest context if authentication fails
|
||||||
|
// Use this for routes that should show personalized content for authenticated users but still work for guests
|
||||||
|
func NewOptionalAuthHandler(securityList *SecurityList, next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Get the security provider
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to authenticate
|
||||||
|
userCtx, err := provider.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
// Authentication failed - set guest context and continue
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authentication succeeded - set user context
|
||||||
|
next.ServeHTTP(w, setUserContext(r, userCtx))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOptionalAuthMiddleware creates authentication middleware that always continues.
|
||||||
|
// On auth failure, a guest user context is set instead of returning 401.
|
||||||
|
// Intended for spec routes where auth enforcement is deferred to a BeforeHandle hook
|
||||||
|
// after model resolution.
|
||||||
|
func NewOptionalAuthMiddleware(securityList *SecurityList) func(http.Handler) http.Handler {
|
||||||
|
return func(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
userCtx, err := provider.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
next.ServeHTTP(w, setUserContext(r, userCtx))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAuthMiddleware creates an authentication middleware with the given security list
|
||||||
|
// This middleware extracts user authentication from the request and adds it to context
|
||||||
|
// Routes can skip authentication by setting SkipAuthKey context value (use SkipAuth helper)
|
||||||
|
// Routes can use optional authentication by setting OptionalAuthKey context value (use OptionalAuth helper)
|
||||||
|
// When authentication is skipped or fails with optional auth, a guest user context is set instead
|
||||||
|
func NewAuthMiddleware(securityList *SecurityList) func(http.Handler) http.Handler {
|
||||||
|
return func(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Check if this route should skip authentication
|
||||||
|
if skip, ok := r.Context().Value(SkipAuthKey).(bool); ok && skip {
|
||||||
|
// Set guest user context for skipped routes
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get the security provider
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this route has optional authentication
|
||||||
|
optional, _ := r.Context().Value(OptionalAuthKey).(bool)
|
||||||
|
|
||||||
|
// Try to authenticate
|
||||||
|
userCtx, err := provider.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
if optional {
|
||||||
|
// Optional auth failed - set guest context and continue
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// Required auth failed - return error
|
||||||
|
http.Error(w, "Authentication failed: "+err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authentication succeeded - set user context
|
||||||
|
next.ServeHTTP(w, setUserContext(r, userCtx))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewModelAuthMiddleware creates authentication middleware that respects ModelRules for the given model name.
|
||||||
|
// It first checks if ModelRules are set for the model:
|
||||||
|
// - If SecurityDisabled is true, authentication is skipped and a guest context is set.
|
||||||
|
// - Otherwise, all checks from NewAuthMiddleware apply (SkipAuthKey, provider check, OptionalAuthKey, Authenticate).
|
||||||
|
//
|
||||||
|
// If the model is not found in any registry, the middleware falls back to standard NewAuthMiddleware behaviour.
|
||||||
|
func NewModelAuthMiddleware(securityList *SecurityList, modelName string) func(http.Handler) http.Handler {
|
||||||
|
return func(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Check ModelRules first
|
||||||
|
if rules, err := modelregistry.GetModelRulesByName(modelName); err == nil {
|
||||||
|
// Store rules in context for downstream use (e.g., security hooks)
|
||||||
|
r = r.WithContext(context.WithValue(r.Context(), ModelRulesKey, rules))
|
||||||
|
|
||||||
|
if rules.SecurityDisabled {
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
isRead := r.Method == http.MethodGet || r.Method == http.MethodHead
|
||||||
|
isUpdate := r.Method == http.MethodPut || r.Method == http.MethodPatch
|
||||||
|
if (isRead && rules.CanPublicRead) || (isUpdate && rules.CanPublicUpdate) {
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this route should skip authentication
|
||||||
|
if skip, ok := r.Context().Value(SkipAuthKey).(bool); ok && skip {
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get the security provider
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this route has optional authentication
|
||||||
|
optional, _ := r.Context().Value(OptionalAuthKey).(bool)
|
||||||
|
|
||||||
|
// Try to authenticate
|
||||||
|
userCtx, err := provider.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
if optional {
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
next.ServeHTTP(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
http.Error(w, "Authentication failed: "+err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
next.ServeHTTP(w, setUserContext(r, userCtx))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetSecurityMiddleware adds security context to requests
|
||||||
|
// This middleware should be applied after AuthMiddleware
|
||||||
|
func SetSecurityMiddleware(securityList *SecurityList) func(http.Handler) http.Handler {
|
||||||
|
return func(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx := context.WithValue(r.Context(), SECURITY_CONTEXT_KEY, securityList)
|
||||||
|
next.ServeHTTP(w, r.WithContext(ctx))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithAuth wraps an HTTPFuncType handler with required authentication
|
||||||
|
// This function performs authentication and returns 401 if authentication fails
|
||||||
|
// Use this for handlers that require authenticated users
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// handler := funcspec.NewHandler(db)
|
||||||
|
// wrappedHandler := security.WithAuth(handler.SqlQueryList("SELECT * FROM orders WHERE user_id = [rid_user]", false, false, false), securityList)
|
||||||
|
// router.HandleFunc("/api/orders", wrappedHandler)
|
||||||
|
func WithAuth(handler func(http.ResponseWriter, *http.Request), securityList *SecurityList) func(http.ResponseWriter, *http.Request) {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Get the security provider
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticate the request
|
||||||
|
authenticatedReq, ok := authenticateRequest(w, r, provider)
|
||||||
|
if !ok {
|
||||||
|
return // authenticateRequest already wrote the error response
|
||||||
|
}
|
||||||
|
|
||||||
|
// Continue with authenticated context
|
||||||
|
handler(w, authenticatedReq)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithOptionalAuth wraps an HTTPFuncType handler with optional authentication
|
||||||
|
// This function tries to authenticate but falls back to guest context if authentication fails
|
||||||
|
// Use this for handlers that should show personalized content for authenticated users but still work for guests
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// handler := funcspec.NewHandler(db)
|
||||||
|
// wrappedHandler := security.WithOptionalAuth(handler.SqlQueryList("SELECT * FROM products", false, false, false), securityList)
|
||||||
|
// router.HandleFunc("/api/products", wrappedHandler)
|
||||||
|
func WithOptionalAuth(handler func(http.ResponseWriter, *http.Request), securityList *SecurityList) func(http.ResponseWriter, *http.Request) {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Get the security provider
|
||||||
|
provider := securityList.Provider()
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to authenticate
|
||||||
|
userCtx, err := provider.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
// Authentication failed - set guest context and continue
|
||||||
|
guestCtx := createGuestContext(r)
|
||||||
|
handler(w, setUserContext(r, guestCtx))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authentication succeeded - set user context
|
||||||
|
handler(w, setUserContext(r, userCtx))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithSecurityContext wraps an HTTPFuncType handler with security context
|
||||||
|
// This function allows you to add security context to specific handler functions
|
||||||
|
// without needing to apply middleware globally
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// handler := funcspec.NewHandler(db)
|
||||||
|
// wrappedHandler := security.WithSecurityContext(handler.SqlQueryList("SELECT * FROM users", false, false, false), securityList)
|
||||||
|
// router.HandleFunc("/api/users", wrappedHandler)
|
||||||
|
func WithSecurityContext(handler func(http.ResponseWriter, *http.Request), securityList *SecurityList) func(http.ResponseWriter, *http.Request) {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx := context.WithValue(r.Context(), SECURITY_CONTEXT_KEY, securityList)
|
||||||
|
handler(w, r.WithContext(ctx))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithAuthAndSecurity wraps an HTTPFuncType handler with both authentication and security context
|
||||||
|
// This is a convenience function that combines WithAuth and WithSecurityContext
|
||||||
|
// Use this when you need both authentication and security context for a handler
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// handler := funcspec.NewHandler(db)
|
||||||
|
// wrappedHandler := security.WithAuthAndSecurity(handler.SqlQueryList("SELECT * FROM users", false, false, false), securityList)
|
||||||
|
// router.HandleFunc("/api/users", wrappedHandler)
|
||||||
|
func WithAuthAndSecurity(handler func(http.ResponseWriter, *http.Request), securityList *SecurityList) func(http.ResponseWriter, *http.Request) {
|
||||||
|
return WithAuth(WithSecurityContext(handler, securityList), securityList)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithOptionalAuthAndSecurity wraps an HTTPFuncType handler with optional authentication and security context
|
||||||
|
// This is a convenience function that combines WithOptionalAuth and WithSecurityContext
|
||||||
|
// Use this when you want optional authentication and security context for a handler
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// handler := funcspec.NewHandler(db)
|
||||||
|
// wrappedHandler := security.WithOptionalAuthAndSecurity(handler.SqlQueryList("SELECT * FROM products", false, false, false), securityList)
|
||||||
|
// router.HandleFunc("/api/products", wrappedHandler)
|
||||||
|
func WithOptionalAuthAndSecurity(handler func(http.ResponseWriter, *http.Request), securityList *SecurityList) func(http.ResponseWriter, *http.Request) {
|
||||||
|
return WithOptionalAuth(WithSecurityContext(handler, securityList), securityList)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSecurityList extracts the SecurityList from request context
|
||||||
|
func GetSecurityList(ctx context.Context) (*SecurityList, bool) {
|
||||||
|
securityList, ok := ctx.Value(SECURITY_CONTEXT_KEY).(*SecurityList)
|
||||||
|
return securityList, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserContext extracts the full user context from request context
|
||||||
|
func GetUserContext(ctx context.Context) (*UserContext, bool) {
|
||||||
|
userCtx, ok := ctx.Value(UserContextKey).(*UserContext)
|
||||||
|
return userCtx, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserID extracts the user ID from context
|
||||||
|
func GetUserID(ctx context.Context) (int, bool) {
|
||||||
|
userID, ok := ctx.Value(UserIDKey).(int)
|
||||||
|
return userID, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserName extracts the user name from context
|
||||||
|
func GetUserName(ctx context.Context) (string, bool) {
|
||||||
|
userName, ok := ctx.Value(UserNameKey).(string)
|
||||||
|
return userName, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserLevel extracts the user level from context
|
||||||
|
func GetUserLevel(ctx context.Context) (int, bool) {
|
||||||
|
userLevel, ok := ctx.Value(UserLevelKey).(int)
|
||||||
|
return userLevel, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSessionID extracts the session ID from context
|
||||||
|
func GetSessionID(ctx context.Context) (string, bool) {
|
||||||
|
sessionID, ok := ctx.Value(SessionIDKey).(string)
|
||||||
|
return sessionID, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSessionID extracts the session ID from context
|
||||||
|
func GetSessionRID(ctx context.Context) (int64, bool) {
|
||||||
|
sessionRIDStr, ok := ctx.Value(SessionRIDKey).(string)
|
||||||
|
sessionRID, err := strconv.ParseInt(sessionRIDStr, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
return sessionRID, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRemoteID extracts the remote ID from context
|
||||||
|
func GetRemoteID(ctx context.Context) (string, bool) {
|
||||||
|
remoteID, ok := ctx.Value(RemoteIDKey).(string)
|
||||||
|
return remoteID, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserRoles extracts user roles from context
|
||||||
|
func GetUserRoles(ctx context.Context) ([]string, bool) {
|
||||||
|
roles, ok := ctx.Value(UserRolesKey).([]string)
|
||||||
|
return roles, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserEmail extracts user email from context
|
||||||
|
func GetUserEmail(ctx context.Context) (string, bool) {
|
||||||
|
email, ok := ctx.Value(UserEmailKey).(string)
|
||||||
|
return email, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserMeta extracts user metadata from context
|
||||||
|
func GetUserMeta(ctx context.Context) (map[string]any, bool) {
|
||||||
|
meta, ok := ctx.Value(UserMetaKey).(map[string]any)
|
||||||
|
return meta, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// SessionCookieOptions configures the session cookie set by SetSessionCookie.
|
||||||
|
// All fields are optional; sensible secure defaults are applied when omitted.
|
||||||
|
type SessionCookieOptions struct {
|
||||||
|
// Name is the cookie name. Defaults to "session_token".
|
||||||
|
Name string
|
||||||
|
// Path is the cookie path. Defaults to "/".
|
||||||
|
Path string
|
||||||
|
// Domain restricts the cookie to a specific domain. Empty means current host.
|
||||||
|
Domain string
|
||||||
|
// Secure sets the Secure flag. Defaults to true.
|
||||||
|
// Set to false only in local development over HTTP.
|
||||||
|
Secure *bool
|
||||||
|
// SameSite sets the SameSite policy. Defaults to http.SameSiteLaxMode.
|
||||||
|
SameSite http.SameSite
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o SessionCookieOptions) name() string {
|
||||||
|
if o.Name != "" {
|
||||||
|
return o.Name
|
||||||
|
}
|
||||||
|
return "session_token"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o SessionCookieOptions) path() string {
|
||||||
|
if o.Path != "" {
|
||||||
|
return o.Path
|
||||||
|
}
|
||||||
|
return "/"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o SessionCookieOptions) secure() bool {
|
||||||
|
if o.Secure != nil {
|
||||||
|
return *o.Secure
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o SessionCookieOptions) sameSite() http.SameSite {
|
||||||
|
if o.SameSite != 0 {
|
||||||
|
return o.SameSite
|
||||||
|
}
|
||||||
|
return http.SameSiteLaxMode
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetSessionCookie writes the session_token cookie to the response after a successful login.
|
||||||
|
// Call this immediately after a successful Authenticator.Login() call.
|
||||||
|
//
|
||||||
|
// Example:
|
||||||
|
//
|
||||||
|
// resp, err := auth.Login(r.Context(), req)
|
||||||
|
// if err != nil { ... }
|
||||||
|
// security.SetSessionCookie(w, resp)
|
||||||
|
// json.NewEncoder(w).Encode(resp)
|
||||||
|
func SetSessionCookie(w http.ResponseWriter, loginResp *LoginResponse, opts ...SessionCookieOptions) {
|
||||||
|
var o SessionCookieOptions
|
||||||
|
if len(opts) > 0 {
|
||||||
|
o = opts[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
maxAge := 0
|
||||||
|
if loginResp.ExpiresIn > 0 {
|
||||||
|
maxAge = int(loginResp.ExpiresIn)
|
||||||
|
}
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: o.name(),
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: o.path(),
|
||||||
|
Domain: o.Domain,
|
||||||
|
MaxAge: maxAge,
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: o.secure(),
|
||||||
|
SameSite: o.sameSite(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSessionCookie returns the session token value from the request cookie, or empty string if not present.
|
||||||
|
//
|
||||||
|
// Example:
|
||||||
|
//
|
||||||
|
// token := security.GetSessionCookie(r)
|
||||||
|
func GetSessionCookie(r *http.Request, opts ...SessionCookieOptions) string {
|
||||||
|
var o SessionCookieOptions
|
||||||
|
if len(opts) > 0 {
|
||||||
|
o = opts[0]
|
||||||
|
}
|
||||||
|
cookie, err := r.Cookie(o.name())
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return cookie.Value
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearSessionCookie expires the session_token cookie, effectively logging the user out on the browser side.
|
||||||
|
// Call this after a successful Authenticator.Logout() call.
|
||||||
|
//
|
||||||
|
// Example:
|
||||||
|
//
|
||||||
|
// err := auth.Logout(r.Context(), req)
|
||||||
|
// if err != nil { ... }
|
||||||
|
// security.ClearSessionCookie(w)
|
||||||
|
func ClearSessionCookie(w http.ResponseWriter, opts ...SessionCookieOptions) {
|
||||||
|
var o SessionCookieOptions
|
||||||
|
if len(opts) > 0 {
|
||||||
|
o = opts[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: o.name(),
|
||||||
|
Value: "",
|
||||||
|
Path: o.path(),
|
||||||
|
Domain: o.Domain,
|
||||||
|
MaxAge: -1,
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: o.secure(),
|
||||||
|
SameSite: o.sameSite(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetModelRulesFromContext extracts ModelRules stored by NewModelAuthMiddleware
|
||||||
|
func GetModelRulesFromContext(ctx context.Context) (modelregistry.ModelRules, bool) {
|
||||||
|
rules, ok := ctx.Value(ModelRulesKey).(modelregistry.ModelRules)
|
||||||
|
return rules, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// // Handler adapters for resolvespec/restheadspec compatibility
|
||||||
|
// // These functions allow using NewAuthHandler and NewOptionalAuthHandler with custom handler abstractions
|
||||||
|
|
||||||
|
// // SpecHandlerAdapter is an interface for handler adapters that need authentication
|
||||||
|
// // Implement this interface to create adapters for custom handler types
|
||||||
|
// type SpecHandlerAdapter interface {
|
||||||
|
// // AdaptToHTTPHandler converts the custom handler to a standard http.Handler
|
||||||
|
// AdaptToHTTPHandler() http.Handler
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // ResolveSpecHandlerAdapter adapts a resolvespec/restheadspec handler method to http.Handler
|
||||||
|
// type ResolveSpecHandlerAdapter struct {
|
||||||
|
// // HandlerMethod is the method to call (e.g., handler.Handle, handler.HandleGet)
|
||||||
|
// HandlerMethod func(w any, r any, params map[string]string)
|
||||||
|
// // Params are the route parameters (e.g., {"schema": "public", "entity": "users"})
|
||||||
|
// Params map[string]string
|
||||||
|
// // RequestAdapter converts *http.Request to the custom Request interface
|
||||||
|
// // Use router.NewHTTPRequest from pkg/common/adapters/router
|
||||||
|
// RequestAdapter func(*http.Request) any
|
||||||
|
// // ResponseAdapter converts http.ResponseWriter to the custom ResponseWriter interface
|
||||||
|
// // Use router.NewHTTPResponseWriter from pkg/common/adapters/router
|
||||||
|
// ResponseAdapter func(http.ResponseWriter) any
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // AdaptToHTTPHandler implements SpecHandlerAdapter
|
||||||
|
// func (a *ResolveSpecHandlerAdapter) AdaptToHTTPHandler() http.Handler {
|
||||||
|
// return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// req := a.RequestAdapter(r)
|
||||||
|
// resp := a.ResponseAdapter(w)
|
||||||
|
// a.HandlerMethod(resp, req, a.Params)
|
||||||
|
// })
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // WrapSpecHandler wraps a spec handler adapter with authentication
|
||||||
|
// // Use this to apply NewAuthHandler or NewOptionalAuthHandler to resolvespec/restheadspec handlers
|
||||||
|
// //
|
||||||
|
// // Example with required auth:
|
||||||
|
// //
|
||||||
|
// // adapter := &security.ResolveSpecHandlerAdapter{
|
||||||
|
// // HandlerMethod: handler.Handle,
|
||||||
|
// // Params: map[string]string{"schema": "public", "entity": "users"},
|
||||||
|
// // RequestAdapter: func(r *http.Request) any { return router.NewHTTPRequest(r) },
|
||||||
|
// // ResponseAdapter: func(w http.ResponseWriter) any { return router.NewHTTPResponseWriter(w) },
|
||||||
|
// // }
|
||||||
|
// // authHandler := security.WrapSpecHandler(securityList, adapter, false)
|
||||||
|
// // muxRouter.Handle("/api/users", authHandler)
|
||||||
|
// func WrapSpecHandler(securityList *SecurityList, adapter SpecHandlerAdapter, optional bool) http.Handler {
|
||||||
|
// httpHandler := adapter.AdaptToHTTPHandler()
|
||||||
|
// if optional {
|
||||||
|
// return NewOptionalAuthHandler(securityList, httpHandler)
|
||||||
|
// }
|
||||||
|
// return NewAuthHandler(securityList, httpHandler)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // MuxRouteBuilder helps build authenticated routes with Gorilla Mux
|
||||||
|
// type MuxRouteBuilder struct {
|
||||||
|
// securityList *SecurityList
|
||||||
|
// requestAdapter func(*http.Request) any
|
||||||
|
// responseAdapter func(http.ResponseWriter) any
|
||||||
|
// paramExtractor func(*http.Request) map[string]string
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // NewMuxRouteBuilder creates a route builder for Gorilla Mux with standard router adapters
|
||||||
|
// // Usage:
|
||||||
|
// //
|
||||||
|
// // builder := security.NewMuxRouteBuilder(securityList, router.NewHTTPRequest, router.NewHTTPResponseWriter)
|
||||||
|
// func NewMuxRouteBuilder(
|
||||||
|
// securityList *SecurityList,
|
||||||
|
// requestAdapter func(*http.Request) any,
|
||||||
|
// responseAdapter func(http.ResponseWriter) any,
|
||||||
|
// ) *MuxRouteBuilder {
|
||||||
|
// return &MuxRouteBuilder{
|
||||||
|
// securityList: securityList,
|
||||||
|
// requestAdapter: requestAdapter,
|
||||||
|
// responseAdapter: responseAdapter,
|
||||||
|
// paramExtractor: nil, // Will be set per route using mux.Vars
|
||||||
|
// }
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // HandleAuth creates an authenticated route handler
|
||||||
|
// // pattern: the route pattern (e.g., "/{schema}/{entity}")
|
||||||
|
// // handler: the handler method to call (e.g., handler.Handle)
|
||||||
|
// // optional: true for optional auth (guest fallback), false for required auth (401 on failure)
|
||||||
|
// // methods: HTTP methods (e.g., "GET", "POST")
|
||||||
|
// //
|
||||||
|
// // Usage:
|
||||||
|
// //
|
||||||
|
// // builder.HandleAuth(router, "/{schema}/{entity}", handler.Handle, false, "POST")
|
||||||
|
// func (b *MuxRouteBuilder) HandleAuth(
|
||||||
|
// router interface {
|
||||||
|
// HandleFunc(pattern string, f func(http.ResponseWriter, *http.Request)) interface{ Methods(...string) interface{} }
|
||||||
|
// },
|
||||||
|
// pattern string,
|
||||||
|
// handlerMethod func(w any, r any, params map[string]string),
|
||||||
|
// optional bool,
|
||||||
|
// methods ...string,
|
||||||
|
// ) {
|
||||||
|
// router.HandleFunc(pattern, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// // Extract params using the registered extractor or default to empty map
|
||||||
|
// var params map[string]string
|
||||||
|
// if b.paramExtractor != nil {
|
||||||
|
// params = b.paramExtractor(r)
|
||||||
|
// } else {
|
||||||
|
// params = make(map[string]string)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// adapter := &ResolveSpecHandlerAdapter{
|
||||||
|
// HandlerMethod: handlerMethod,
|
||||||
|
// Params: params,
|
||||||
|
// RequestAdapter: b.requestAdapter,
|
||||||
|
// ResponseAdapter: b.responseAdapter,
|
||||||
|
// }
|
||||||
|
// authHandler := WrapSpecHandler(b.securityList, adapter, optional)
|
||||||
|
// authHandler.ServeHTTP(w, r)
|
||||||
|
// }).Methods(methods...)
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // SetParamExtractor sets a custom parameter extractor function
|
||||||
|
// // For Gorilla Mux, you would use: builder.SetParamExtractor(mux.Vars)
|
||||||
|
// func (b *MuxRouteBuilder) SetParamExtractor(extractor func(*http.Request) map[string]string) {
|
||||||
|
// b.paramExtractor = extractor
|
||||||
|
// }
|
||||||
|
|
||||||
|
// // SetupAuthenticatedSpecRoutes sets up all standard resolvespec/restheadspec routes with authentication
|
||||||
|
// // This is a convenience function that sets up the common route patterns
|
||||||
|
// //
|
||||||
|
// // Usage:
|
||||||
|
// //
|
||||||
|
// // security.SetupAuthenticatedSpecRoutes(router, handler, securityList, router.NewHTTPRequest, router.NewHTTPResponseWriter, mux.Vars)
|
||||||
|
// func SetupAuthenticatedSpecRoutes(
|
||||||
|
// router interface {
|
||||||
|
// HandleFunc(pattern string, f func(http.ResponseWriter, *http.Request)) interface{ Methods(...string) interface{} }
|
||||||
|
// },
|
||||||
|
// handler interface {
|
||||||
|
// Handle(w any, r any, params map[string]string)
|
||||||
|
// HandleGet(w any, r any, params map[string]string)
|
||||||
|
// },
|
||||||
|
// securityList *SecurityList,
|
||||||
|
// requestAdapter func(*http.Request) any,
|
||||||
|
// responseAdapter func(http.ResponseWriter) any,
|
||||||
|
// paramExtractor func(*http.Request) map[string]string,
|
||||||
|
// ) {
|
||||||
|
// builder := NewMuxRouteBuilder(securityList, requestAdapter, responseAdapter)
|
||||||
|
// builder.SetParamExtractor(paramExtractor)
|
||||||
|
|
||||||
|
// // POST /{schema}/{entity}
|
||||||
|
// builder.HandleAuth(router, "/{schema}/{entity}", handler.Handle, false, "POST")
|
||||||
|
|
||||||
|
// // POST /{schema}/{entity}/{id}
|
||||||
|
// builder.HandleAuth(router, "/{schema}/{entity}/{id}", handler.Handle, false, "POST")
|
||||||
|
|
||||||
|
// // GET /{schema}/{entity}
|
||||||
|
// builder.HandleAuth(router, "/{schema}/{entity}", handler.HandleGet, false, "GET")
|
||||||
|
// }
|
||||||
+615
@@ -0,0 +1,615 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gorilla/mux"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Example: OAuth2 Authentication with Google
|
||||||
|
func ExampleOAuth2Google() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
// Create OAuth2 authenticator for Google
|
||||||
|
oauth2Auth := NewGoogleAuthenticator(
|
||||||
|
"your-client-id",
|
||||||
|
"your-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Login endpoint - redirects to Google
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := oauth2Auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := oauth2Auth.OAuth2GetAuthURL("google", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Callback endpoint - handles Google response
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := oauth2Auth.OAuth2HandleCallback(r.Context(), "google", code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set session cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: true,
|
||||||
|
SameSite: http.SameSiteLaxMode,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Return user info as JSON
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: OAuth2 Authentication with GitHub
|
||||||
|
func ExampleOAuth2GitHub() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
oauth2Auth := NewGitHubAuthenticator(
|
||||||
|
"your-github-client-id",
|
||||||
|
"your-github-client-secret",
|
||||||
|
"http://localhost:8080/auth/github/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/github/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := oauth2Auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := oauth2Auth.OAuth2GetAuthURL("github", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/github/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := oauth2Auth.OAuth2HandleCallback(r.Context(), "github", code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: Custom OAuth2 Provider
|
||||||
|
func ExampleOAuth2Custom() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
// Custom OAuth2 provider configuration
|
||||||
|
oauth2Auth := NewDatabaseAuthenticator(db).WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: "your-client-id",
|
||||||
|
ClientSecret: "your-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://your-provider.com/oauth/authorize",
|
||||||
|
TokenURL: "https://your-provider.com/oauth/token",
|
||||||
|
UserInfoURL: "https://your-provider.com/oauth/userinfo",
|
||||||
|
ProviderName: "custom-provider",
|
||||||
|
|
||||||
|
// Custom user info parser
|
||||||
|
UserInfoParser: func(userInfo map[string]any) (*UserContext, error) {
|
||||||
|
// Extract custom fields from your provider
|
||||||
|
return &UserContext{
|
||||||
|
UserName: userInfo["username"].(string),
|
||||||
|
Email: userInfo["email"].(string),
|
||||||
|
RemoteID: userInfo["id"].(string),
|
||||||
|
UserLevel: 1,
|
||||||
|
Roles: []string{"user"},
|
||||||
|
Claims: userInfo,
|
||||||
|
}, nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := oauth2Auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := oauth2Auth.OAuth2GetAuthURL("custom-provider", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := oauth2Auth.OAuth2HandleCallback(r.Context(), "custom-provider", code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: Multi-Provider OAuth2 with Security Integration
|
||||||
|
func ExampleOAuth2MultiProvider() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
// Create OAuth2 authenticators for multiple providers
|
||||||
|
googleAuth := NewGoogleAuthenticator(
|
||||||
|
"google-client-id",
|
||||||
|
"google-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
githubAuth := NewGitHubAuthenticator(
|
||||||
|
"github-client-id",
|
||||||
|
"github-client-secret",
|
||||||
|
"http://localhost:8080/auth/github/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Create column and row security providers
|
||||||
|
colSec := NewDatabaseColumnSecurityProvider(db)
|
||||||
|
rowSec := NewDatabaseRowSecurityProvider(db)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Google OAuth2 routes
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := googleAuth.OAuth2GenerateState()
|
||||||
|
authURL, _ := googleAuth.OAuth2GetAuthURL("google", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := googleAuth.OAuth2HandleCallback(r.Context(), "google", code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.Redirect(w, r, "/dashboard", http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
// GitHub OAuth2 routes
|
||||||
|
router.HandleFunc("/auth/github/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := githubAuth.OAuth2GenerateState()
|
||||||
|
authURL, _ := githubAuth.OAuth2GetAuthURL("github", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/github/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := githubAuth.OAuth2HandleCallback(r.Context(), "github", code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.Redirect(w, r, "/dashboard", http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Use Google auth for protected routes (or GitHub - both work)
|
||||||
|
provider, _ := NewCompositeSecurityProvider(googleAuth, colSec, rowSec)
|
||||||
|
securityList, _ := NewSecurityList(provider)
|
||||||
|
|
||||||
|
// Protected route with authentication
|
||||||
|
protectedRouter := router.PathPrefix("/api").Subrouter()
|
||||||
|
protectedRouter.Use(NewAuthMiddleware(securityList))
|
||||||
|
protectedRouter.Use(SetSecurityMiddleware(securityList))
|
||||||
|
|
||||||
|
protectedRouter.HandleFunc("/profile", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := GetUserContext(r.Context())
|
||||||
|
_ = json.NewEncoder(w).Encode(userCtx)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: OAuth2 with Token Refresh
|
||||||
|
func ExampleOAuth2TokenRefresh() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
oauth2Auth := NewGoogleAuthenticator(
|
||||||
|
"your-client-id",
|
||||||
|
"your-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Refresh token endpoint
|
||||||
|
router.HandleFunc("/auth/refresh", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
RefreshToken string `json:"refresh_token"`
|
||||||
|
Provider string `json:"provider"` // "google", "github", etc.
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
http.Error(w, "Invalid request", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default to google if not specified
|
||||||
|
if req.Provider == "" {
|
||||||
|
req.Provider = "google"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use OAuth2-specific refresh method
|
||||||
|
loginResp, err := oauth2Auth.OAuth2RefreshToken(r.Context(), req.RefreshToken, req.Provider)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set new session cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: true,
|
||||||
|
SameSite: http.SameSiteLaxMode,
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: OAuth2 Logout
|
||||||
|
func ExampleOAuth2Logout() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
oauth2Auth := NewGoogleAuthenticator(
|
||||||
|
"your-client-id",
|
||||||
|
"your-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/logout", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
token := r.Header.Get("Authorization")
|
||||||
|
if token == "" {
|
||||||
|
cookie, err := r.Cookie("session_token")
|
||||||
|
if err == nil {
|
||||||
|
token = cookie.Value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if token != "" {
|
||||||
|
// Get user ID from session
|
||||||
|
userCtx, err := oauth2Auth.Authenticate(r)
|
||||||
|
if err == nil {
|
||||||
|
_ = oauth2Auth.Logout(r.Context(), LogoutRequest{
|
||||||
|
Token: token,
|
||||||
|
UserID: userCtx.UserID,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: "",
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: -1,
|
||||||
|
HttpOnly: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte("Logged out successfully"))
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: Complete OAuth2 Integration with Database Setup
|
||||||
|
func ExampleOAuth2Complete() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
// Create tables (run once)
|
||||||
|
setupOAuth2Tables(db)
|
||||||
|
|
||||||
|
// Create OAuth2 authenticator
|
||||||
|
oauth2Auth := NewGoogleAuthenticator(
|
||||||
|
"your-client-id",
|
||||||
|
"your-client-secret",
|
||||||
|
"http://localhost:8080/auth/google/callback",
|
||||||
|
db,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Create security providers
|
||||||
|
colSec := NewDatabaseColumnSecurityProvider(db)
|
||||||
|
rowSec := NewDatabaseRowSecurityProvider(db)
|
||||||
|
provider, _ := NewCompositeSecurityProvider(oauth2Auth, colSec, rowSec)
|
||||||
|
securityList, _ := NewSecurityList(provider)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Public routes
|
||||||
|
router.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
_, _ = w.Write([]byte("Welcome! <a href='/auth/google/login'>Login with Google</a>"))
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := oauth2Auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := oauth2Auth.OAuth2GetAuthURL("github", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
state := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
loginResp, err := oauth2Auth.OAuth2HandleCallback(r.Context(), "github", code, state)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResp.Token,
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: int(loginResp.ExpiresIn),
|
||||||
|
HttpOnly: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.Redirect(w, r, "/dashboard", http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Protected routes
|
||||||
|
protectedRouter := router.PathPrefix("/").Subrouter()
|
||||||
|
protectedRouter.Use(NewAuthMiddleware(securityList))
|
||||||
|
protectedRouter.Use(SetSecurityMiddleware(securityList))
|
||||||
|
|
||||||
|
protectedRouter.HandleFunc("/dashboard", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := GetUserContext(r.Context())
|
||||||
|
_, _ = fmt.Fprintf(w, "Welcome, %s! Your email: %s", userCtx.UserName, userCtx.Email)
|
||||||
|
})
|
||||||
|
|
||||||
|
protectedRouter.HandleFunc("/api/profile", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := GetUserContext(r.Context())
|
||||||
|
_ = json.NewEncoder(w).Encode(userCtx)
|
||||||
|
})
|
||||||
|
|
||||||
|
protectedRouter.HandleFunc("/auth/logout", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := GetUserContext(r.Context())
|
||||||
|
_ = oauth2Auth.Logout(r.Context(), LogoutRequest{
|
||||||
|
Token: userCtx.SessionID,
|
||||||
|
UserID: userCtx.UserID,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: "",
|
||||||
|
Path: "/",
|
||||||
|
MaxAge: -1,
|
||||||
|
HttpOnly: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
http.Redirect(w, r, "/", http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupOAuth2Tables(db *sql.DB) {
|
||||||
|
// Create tables from database_schema.sql
|
||||||
|
// This is a helper function - in production, use migrations
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Create users table if not exists
|
||||||
|
_, _ = db.ExecContext(ctx, `
|
||||||
|
CREATE TABLE IF NOT EXISTS users (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
username VARCHAR(255) NOT NULL UNIQUE,
|
||||||
|
email VARCHAR(255) NOT NULL UNIQUE,
|
||||||
|
password VARCHAR(255),
|
||||||
|
user_level INTEGER DEFAULT 0,
|
||||||
|
roles VARCHAR(500),
|
||||||
|
is_active BOOLEAN DEFAULT true,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_login_at TIMESTAMP,
|
||||||
|
remote_id VARCHAR(255),
|
||||||
|
auth_provider VARCHAR(50)
|
||||||
|
)
|
||||||
|
`)
|
||||||
|
|
||||||
|
// Create user_sessions table (used for both regular and OAuth2 sessions)
|
||||||
|
_, _ = db.ExecContext(ctx, `
|
||||||
|
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||||
|
id SERIAL PRIMARY KEY,
|
||||||
|
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||||
|
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
expires_at TIMESTAMP NOT NULL,
|
||||||
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_activity_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
ip_address VARCHAR(45),
|
||||||
|
user_agent TEXT,
|
||||||
|
access_token TEXT,
|
||||||
|
refresh_token TEXT,
|
||||||
|
token_type VARCHAR(50) DEFAULT 'Bearer',
|
||||||
|
auth_provider VARCHAR(50)
|
||||||
|
)
|
||||||
|
`)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Example: All OAuth2 Providers at Once
|
||||||
|
func ExampleOAuth2AllProviders() {
|
||||||
|
db, _ := sql.Open("postgres", "connection-string")
|
||||||
|
|
||||||
|
// Create authenticator with ALL OAuth2 providers
|
||||||
|
auth := NewDatabaseAuthenticator(db).
|
||||||
|
WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: "google-client-id",
|
||||||
|
ClientSecret: "google-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/google/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://accounts.google.com/o/oauth2/auth",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
||||||
|
ProviderName: "google",
|
||||||
|
}).
|
||||||
|
WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: "github-client-id",
|
||||||
|
ClientSecret: "github-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/github/callback",
|
||||||
|
Scopes: []string{"user:email"},
|
||||||
|
AuthURL: "https://github.com/login/oauth/authorize",
|
||||||
|
TokenURL: "https://github.com/login/oauth/access_token",
|
||||||
|
UserInfoURL: "https://api.github.com/user",
|
||||||
|
ProviderName: "github",
|
||||||
|
}).
|
||||||
|
WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: "microsoft-client-id",
|
||||||
|
ClientSecret: "microsoft-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/microsoft/callback",
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize",
|
||||||
|
TokenURL: "https://login.microsoftonline.com/common/oauth2/v2.0/token",
|
||||||
|
UserInfoURL: "https://graph.microsoft.com/v1.0/me",
|
||||||
|
ProviderName: "microsoft",
|
||||||
|
}).
|
||||||
|
WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: "facebook-client-id",
|
||||||
|
ClientSecret: "facebook-client-secret",
|
||||||
|
RedirectURL: "http://localhost:8080/auth/facebook/callback",
|
||||||
|
Scopes: []string{"email"},
|
||||||
|
AuthURL: "https://www.facebook.com/v12.0/dialog/oauth",
|
||||||
|
TokenURL: "https://graph.facebook.com/v12.0/oauth/access_token",
|
||||||
|
UserInfoURL: "https://graph.facebook.com/me?fields=id,name,email",
|
||||||
|
ProviderName: "facebook",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Get list of configured providers
|
||||||
|
providers := auth.OAuth2GetProviders()
|
||||||
|
fmt.Printf("Configured OAuth2 providers: %v\n", providers)
|
||||||
|
|
||||||
|
router := mux.NewRouter()
|
||||||
|
|
||||||
|
// Google routes
|
||||||
|
router.HandleFunc("/auth/google/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("google", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
router.HandleFunc("/auth/google/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, err := auth.OAuth2HandleCallback(r.Context(), "google", r.URL.Query().Get("code"), r.URL.Query().Get("state"))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
// GitHub routes
|
||||||
|
router.HandleFunc("/auth/github/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("github", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
router.HandleFunc("/auth/github/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, err := auth.OAuth2HandleCallback(r.Context(), "github", r.URL.Query().Get("code"), r.URL.Query().Get("state"))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Microsoft routes
|
||||||
|
router.HandleFunc("/auth/microsoft/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("microsoft", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
router.HandleFunc("/auth/microsoft/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, err := auth.OAuth2HandleCallback(r.Context(), "microsoft", r.URL.Query().Get("code"), r.URL.Query().Get("state"))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Facebook routes
|
||||||
|
router.HandleFunc("/auth/facebook/login", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
state, _ := auth.OAuth2GenerateState()
|
||||||
|
authURL, _ := auth.OAuth2GetAuthURL("facebook", state)
|
||||||
|
http.Redirect(w, r, authURL, http.StatusTemporaryRedirect)
|
||||||
|
})
|
||||||
|
router.HandleFunc("/auth/facebook/callback", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
loginResp, err := auth.OAuth2HandleCallback(r.Context(), "facebook", r.URL.Query().Get("code"), r.URL.Query().Get("state"))
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResp)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Create security list for protected routes
|
||||||
|
colSec := NewDatabaseColumnSecurityProvider(db)
|
||||||
|
rowSec := NewDatabaseRowSecurityProvider(db)
|
||||||
|
provider, _ := NewCompositeSecurityProvider(auth, colSec, rowSec)
|
||||||
|
securityList, _ := NewSecurityList(provider)
|
||||||
|
|
||||||
|
// Protected routes work for ALL OAuth2 providers + regular sessions
|
||||||
|
protectedRouter := router.PathPrefix("/api").Subrouter()
|
||||||
|
protectedRouter.Use(NewAuthMiddleware(securityList))
|
||||||
|
protectedRouter.Use(SetSecurityMiddleware(securityList))
|
||||||
|
|
||||||
|
protectedRouter.HandleFunc("/profile", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, _ := GetUserContext(r.Context())
|
||||||
|
_ = json.NewEncoder(w).Encode(userCtx)
|
||||||
|
})
|
||||||
|
|
||||||
|
_ = http.ListenAndServe(":8080", router)
|
||||||
|
}
|
||||||
+579
@@ -0,0 +1,579 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"golang.org/x/oauth2"
|
||||||
|
)
|
||||||
|
|
||||||
|
// OAuth2Config contains configuration for OAuth2 authentication
|
||||||
|
type OAuth2Config struct {
|
||||||
|
ClientID string
|
||||||
|
ClientSecret string
|
||||||
|
RedirectURL string
|
||||||
|
Scopes []string
|
||||||
|
AuthURL string
|
||||||
|
TokenURL string
|
||||||
|
UserInfoURL string
|
||||||
|
ProviderName string
|
||||||
|
|
||||||
|
// Optional: Custom user info parser
|
||||||
|
// If not provided, will use standard claims (sub, email, name)
|
||||||
|
UserInfoParser func(userInfo map[string]any) (*UserContext, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuth2Provider holds configuration and state for a single OAuth2 provider
|
||||||
|
type OAuth2Provider struct {
|
||||||
|
config *oauth2.Config
|
||||||
|
userInfoURL string
|
||||||
|
userInfoParser func(userInfo map[string]any) (*UserContext, error)
|
||||||
|
providerName string
|
||||||
|
states map[string]time.Time // state -> expiry time
|
||||||
|
statesMutex sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithOAuth2 configures OAuth2 support for the DatabaseAuthenticator
|
||||||
|
// Can be called multiple times to add multiple OAuth2 providers
|
||||||
|
// Returns the same DatabaseAuthenticator instance for method chaining
|
||||||
|
func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthenticator {
|
||||||
|
if cfg.ProviderName == "" {
|
||||||
|
cfg.ProviderName = "oauth2"
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.UserInfoParser == nil {
|
||||||
|
cfg.UserInfoParser = defaultOAuth2UserInfoParser
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := &OAuth2Provider{
|
||||||
|
config: &oauth2.Config{
|
||||||
|
ClientID: cfg.ClientID,
|
||||||
|
ClientSecret: cfg.ClientSecret,
|
||||||
|
RedirectURL: cfg.RedirectURL,
|
||||||
|
Scopes: cfg.Scopes,
|
||||||
|
Endpoint: oauth2.Endpoint{
|
||||||
|
AuthURL: cfg.AuthURL,
|
||||||
|
TokenURL: cfg.TokenURL,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
userInfoURL: cfg.UserInfoURL,
|
||||||
|
userInfoParser: cfg.UserInfoParser,
|
||||||
|
providerName: cfg.ProviderName,
|
||||||
|
states: make(map[string]time.Time),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize providers map if needed
|
||||||
|
a.oauth2ProvidersMutex.Lock()
|
||||||
|
if a.oauth2Providers == nil {
|
||||||
|
a.oauth2Providers = make(map[string]*OAuth2Provider)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register provider
|
||||||
|
a.oauth2Providers[cfg.ProviderName] = provider
|
||||||
|
a.oauth2ProvidersMutex.Unlock()
|
||||||
|
|
||||||
|
// Start state cleanup goroutine for this provider
|
||||||
|
go provider.cleanupStates()
|
||||||
|
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuth2GetAuthURL returns the OAuth2 authorization URL for redirecting users
|
||||||
|
func (a *DatabaseAuthenticator) OAuth2GetAuthURL(providerName, state string) (string, error) {
|
||||||
|
provider, err := a.getOAuth2Provider(providerName)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store state for validation
|
||||||
|
provider.statesMutex.Lock()
|
||||||
|
provider.states[state] = time.Now().Add(10 * time.Minute)
|
||||||
|
provider.statesMutex.Unlock()
|
||||||
|
|
||||||
|
return provider.config.AuthCodeURL(state), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuth2GenerateState generates a random state string for CSRF protection
|
||||||
|
func (a *DatabaseAuthenticator) OAuth2GenerateState() (string, error) {
|
||||||
|
b := make([]byte, 32)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return base64.URLEncoding.EncodeToString(b), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuth2HandleCallback handles the OAuth2 callback and exchanges code for token
|
||||||
|
func (a *DatabaseAuthenticator) OAuth2HandleCallback(ctx context.Context, providerName, code, state string) (*LoginResponse, error) {
|
||||||
|
provider, err := a.getOAuth2Provider(providerName)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate state
|
||||||
|
if !provider.validateState(state) {
|
||||||
|
return nil, fmt.Errorf("invalid state parameter")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exchange code for token
|
||||||
|
token, err := provider.config.Exchange(ctx, code)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fetch user info
|
||||||
|
client := provider.config.Client(ctx, token)
|
||||||
|
resp, err := client.Get(provider.userInfoURL)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to fetch user info: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to read user info: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var userInfo map[string]any
|
||||||
|
if err := json.Unmarshal(body, &userInfo); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user info: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse user info
|
||||||
|
userCtx, err := provider.userInfoParser(userInfo)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get or create user in database
|
||||||
|
userID, err := a.oauth2GetOrCreateUser(ctx, userCtx, providerName)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get or create user: %w", err)
|
||||||
|
}
|
||||||
|
userCtx.UserID = userID
|
||||||
|
|
||||||
|
// Create session token
|
||||||
|
sessionToken, err := a.OAuth2GenerateState()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
expiresAt := time.Now().Add(24 * time.Hour)
|
||||||
|
if token.Expiry.After(time.Now()) {
|
||||||
|
expiresAt = token.Expiry
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store session in database
|
||||||
|
err = a.oauth2CreateSession(ctx, sessionToken, userCtx.UserID, token, expiresAt, providerName)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
userCtx.SessionID = sessionToken
|
||||||
|
|
||||||
|
return &LoginResponse{
|
||||||
|
Token: sessionToken,
|
||||||
|
RefreshToken: token.RefreshToken,
|
||||||
|
User: userCtx,
|
||||||
|
ExpiresIn: int64(time.Until(expiresAt).Seconds()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuth2GetProviders returns list of configured OAuth2 provider names
|
||||||
|
func (a *DatabaseAuthenticator) OAuth2GetProviders() []string {
|
||||||
|
a.oauth2ProvidersMutex.RLock()
|
||||||
|
defer a.oauth2ProvidersMutex.RUnlock()
|
||||||
|
|
||||||
|
if a.oauth2Providers == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
providers := make([]string, 0, len(a.oauth2Providers))
|
||||||
|
for name := range a.oauth2Providers {
|
||||||
|
providers = append(providers, name)
|
||||||
|
}
|
||||||
|
return providers
|
||||||
|
}
|
||||||
|
|
||||||
|
// getOAuth2Provider retrieves a registered OAuth2 provider by name
|
||||||
|
func (a *DatabaseAuthenticator) getOAuth2Provider(providerName string) (*OAuth2Provider, error) {
|
||||||
|
a.oauth2ProvidersMutex.RLock()
|
||||||
|
defer a.oauth2ProvidersMutex.RUnlock()
|
||||||
|
|
||||||
|
if a.oauth2Providers == nil {
|
||||||
|
return nil, fmt.Errorf("OAuth2 not configured - call WithOAuth2() first")
|
||||||
|
}
|
||||||
|
|
||||||
|
provider, ok := a.oauth2Providers[providerName]
|
||||||
|
if !ok {
|
||||||
|
// Build provider list without calling OAuth2GetProviders to avoid recursion
|
||||||
|
providerNames := make([]string, 0, len(a.oauth2Providers))
|
||||||
|
for name := range a.oauth2Providers {
|
||||||
|
providerNames = append(providerNames, name)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("OAuth2 provider '%s' not found - available providers: %v", providerName, providerNames)
|
||||||
|
}
|
||||||
|
|
||||||
|
return provider, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// oauth2GetOrCreateUser finds or creates a user based on OAuth2 info using stored procedure
|
||||||
|
func (a *DatabaseAuthenticator) oauth2GetOrCreateUser(ctx context.Context, userCtx *UserContext, providerName string) (int, error) {
|
||||||
|
userData := map[string]interface{}{
|
||||||
|
"username": userCtx.UserName,
|
||||||
|
"email": userCtx.Email,
|
||||||
|
"remote_id": userCtx.RemoteID,
|
||||||
|
"user_level": userCtx.UserLevel,
|
||||||
|
"roles": userCtx.Roles,
|
||||||
|
"auth_provider": providerName,
|
||||||
|
}
|
||||||
|
|
||||||
|
userJSON, err := json.Marshal(userData)
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to marshal user data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
var userID *int
|
||||||
|
|
||||||
|
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_user_id
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, a.sqlNames.OAuthGetOrCreateUser), userJSON).Scan(&success, &errMsg, &userID)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to get or create user: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return 0, fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return 0, fmt.Errorf("failed to get or create user")
|
||||||
|
}
|
||||||
|
|
||||||
|
if userID == nil {
|
||||||
|
return 0, fmt.Errorf("user ID not returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
return *userID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// oauth2CreateSession creates a new OAuth2 session using stored procedure
|
||||||
|
func (a *DatabaseAuthenticator) oauth2CreateSession(ctx context.Context, sessionToken string, userID int, token *oauth2.Token, expiresAt time.Time, providerName string) error {
|
||||||
|
sessionData := map[string]interface{}{
|
||||||
|
"session_token": sessionToken,
|
||||||
|
"user_id": userID,
|
||||||
|
"access_token": token.AccessToken,
|
||||||
|
"refresh_token": token.RefreshToken,
|
||||||
|
"token_type": token.TokenType,
|
||||||
|
"expires_at": expiresAt,
|
||||||
|
"auth_provider": providerName,
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionJSON, err := json.Marshal(sessionData)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal session data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
|
||||||
|
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, a.sqlNames.OAuthCreateSession), sessionJSON).Scan(&success, &errMsg)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to create session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to create session")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateState validates state using in-memory storage
|
||||||
|
func (p *OAuth2Provider) validateState(state string) bool {
|
||||||
|
p.statesMutex.Lock()
|
||||||
|
defer p.statesMutex.Unlock()
|
||||||
|
|
||||||
|
expiry, ok := p.states[state]
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if time.Now().After(expiry) {
|
||||||
|
delete(p.states, state)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(p.states, state) // One-time use
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// cleanupStates removes expired states periodically
|
||||||
|
func (p *OAuth2Provider) cleanupStates() {
|
||||||
|
ticker := time.NewTicker(5 * time.Minute)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for range ticker.C {
|
||||||
|
p.statesMutex.Lock()
|
||||||
|
now := time.Now()
|
||||||
|
for state, expiry := range p.states {
|
||||||
|
if now.After(expiry) {
|
||||||
|
delete(p.states, state)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
p.statesMutex.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// defaultOAuth2UserInfoParser parses standard OAuth2 user info claims
|
||||||
|
func defaultOAuth2UserInfoParser(userInfo map[string]any) (*UserContext, error) {
|
||||||
|
ctx := &UserContext{
|
||||||
|
Claims: userInfo,
|
||||||
|
Roles: []string{"user"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract standard claims
|
||||||
|
if sub, ok := userInfo["sub"].(string); ok {
|
||||||
|
ctx.RemoteID = sub
|
||||||
|
}
|
||||||
|
if email, ok := userInfo["email"].(string); ok {
|
||||||
|
ctx.Email = email
|
||||||
|
// Use email as username if name not available
|
||||||
|
ctx.UserName = strings.Split(email, "@")[0]
|
||||||
|
}
|
||||||
|
if name, ok := userInfo["name"].(string); ok {
|
||||||
|
ctx.UserName = name
|
||||||
|
}
|
||||||
|
if login, ok := userInfo["login"].(string); ok {
|
||||||
|
ctx.UserName = login // GitHub uses "login"
|
||||||
|
}
|
||||||
|
|
||||||
|
if ctx.UserName == "" {
|
||||||
|
return nil, fmt.Errorf("could not extract username from user info")
|
||||||
|
}
|
||||||
|
|
||||||
|
return ctx, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuth2RefreshToken refreshes an expired OAuth2 access token using the refresh token
|
||||||
|
// Takes the refresh token and returns a new LoginResponse with updated tokens
|
||||||
|
func (a *DatabaseAuthenticator) OAuth2RefreshToken(ctx context.Context, refreshToken, providerName string) (*LoginResponse, error) {
|
||||||
|
provider, err := a.getOAuth2Provider(providerName)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get session by refresh token from database
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
var sessionData []byte
|
||||||
|
|
||||||
|
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, a.sqlNames.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &sessionData)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("invalid or expired refresh token")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse session data
|
||||||
|
var session struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
AccessToken string `json:"access_token"`
|
||||||
|
TokenType string `json:"token_type"`
|
||||||
|
Expiry time.Time `json:"expiry"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(sessionData, &session); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse session data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create oauth2.Token from stored data
|
||||||
|
oldToken := &oauth2.Token{
|
||||||
|
AccessToken: session.AccessToken,
|
||||||
|
TokenType: session.TokenType,
|
||||||
|
RefreshToken: refreshToken,
|
||||||
|
Expiry: session.Expiry,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use OAuth2 provider to refresh the token
|
||||||
|
tokenSource := provider.config.TokenSource(ctx, oldToken)
|
||||||
|
newToken, err := tokenSource.Token()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to refresh token with provider: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate new session token
|
||||||
|
newSessionToken, err := a.OAuth2GenerateState()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate new session token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update session in database with new tokens
|
||||||
|
updateData := map[string]interface{}{
|
||||||
|
"user_id": session.UserID,
|
||||||
|
"old_refresh_token": refreshToken,
|
||||||
|
"new_session_token": newSessionToken,
|
||||||
|
"new_access_token": newToken.AccessToken,
|
||||||
|
"new_refresh_token": newToken.RefreshToken,
|
||||||
|
"expires_at": newToken.Expiry,
|
||||||
|
}
|
||||||
|
|
||||||
|
updateJSON, err := json.Marshal(updateData)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal update data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var updateSuccess bool
|
||||||
|
var updateErrMsg *string
|
||||||
|
|
||||||
|
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, a.sqlNames.OAuthUpdateRefreshToken), updateJSON).Scan(&updateSuccess, &updateErrMsg)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to update session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !updateSuccess {
|
||||||
|
if updateErrMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *updateErrMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to update session")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get user data
|
||||||
|
var userSuccess bool
|
||||||
|
var userErrMsg *string
|
||||||
|
var userData []byte
|
||||||
|
|
||||||
|
err = a.getDB().QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, a.sqlNames.OAuthGetUser), session.UserID).Scan(&userSuccess, &userErrMsg, &userData)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get user data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !userSuccess {
|
||||||
|
if userErrMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *userErrMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to get user data")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse user context
|
||||||
|
var userCtx UserContext
|
||||||
|
if err := json.Unmarshal(userData, &userCtx); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
userCtx.SessionID = newSessionToken
|
||||||
|
|
||||||
|
return &LoginResponse{
|
||||||
|
Token: newSessionToken,
|
||||||
|
RefreshToken: newToken.RefreshToken,
|
||||||
|
User: &userCtx,
|
||||||
|
ExpiresIn: int64(time.Until(newToken.Expiry).Seconds()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pre-configured OAuth2 factory methods
|
||||||
|
|
||||||
|
// NewGoogleAuthenticator creates a DatabaseAuthenticator configured for Google OAuth2
|
||||||
|
func NewGoogleAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
||||||
|
auth := NewDatabaseAuthenticator(db)
|
||||||
|
return auth.WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: clientSecret,
|
||||||
|
RedirectURL: redirectURL,
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://accounts.google.com/o/oauth2/auth",
|
||||||
|
TokenURL: "https://oauth2.googleapis.com/token",
|
||||||
|
UserInfoURL: "https://www.googleapis.com/oauth2/v2/userinfo",
|
||||||
|
ProviderName: "google",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewGitHubAuthenticator creates a DatabaseAuthenticator configured for GitHub OAuth2
|
||||||
|
func NewGitHubAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
||||||
|
auth := NewDatabaseAuthenticator(db)
|
||||||
|
return auth.WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: clientSecret,
|
||||||
|
RedirectURL: redirectURL,
|
||||||
|
Scopes: []string{"user:email"},
|
||||||
|
AuthURL: "https://github.com/login/oauth/authorize",
|
||||||
|
TokenURL: "https://github.com/login/oauth/access_token",
|
||||||
|
UserInfoURL: "https://api.github.com/user",
|
||||||
|
ProviderName: "github",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMicrosoftAuthenticator creates a DatabaseAuthenticator configured for Microsoft OAuth2
|
||||||
|
func NewMicrosoftAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
||||||
|
auth := NewDatabaseAuthenticator(db)
|
||||||
|
return auth.WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: clientSecret,
|
||||||
|
RedirectURL: redirectURL,
|
||||||
|
Scopes: []string{"openid", "profile", "email"},
|
||||||
|
AuthURL: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize",
|
||||||
|
TokenURL: "https://login.microsoftonline.com/common/oauth2/v2.0/token",
|
||||||
|
UserInfoURL: "https://graph.microsoft.com/v1.0/me",
|
||||||
|
ProviderName: "microsoft",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewFacebookAuthenticator creates a DatabaseAuthenticator configured for Facebook OAuth2
|
||||||
|
func NewFacebookAuthenticator(clientID, clientSecret, redirectURL string, db *sql.DB) *DatabaseAuthenticator {
|
||||||
|
auth := NewDatabaseAuthenticator(db)
|
||||||
|
return auth.WithOAuth2(OAuth2Config{
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: clientSecret,
|
||||||
|
RedirectURL: redirectURL,
|
||||||
|
Scopes: []string{"email"},
|
||||||
|
AuthURL: "https://www.facebook.com/v12.0/dialog/oauth",
|
||||||
|
TokenURL: "https://graph.facebook.com/v12.0/oauth/access_token",
|
||||||
|
UserInfoURL: "https://graph.facebook.com/me?fields=id,name,email",
|
||||||
|
ProviderName: "facebook",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMultiProviderAuthenticator creates a DatabaseAuthenticator with all major OAuth2 providers configured
|
||||||
|
func NewMultiProviderAuthenticator(db *sql.DB, configs map[string]OAuth2Config) *DatabaseAuthenticator {
|
||||||
|
auth := NewDatabaseAuthenticator(db)
|
||||||
|
|
||||||
|
//nolint:gocritic // OAuth2Config is copied but kept for API simplicity
|
||||||
|
for _, cfg := range configs {
|
||||||
|
auth.WithOAuth2(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
return auth
|
||||||
|
}
|
||||||
+917
@@ -0,0 +1,917 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// OAuthServerConfig configures the MCP-standard OAuth2 authorization server.
|
||||||
|
type OAuthServerConfig struct {
|
||||||
|
// Issuer is the public base URL of this server (e.g. "https://api.example.com").
|
||||||
|
// Used in /.well-known/oauth-authorization-server and to build endpoint URLs.
|
||||||
|
Issuer string
|
||||||
|
|
||||||
|
// ProviderCallbackPath is the path on this server that external OAuth2 providers
|
||||||
|
// redirect back to. Defaults to "/oauth/provider/callback".
|
||||||
|
ProviderCallbackPath string
|
||||||
|
|
||||||
|
// LoginTitle is shown on the built-in login form when the server acts as its own
|
||||||
|
// identity provider. Defaults to "Sign in".
|
||||||
|
LoginTitle string
|
||||||
|
|
||||||
|
// PersistClients stores registered clients in the database when a DatabaseAuthenticator is provided.
|
||||||
|
// Clients registered during a session survive server restarts.
|
||||||
|
PersistClients bool
|
||||||
|
|
||||||
|
// PersistCodes stores authorization codes in the database.
|
||||||
|
// Useful for multi-instance deployments. Defaults to in-memory.
|
||||||
|
PersistCodes bool
|
||||||
|
|
||||||
|
// DefaultScopes lists scopes advertised in server metadata. Defaults to ["openid","profile","email"].
|
||||||
|
DefaultScopes []string
|
||||||
|
|
||||||
|
// AccessTokenTTL is the issued token lifetime. Defaults to 24h.
|
||||||
|
AccessTokenTTL time.Duration
|
||||||
|
|
||||||
|
// AuthCodeTTL is the auth code lifetime. Defaults to 2 minutes.
|
||||||
|
AuthCodeTTL time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
// oauthClient is a dynamically registered OAuth2 client (RFC 7591).
|
||||||
|
type oauthClient struct {
|
||||||
|
ClientID string `json:"client_id"`
|
||||||
|
RedirectURIs []string `json:"redirect_uris"`
|
||||||
|
ClientName string `json:"client_name,omitempty"`
|
||||||
|
GrantTypes []string `json:"grant_types"`
|
||||||
|
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// pendingAuth tracks an in-progress authorization code exchange.
|
||||||
|
type pendingAuth struct {
|
||||||
|
ClientID string
|
||||||
|
RedirectURI string
|
||||||
|
ClientState string
|
||||||
|
CodeChallenge string
|
||||||
|
CodeChallengeMethod string
|
||||||
|
ProviderName string // empty = password login
|
||||||
|
ExpiresAt time.Time
|
||||||
|
SessionToken string // set after authentication completes
|
||||||
|
RefreshToken string // set after authentication completes when refresh tokens are issued
|
||||||
|
Scopes []string // requested scopes
|
||||||
|
}
|
||||||
|
|
||||||
|
// externalProvider pairs a DatabaseAuthenticator with its provider name.
|
||||||
|
type externalProvider struct {
|
||||||
|
auth *DatabaseAuthenticator
|
||||||
|
providerName string
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthServer implements the MCP-standard OAuth2 authorization server (OAuth 2.1 + PKCE).
|
||||||
|
//
|
||||||
|
// It can act as both:
|
||||||
|
// - A direct identity provider using DatabaseAuthenticator username/password login
|
||||||
|
// - A federation layer that delegates authentication to external OAuth2 providers
|
||||||
|
// (Google, GitHub, Microsoft, etc.) registered via RegisterExternalProvider
|
||||||
|
//
|
||||||
|
// The server exposes these RFC-compliant endpoints:
|
||||||
|
//
|
||||||
|
// GET /.well-known/oauth-authorization-server RFC 8414 — server metadata discovery
|
||||||
|
// POST /oauth/register RFC 7591 — dynamic client registration
|
||||||
|
// GET /oauth/authorize OAuth 2.1 + PKCE — start authorization
|
||||||
|
// POST /oauth/authorize Direct login form submission
|
||||||
|
// POST /oauth/token Token exchange and refresh
|
||||||
|
// POST /oauth/revoke RFC 7009 — token revocation
|
||||||
|
// POST /oauth/introspect RFC 7662 — token introspection
|
||||||
|
// GET {ProviderCallbackPath} Internal — external provider callback
|
||||||
|
type OAuthServer struct {
|
||||||
|
cfg OAuthServerConfig
|
||||||
|
auth *DatabaseAuthenticator // nil = only external providers
|
||||||
|
providers []externalProvider
|
||||||
|
|
||||||
|
mu sync.RWMutex
|
||||||
|
clients map[string]*oauthClient
|
||||||
|
pending map[string]*pendingAuth // provider_state → pending (external flow)
|
||||||
|
codes map[string]*pendingAuth // auth_code → pending (post-auth)
|
||||||
|
|
||||||
|
done chan struct{} // closed by Close() to stop background goroutines
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOAuthServer creates a new MCP OAuth2 authorization server.
|
||||||
|
//
|
||||||
|
// Pass a DatabaseAuthenticator to enable direct username/password login (the server
|
||||||
|
// acts as its own identity provider). Pass nil to use only external providers.
|
||||||
|
// External providers are added separately via RegisterExternalProvider.
|
||||||
|
//
|
||||||
|
// Call Close() to stop background goroutines when the server is no longer needed.
|
||||||
|
func NewOAuthServer(cfg OAuthServerConfig, auth *DatabaseAuthenticator) *OAuthServer {
|
||||||
|
if cfg.ProviderCallbackPath == "" {
|
||||||
|
cfg.ProviderCallbackPath = "/oauth/provider/callback"
|
||||||
|
}
|
||||||
|
if cfg.LoginTitle == "" {
|
||||||
|
cfg.LoginTitle = "Sign in"
|
||||||
|
}
|
||||||
|
if len(cfg.DefaultScopes) == 0 {
|
||||||
|
cfg.DefaultScopes = []string{"openid", "profile", "email"}
|
||||||
|
}
|
||||||
|
if cfg.AccessTokenTTL == 0 {
|
||||||
|
cfg.AccessTokenTTL = 24 * time.Hour
|
||||||
|
}
|
||||||
|
if cfg.AuthCodeTTL == 0 {
|
||||||
|
cfg.AuthCodeTTL = 2 * time.Minute
|
||||||
|
}
|
||||||
|
// Normalize issuer: remove trailing slash to ensure consistent endpoint URL construction.
|
||||||
|
cfg.Issuer = strings.TrimSuffix(cfg.Issuer, "/")
|
||||||
|
s := &OAuthServer{
|
||||||
|
cfg: cfg,
|
||||||
|
auth: auth,
|
||||||
|
clients: make(map[string]*oauthClient),
|
||||||
|
pending: make(map[string]*pendingAuth),
|
||||||
|
codes: make(map[string]*pendingAuth),
|
||||||
|
done: make(chan struct{}),
|
||||||
|
}
|
||||||
|
go s.cleanupExpired()
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close stops the background goroutines started by NewOAuthServer.
|
||||||
|
// It is safe to call Close multiple times.
|
||||||
|
func (s *OAuthServer) Close() {
|
||||||
|
select {
|
||||||
|
case <-s.done:
|
||||||
|
// already closed
|
||||||
|
default:
|
||||||
|
close(s.done)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterExternalProvider adds an external OAuth2 provider (Google, GitHub, Microsoft, etc.)
|
||||||
|
// that handles user authentication via redirect. The DatabaseAuthenticator must have been
|
||||||
|
// configured with WithOAuth2(providerName, ...) before calling this.
|
||||||
|
// Multiple providers can be registered; the first is used as the default.
|
||||||
|
// All providers must be registered before the server starts serving requests.
|
||||||
|
func (s *OAuthServer) RegisterExternalProvider(auth *DatabaseAuthenticator, providerName string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.providers = append(s.providers, externalProvider{auth: auth, providerName: providerName})
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProviderCallbackPath returns the configured path for external provider callbacks.
|
||||||
|
func (s *OAuthServer) ProviderCallbackPath() string {
|
||||||
|
return s.cfg.ProviderCallbackPath
|
||||||
|
}
|
||||||
|
|
||||||
|
// HTTPHandler returns an http.Handler that serves all RFC-required OAuth2 endpoints.
|
||||||
|
// Mount it at the root of your HTTP server alongside the MCP transport.
|
||||||
|
//
|
||||||
|
// mux := http.NewServeMux()
|
||||||
|
// mux.Handle("/", oauthServer.HTTPHandler())
|
||||||
|
// mux.Handle("/mcp/", mcpTransport)
|
||||||
|
func (s *OAuthServer) HTTPHandler() http.Handler {
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.HandleFunc("/.well-known/oauth-authorization-server", s.metadataHandler)
|
||||||
|
mux.HandleFunc("/oauth/register", s.registerHandler)
|
||||||
|
mux.HandleFunc("/oauth/authorize", s.authorizeHandler)
|
||||||
|
mux.HandleFunc("/oauth/token", s.tokenHandler)
|
||||||
|
mux.HandleFunc("/oauth/revoke", s.revokeHandler)
|
||||||
|
mux.HandleFunc("/oauth/introspect", s.introspectHandler)
|
||||||
|
mux.HandleFunc(s.cfg.ProviderCallbackPath, s.providerCallbackHandler)
|
||||||
|
return mux
|
||||||
|
}
|
||||||
|
|
||||||
|
// cleanupExpired removes stale pending auths and codes every 5 minutes.
|
||||||
|
func (s *OAuthServer) cleanupExpired() {
|
||||||
|
ticker := time.NewTicker(5 * time.Minute)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-s.done:
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
now := time.Now()
|
||||||
|
s.mu.Lock()
|
||||||
|
for k, p := range s.pending {
|
||||||
|
if now.After(p.ExpiresAt) {
|
||||||
|
delete(s.pending, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for k, p := range s.codes {
|
||||||
|
if now.After(p.ExpiresAt) {
|
||||||
|
delete(s.codes, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// RFC 8414 — Server metadata
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) metadataHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
issuer := s.cfg.Issuer
|
||||||
|
meta := map[string]interface{}{
|
||||||
|
"issuer": issuer,
|
||||||
|
"authorization_endpoint": issuer + "/oauth/authorize",
|
||||||
|
"token_endpoint": issuer + "/oauth/token",
|
||||||
|
"registration_endpoint": issuer + "/oauth/register",
|
||||||
|
"revocation_endpoint": issuer + "/oauth/revoke",
|
||||||
|
"introspection_endpoint": issuer + "/oauth/introspect",
|
||||||
|
"scopes_supported": s.cfg.DefaultScopes,
|
||||||
|
"response_types_supported": []string{"code"},
|
||||||
|
"grant_types_supported": []string{"authorization_code", "refresh_token"},
|
||||||
|
"code_challenge_methods_supported": []string{"S256"},
|
||||||
|
"token_endpoint_auth_methods_supported": []string{"none"},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(meta) //nolint:errcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// RFC 7591 — Dynamic client registration
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) registerHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var req struct {
|
||||||
|
RedirectURIs []string `json:"redirect_uris"`
|
||||||
|
ClientName string `json:"client_name"`
|
||||||
|
GrantTypes []string `json:"grant_types"`
|
||||||
|
AllowedScopes []string `json:"allowed_scopes"`
|
||||||
|
}
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
writeOAuthError(w, "invalid_request", "malformed JSON", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if len(req.RedirectURIs) == 0 {
|
||||||
|
writeOAuthError(w, "invalid_request", "redirect_uris required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
grantTypes := req.GrantTypes
|
||||||
|
if len(grantTypes) == 0 {
|
||||||
|
grantTypes = []string{"authorization_code"}
|
||||||
|
}
|
||||||
|
allowedScopes := req.AllowedScopes
|
||||||
|
if len(allowedScopes) == 0 {
|
||||||
|
allowedScopes = s.cfg.DefaultScopes
|
||||||
|
}
|
||||||
|
clientID, err := randomOAuthToken()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "server error", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
client := &oauthClient{
|
||||||
|
ClientID: clientID,
|
||||||
|
RedirectURIs: req.RedirectURIs,
|
||||||
|
ClientName: req.ClientName,
|
||||||
|
GrantTypes: grantTypes,
|
||||||
|
AllowedScopes: allowedScopes,
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.cfg.PersistClients && s.auth != nil {
|
||||||
|
dbClient := &OAuthServerClient{
|
||||||
|
ClientID: client.ClientID,
|
||||||
|
RedirectURIs: client.RedirectURIs,
|
||||||
|
ClientName: client.ClientName,
|
||||||
|
GrantTypes: client.GrantTypes,
|
||||||
|
AllowedScopes: client.AllowedScopes,
|
||||||
|
}
|
||||||
|
if _, err := s.auth.OAuthRegisterClient(r.Context(), dbClient); err != nil {
|
||||||
|
http.Error(w, "server error", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
s.clients[clientID] = client
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusCreated)
|
||||||
|
json.NewEncoder(w).Encode(client) //nolint:errcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Authorization endpoint — GET + POST /oauth/authorize
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) authorizeHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch r.Method {
|
||||||
|
case http.MethodGet:
|
||||||
|
s.authorizeGet(w, r)
|
||||||
|
case http.MethodPost:
|
||||||
|
s.authorizePost(w, r)
|
||||||
|
default:
|
||||||
|
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// authorizeGet validates the request and either:
|
||||||
|
// - Redirects to an external provider (if providers are registered)
|
||||||
|
// - Renders a login form (if the server is its own identity provider)
|
||||||
|
func (s *OAuthServer) authorizeGet(w http.ResponseWriter, r *http.Request) {
|
||||||
|
q := r.URL.Query()
|
||||||
|
clientID := q.Get("client_id")
|
||||||
|
redirectURI := q.Get("redirect_uri")
|
||||||
|
clientState := q.Get("state")
|
||||||
|
codeChallenge := q.Get("code_challenge")
|
||||||
|
codeChallengeMethod := q.Get("code_challenge_method")
|
||||||
|
providerName := q.Get("provider")
|
||||||
|
scopeStr := q.Get("scope")
|
||||||
|
var scopes []string
|
||||||
|
if scopeStr != "" {
|
||||||
|
scopes = strings.Fields(scopeStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if q.Get("response_type") != "code" {
|
||||||
|
writeOAuthError(w, "unsupported_response_type", "only 'code' is supported", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if codeChallenge == "" {
|
||||||
|
writeOAuthError(w, "invalid_request", "code_challenge required (PKCE S256)", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if codeChallengeMethod != "" && codeChallengeMethod != "S256" {
|
||||||
|
writeOAuthError(w, "invalid_request", "only S256 code_challenge_method is supported", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
client, ok := s.lookupOrFetchClient(r.Context(), clientID)
|
||||||
|
if !ok {
|
||||||
|
writeOAuthError(w, "invalid_client", "unknown client_id", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !oauthSliceContains(client.RedirectURIs, redirectURI) {
|
||||||
|
writeOAuthError(w, "invalid_request", "redirect_uri not registered", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// External provider path
|
||||||
|
if len(s.providers) > 0 {
|
||||||
|
s.redirectToExternalProvider(w, r, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, providerName, scopes)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Direct login form path (server is its own identity provider)
|
||||||
|
if s.auth == nil {
|
||||||
|
http.Error(w, "no authentication provider configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.renderLoginForm(w, r, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, scopeStr, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
// authorizePost handles login form submission for the direct login flow.
|
||||||
|
func (s *OAuthServer) authorizePost(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if err := r.ParseForm(); err != nil {
|
||||||
|
http.Error(w, "invalid form", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
clientID := r.FormValue("client_id")
|
||||||
|
redirectURI := r.FormValue("redirect_uri")
|
||||||
|
clientState := r.FormValue("client_state")
|
||||||
|
codeChallenge := r.FormValue("code_challenge")
|
||||||
|
codeChallengeMethod := r.FormValue("code_challenge_method")
|
||||||
|
username := r.FormValue("username")
|
||||||
|
password := r.FormValue("password")
|
||||||
|
scopeStr := r.FormValue("scope")
|
||||||
|
var scopes []string
|
||||||
|
if scopeStr != "" {
|
||||||
|
scopes = strings.Fields(scopeStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
client, ok := s.lookupOrFetchClient(r.Context(), clientID)
|
||||||
|
if !ok || !oauthSliceContains(client.RedirectURIs, redirectURI) {
|
||||||
|
http.Error(w, "invalid client or redirect_uri", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.auth == nil {
|
||||||
|
http.Error(w, "no authentication provider configured", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
loginResp, err := s.auth.Login(r.Context(), LoginRequest{
|
||||||
|
Username: username,
|
||||||
|
Password: password,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
s.renderLoginForm(w, r, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, scopeStr, "Invalid username or password")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s.issueCodeAndRedirect(w, r, loginResp.Token, loginResp.RefreshToken, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, "", scopes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// redirectToExternalProvider stores the pending auth and redirects to the configured provider.
|
||||||
|
func (s *OAuthServer) redirectToExternalProvider(w http.ResponseWriter, r *http.Request, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, providerName string, scopes []string) {
|
||||||
|
var provider *externalProvider
|
||||||
|
if providerName != "" {
|
||||||
|
for i := range s.providers {
|
||||||
|
if s.providers[i].providerName == providerName {
|
||||||
|
provider = &s.providers[i]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, fmt.Sprintf("provider %q not found", providerName), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
provider = &s.providers[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
providerState, err := randomOAuthToken()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "server error", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pending := &pendingAuth{
|
||||||
|
ClientID: clientID,
|
||||||
|
RedirectURI: redirectURI,
|
||||||
|
ClientState: clientState,
|
||||||
|
CodeChallenge: codeChallenge,
|
||||||
|
CodeChallengeMethod: codeChallengeMethod,
|
||||||
|
ProviderName: provider.providerName,
|
||||||
|
ExpiresAt: time.Now().Add(10 * time.Minute),
|
||||||
|
Scopes: scopes,
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
s.pending[providerState] = pending
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
authURL, err := provider.auth.OAuth2GetAuthURL(provider.providerName, providerState)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
http.Redirect(w, r, authURL, http.StatusFound)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// External provider callback — GET {ProviderCallbackPath}
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) providerCallbackHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.URL.Query().Get("code")
|
||||||
|
providerState := r.URL.Query().Get("state")
|
||||||
|
|
||||||
|
if code == "" {
|
||||||
|
http.Error(w, "missing code", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
pending, ok := s.pending[providerState]
|
||||||
|
if ok {
|
||||||
|
delete(s.pending, providerState)
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if !ok || time.Now().After(pending.ExpiresAt) {
|
||||||
|
http.Error(w, "invalid or expired state", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
provider := s.providerByName(pending.ProviderName)
|
||||||
|
if provider == nil {
|
||||||
|
http.Error(w, fmt.Sprintf("provider %q not found", pending.ProviderName), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
loginResp, err := provider.auth.OAuth2HandleCallback(r.Context(), pending.ProviderName, code, providerState)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s.issueCodeAndRedirect(w, r, loginResp.Token, loginResp.RefreshToken,
|
||||||
|
pending.ClientID, pending.RedirectURI, pending.ClientState,
|
||||||
|
pending.CodeChallenge, pending.CodeChallengeMethod, pending.ProviderName, pending.Scopes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// issueCodeAndRedirect generates a short-lived auth code and redirects to the MCP client.
|
||||||
|
func (s *OAuthServer) issueCodeAndRedirect(w http.ResponseWriter, r *http.Request, sessionToken, refreshToken, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, providerName string, scopes []string) {
|
||||||
|
authCode, err := randomOAuthToken()
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "server error", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
pending := &pendingAuth{
|
||||||
|
ClientID: clientID,
|
||||||
|
RedirectURI: redirectURI,
|
||||||
|
ClientState: clientState,
|
||||||
|
CodeChallenge: codeChallenge,
|
||||||
|
CodeChallengeMethod: codeChallengeMethod,
|
||||||
|
ProviderName: providerName,
|
||||||
|
SessionToken: sessionToken,
|
||||||
|
RefreshToken: refreshToken,
|
||||||
|
ExpiresAt: time.Now().Add(s.cfg.AuthCodeTTL),
|
||||||
|
Scopes: scopes,
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.cfg.PersistCodes && s.auth != nil {
|
||||||
|
oauthCode := &OAuthCode{
|
||||||
|
Code: authCode,
|
||||||
|
ClientID: clientID,
|
||||||
|
RedirectURI: redirectURI,
|
||||||
|
ClientState: clientState,
|
||||||
|
CodeChallenge: codeChallenge,
|
||||||
|
CodeChallengeMethod: codeChallengeMethod,
|
||||||
|
SessionToken: sessionToken,
|
||||||
|
RefreshToken: refreshToken,
|
||||||
|
Scopes: scopes,
|
||||||
|
ExpiresAt: pending.ExpiresAt,
|
||||||
|
}
|
||||||
|
if err := s.auth.OAuthSaveCode(r.Context(), oauthCode); err != nil {
|
||||||
|
http.Error(w, "server error", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.codes[authCode] = pending
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
redirectURL, err := url.Parse(redirectURI)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "invalid redirect_uri", http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
qp := redirectURL.Query()
|
||||||
|
qp.Set("code", authCode)
|
||||||
|
if clientState != "" {
|
||||||
|
qp.Set("state", clientState)
|
||||||
|
}
|
||||||
|
redirectURL.RawQuery = qp.Encode()
|
||||||
|
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Token endpoint — POST /oauth/token
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) tokenHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := r.ParseForm(); err != nil {
|
||||||
|
writeOAuthError(w, "invalid_request", "cannot parse form", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
switch r.FormValue("grant_type") {
|
||||||
|
case "authorization_code":
|
||||||
|
s.handleAuthCodeGrant(w, r)
|
||||||
|
case "refresh_token":
|
||||||
|
s.handleRefreshGrant(w, r)
|
||||||
|
default:
|
||||||
|
writeOAuthError(w, "unsupported_grant_type", "", http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *OAuthServer) handleAuthCodeGrant(w http.ResponseWriter, r *http.Request) {
|
||||||
|
code := r.FormValue("code")
|
||||||
|
redirectURI := r.FormValue("redirect_uri")
|
||||||
|
clientID := r.FormValue("client_id")
|
||||||
|
codeVerifier := r.FormValue("code_verifier")
|
||||||
|
|
||||||
|
if code == "" || codeVerifier == "" {
|
||||||
|
writeOAuthError(w, "invalid_request", "code and code_verifier required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var sessionToken string
|
||||||
|
var refreshToken string
|
||||||
|
var scopes []string
|
||||||
|
|
||||||
|
if s.cfg.PersistCodes && s.auth != nil {
|
||||||
|
oauthCode, err := s.auth.OAuthExchangeCode(r.Context(), code)
|
||||||
|
if err != nil {
|
||||||
|
writeOAuthError(w, "invalid_grant", "code expired or invalid", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if oauthCode.ClientID != clientID {
|
||||||
|
writeOAuthError(w, "invalid_client", "", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if oauthCode.RedirectURI != redirectURI {
|
||||||
|
writeOAuthError(w, "invalid_grant", "redirect_uri mismatch", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !validatePKCESHA256(oauthCode.CodeChallenge, codeVerifier) {
|
||||||
|
writeOAuthError(w, "invalid_grant", "code_verifier invalid", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sessionToken = oauthCode.SessionToken
|
||||||
|
refreshToken = oauthCode.RefreshToken
|
||||||
|
scopes = oauthCode.Scopes
|
||||||
|
} else {
|
||||||
|
s.mu.Lock()
|
||||||
|
pending, ok := s.codes[code]
|
||||||
|
if ok {
|
||||||
|
delete(s.codes, code)
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if !ok || time.Now().After(pending.ExpiresAt) {
|
||||||
|
writeOAuthError(w, "invalid_grant", "code expired or invalid", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if pending.ClientID != clientID {
|
||||||
|
writeOAuthError(w, "invalid_client", "", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if pending.RedirectURI != redirectURI {
|
||||||
|
writeOAuthError(w, "invalid_grant", "redirect_uri mismatch", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !validatePKCESHA256(pending.CodeChallenge, codeVerifier) {
|
||||||
|
writeOAuthError(w, "invalid_grant", "code_verifier invalid", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sessionToken = pending.SessionToken
|
||||||
|
refreshToken = pending.RefreshToken
|
||||||
|
scopes = pending.Scopes
|
||||||
|
}
|
||||||
|
|
||||||
|
s.writeOAuthToken(w, sessionToken, refreshToken, scopes)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *OAuthServer) handleRefreshGrant(w http.ResponseWriter, r *http.Request) {
|
||||||
|
refreshToken := r.FormValue("refresh_token")
|
||||||
|
providerName := r.FormValue("provider")
|
||||||
|
if refreshToken == "" {
|
||||||
|
writeOAuthError(w, "invalid_request", "refresh_token required", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try external providers first, then fall back to DatabaseAuthenticator
|
||||||
|
provider := s.providerByName(providerName)
|
||||||
|
if provider != nil {
|
||||||
|
loginResp, err := provider.auth.OAuth2RefreshToken(r.Context(), refreshToken, providerName)
|
||||||
|
if err != nil {
|
||||||
|
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.auth != nil {
|
||||||
|
loginResp, err := s.auth.RefreshToken(r.Context(), refreshToken)
|
||||||
|
if err != nil {
|
||||||
|
writeOAuthError(w, "invalid_grant", err.Error(), http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.writeOAuthToken(w, loginResp.Token, loginResp.RefreshToken, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
writeOAuthError(w, "invalid_grant", "no provider available for refresh", http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// RFC 7009 — Token revocation
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) revokeHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := r.ParseForm(); err != nil {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
token := r.FormValue("token")
|
||||||
|
if token == "" {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.auth != nil {
|
||||||
|
s.auth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck
|
||||||
|
} else {
|
||||||
|
// In external-provider-only mode, attempt revocation via the first provider's auth.
|
||||||
|
s.mu.RLock()
|
||||||
|
var providerAuth *DatabaseAuthenticator
|
||||||
|
if len(s.providers) > 0 {
|
||||||
|
providerAuth = s.providers[0].auth
|
||||||
|
}
|
||||||
|
s.mu.RUnlock()
|
||||||
|
if providerAuth != nil {
|
||||||
|
providerAuth.OAuthRevokeToken(r.Context(), token) //nolint:errcheck
|
||||||
|
}
|
||||||
|
}
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// RFC 7662 — Token introspection
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) introspectHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := r.ParseForm(); err != nil {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
|
||||||
|
return
|
||||||
|
}
|
||||||
|
token := r.FormValue("token")
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
if token == "" {
|
||||||
|
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resolve the authenticator to use: prefer the primary auth, then the first provider's auth.
|
||||||
|
authToUse := s.auth
|
||||||
|
if authToUse == nil {
|
||||||
|
s.mu.RLock()
|
||||||
|
if len(s.providers) > 0 {
|
||||||
|
authToUse = s.providers[0].auth
|
||||||
|
}
|
||||||
|
s.mu.RUnlock()
|
||||||
|
}
|
||||||
|
if authToUse == nil {
|
||||||
|
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
info, err := authToUse.OAuthIntrospectToken(r.Context(), token)
|
||||||
|
if err != nil {
|
||||||
|
w.Write([]byte(`{"active":false}`)) //nolint:errcheck
|
||||||
|
return
|
||||||
|
}
|
||||||
|
json.NewEncoder(w).Encode(info) //nolint:errcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Login form (direct identity provider mode)
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
func (s *OAuthServer) renderLoginForm(w http.ResponseWriter, r *http.Request, clientID, redirectURI, clientState, codeChallenge, codeChallengeMethod, scope, errMsg string) {
|
||||||
|
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||||
|
errHTML := ""
|
||||||
|
if errMsg != "" {
|
||||||
|
errHTML = `<p style="color:red">` + errMsg + `</p>`
|
||||||
|
}
|
||||||
|
fmt.Fprintf(w, loginFormHTML,
|
||||||
|
s.cfg.LoginTitle,
|
||||||
|
s.cfg.LoginTitle,
|
||||||
|
errHTML,
|
||||||
|
clientID,
|
||||||
|
htmlEscape(redirectURI),
|
||||||
|
htmlEscape(clientState),
|
||||||
|
htmlEscape(codeChallenge),
|
||||||
|
htmlEscape(codeChallengeMethod),
|
||||||
|
htmlEscape(scope),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
const loginFormHTML = `<!DOCTYPE html>
|
||||||
|
<html><head><meta charset="utf-8"><title>%s</title>
|
||||||
|
<style>body{font-family:sans-serif;display:flex;justify-content:center;align-items:center;min-height:100vh;margin:0;background:#f5f5f5}
|
||||||
|
.card{background:#fff;padding:2rem;border-radius:8px;box-shadow:0 2px 8px rgba(0,0,0,.15);width:320px}
|
||||||
|
h2{margin:0 0 1.5rem;font-size:1.25rem}
|
||||||
|
label{display:block;margin-bottom:.25rem;font-size:.875rem;color:#555}
|
||||||
|
input[type=text],input[type=password]{width:100%%;box-sizing:border-box;padding:.5rem;border:1px solid #ccc;border-radius:4px;margin-bottom:1rem;font-size:1rem}
|
||||||
|
button{width:100%%;padding:.6rem;background:#0070f3;color:#fff;border:none;border-radius:4px;font-size:1rem;cursor:pointer}
|
||||||
|
button:hover{background:#005fd4}.err{color:#d32f2f;margin-bottom:1rem;font-size:.875rem}</style>
|
||||||
|
</head><body><div class="card">
|
||||||
|
<h2>%s</h2>%s
|
||||||
|
<form method="POST" action="authorize">
|
||||||
|
<input type="hidden" name="client_id" value="%s">
|
||||||
|
<input type="hidden" name="redirect_uri" value="%s">
|
||||||
|
<input type="hidden" name="client_state" value="%s">
|
||||||
|
<input type="hidden" name="code_challenge" value="%s">
|
||||||
|
<input type="hidden" name="code_challenge_method" value="%s">
|
||||||
|
<input type="hidden" name="scope" value="%s">
|
||||||
|
<label>Username</label><input type="text" name="username" autofocus autocomplete="username">
|
||||||
|
<label>Password</label><input type="password" name="password" autocomplete="current-password">
|
||||||
|
<button type="submit">Sign in</button>
|
||||||
|
</form></div></body></html>`
|
||||||
|
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
// Helpers
|
||||||
|
// --------------------------------------------------------------------------
|
||||||
|
|
||||||
|
// lookupOrFetchClient checks in-memory first, then DB if PersistClients is enabled.
|
||||||
|
func (s *OAuthServer) lookupOrFetchClient(ctx context.Context, clientID string) (*oauthClient, bool) {
|
||||||
|
s.mu.RLock()
|
||||||
|
c, ok := s.clients[clientID]
|
||||||
|
s.mu.RUnlock()
|
||||||
|
if ok {
|
||||||
|
return c, true
|
||||||
|
}
|
||||||
|
|
||||||
|
if !s.cfg.PersistClients || s.auth == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
dbClient, err := s.auth.OAuthGetClient(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
c = &oauthClient{
|
||||||
|
ClientID: dbClient.ClientID,
|
||||||
|
RedirectURIs: dbClient.RedirectURIs,
|
||||||
|
ClientName: dbClient.ClientName,
|
||||||
|
GrantTypes: dbClient.GrantTypes,
|
||||||
|
AllowedScopes: dbClient.AllowedScopes,
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
s.clients[clientID] = c
|
||||||
|
s.mu.Unlock()
|
||||||
|
return c, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *OAuthServer) providerByName(name string) *externalProvider {
|
||||||
|
for i := range s.providers {
|
||||||
|
if s.providers[i].providerName == name {
|
||||||
|
return &s.providers[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// If name is empty and only one provider exists, return it
|
||||||
|
if name == "" && len(s.providers) == 1 {
|
||||||
|
return &s.providers[0]
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validatePKCESHA256(challenge, verifier string) bool {
|
||||||
|
h := sha256.Sum256([]byte(verifier))
|
||||||
|
return base64.RawURLEncoding.EncodeToString(h[:]) == challenge
|
||||||
|
}
|
||||||
|
|
||||||
|
func randomOAuthToken() (string, error) {
|
||||||
|
b := make([]byte, 32)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return base64.RawURLEncoding.EncodeToString(b), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func oauthSliceContains(slice []string, s string) bool {
|
||||||
|
for _, v := range slice {
|
||||||
|
if v == s {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *OAuthServer) writeOAuthToken(w http.ResponseWriter, accessToken, refreshToken string, scopes []string) {
|
||||||
|
expiresIn := int64(s.cfg.AccessTokenTTL.Seconds())
|
||||||
|
resp := map[string]interface{}{
|
||||||
|
"access_token": accessToken,
|
||||||
|
"token_type": "Bearer",
|
||||||
|
"expires_in": expiresIn,
|
||||||
|
}
|
||||||
|
if refreshToken != "" {
|
||||||
|
resp["refresh_token"] = refreshToken
|
||||||
|
}
|
||||||
|
if len(scopes) > 0 {
|
||||||
|
resp["scope"] = strings.Join(scopes, " ")
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.Header().Set("Cache-Control", "no-store")
|
||||||
|
w.Header().Set("Pragma", "no-cache")
|
||||||
|
json.NewEncoder(w).Encode(resp) //nolint:errcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeOAuthError(w http.ResponseWriter, errCode, description string, status int) {
|
||||||
|
resp := map[string]string{"error": errCode}
|
||||||
|
if description != "" {
|
||||||
|
resp["error_description"] = description
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(status)
|
||||||
|
json.NewEncoder(w).Encode(resp) //nolint:errcheck
|
||||||
|
}
|
||||||
|
|
||||||
|
func htmlEscape(s string) string {
|
||||||
|
s = strings.ReplaceAll(s, "&", "&")
|
||||||
|
s = strings.ReplaceAll(s, `"`, """)
|
||||||
|
s = strings.ReplaceAll(s, "<", "<")
|
||||||
|
s = strings.ReplaceAll(s, ">", ">")
|
||||||
|
return s
|
||||||
|
}
|
||||||
+204
@@ -0,0 +1,204 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// OAuthServerClient is a persisted RFC 7591 registered OAuth2 client.
|
||||||
|
type OAuthServerClient struct {
|
||||||
|
ClientID string `json:"client_id"`
|
||||||
|
RedirectURIs []string `json:"redirect_uris"`
|
||||||
|
ClientName string `json:"client_name,omitempty"`
|
||||||
|
GrantTypes []string `json:"grant_types"`
|
||||||
|
AllowedScopes []string `json:"allowed_scopes,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthCode is a short-lived authorization code.
|
||||||
|
type OAuthCode struct {
|
||||||
|
Code string `json:"code"`
|
||||||
|
ClientID string `json:"client_id"`
|
||||||
|
RedirectURI string `json:"redirect_uri"`
|
||||||
|
ClientState string `json:"client_state,omitempty"`
|
||||||
|
CodeChallenge string `json:"code_challenge"`
|
||||||
|
CodeChallengeMethod string `json:"code_challenge_method"`
|
||||||
|
SessionToken string `json:"session_token"`
|
||||||
|
RefreshToken string `json:"refresh_token,omitempty"`
|
||||||
|
Scopes []string `json:"scopes,omitempty"`
|
||||||
|
ExpiresAt time.Time `json:"expires_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthTokenInfo is the RFC 7662 token introspection response.
|
||||||
|
type OAuthTokenInfo struct {
|
||||||
|
Active bool `json:"active"`
|
||||||
|
Sub string `json:"sub,omitempty"`
|
||||||
|
Username string `json:"username,omitempty"`
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
UserLevel int `json:"user_level,omitempty"`
|
||||||
|
Roles []string `json:"roles,omitempty"`
|
||||||
|
Exp int64 `json:"exp,omitempty"`
|
||||||
|
Iat int64 `json:"iat,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthRegisterClient persists an OAuth2 client registration.
|
||||||
|
func (a *DatabaseAuthenticator) OAuthRegisterClient(ctx context.Context, client *OAuthServerClient) (*OAuthServerClient, error) {
|
||||||
|
input, err := json.Marshal(client)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal client: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
var data []byte
|
||||||
|
|
||||||
|
err = a.db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, a.sqlNames.OAuthRegisterClient), input).Scan(&success, &errMsg, &data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to register client")
|
||||||
|
}
|
||||||
|
|
||||||
|
var result OAuthServerClient
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse registered client: %w", err)
|
||||||
|
}
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthGetClient retrieves a registered client by ID.
|
||||||
|
func (a *DatabaseAuthenticator) OAuthGetClient(ctx context.Context, clientID string) (*OAuthServerClient, error) {
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
var data []byte
|
||||||
|
|
||||||
|
err := a.db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, a.sqlNames.OAuthGetClient), clientID).Scan(&success, &errMsg, &data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("client not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
var result OAuthServerClient
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse client: %w", err)
|
||||||
|
}
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthSaveCode persists an authorization code.
|
||||||
|
func (a *DatabaseAuthenticator) OAuthSaveCode(ctx context.Context, code *OAuthCode) error {
|
||||||
|
input, err := json.Marshal(code)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal code: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
|
||||||
|
err = a.db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1::jsonb)
|
||||||
|
`, a.sqlNames.OAuthSaveCode), input).Scan(&success, &errMsg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to save code: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to save code")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthExchangeCode retrieves and deletes an authorization code (single use).
|
||||||
|
func (a *DatabaseAuthenticator) OAuthExchangeCode(ctx context.Context, code string) (*OAuthCode, error) {
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
var data []byte
|
||||||
|
|
||||||
|
err := a.db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, a.sqlNames.OAuthExchangeCode), code).Scan(&success, &errMsg, &data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("invalid or expired code")
|
||||||
|
}
|
||||||
|
|
||||||
|
var result OAuthCode
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse code data: %w", err)
|
||||||
|
}
|
||||||
|
result.Code = code
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthIntrospectToken validates a token and returns its metadata (RFC 7662).
|
||||||
|
func (a *DatabaseAuthenticator) OAuthIntrospectToken(ctx context.Context, token string) (*OAuthTokenInfo, error) {
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
var data []byte
|
||||||
|
|
||||||
|
err := a.db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error, p_data::text
|
||||||
|
FROM %s($1)
|
||||||
|
`, a.sqlNames.OAuthIntrospect), token).Scan(&success, &errMsg, &data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to introspect token: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return nil, fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("introspection failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
var result OAuthTokenInfo
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse token info: %w", err)
|
||||||
|
}
|
||||||
|
return &result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OAuthRevokeToken revokes a token by deleting the session (RFC 7009).
|
||||||
|
func (a *DatabaseAuthenticator) OAuthRevokeToken(ctx context.Context, token string) error {
|
||||||
|
var success bool
|
||||||
|
var errMsg *string
|
||||||
|
|
||||||
|
err := a.db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||||
|
SELECT p_success, p_error
|
||||||
|
FROM %s($1)
|
||||||
|
`, a.sqlNames.OAuthRevoke), token).Scan(&success, &errMsg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to revoke token: %w", err)
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg != nil {
|
||||||
|
return fmt.Errorf("%s", *errMsg)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to revoke token")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+185
@@ -0,0 +1,185 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PasskeyCredential represents a stored WebAuthn/FIDO2 credential
|
||||||
|
type PasskeyCredential struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
CredentialID []byte `json:"credential_id"` // Raw credential ID from authenticator
|
||||||
|
PublicKey []byte `json:"public_key"` // COSE public key
|
||||||
|
AttestationType string `json:"attestation_type"` // none, indirect, direct
|
||||||
|
AAGUID []byte `json:"aaguid"` // Authenticator AAGUID
|
||||||
|
SignCount uint32 `json:"sign_count"` // Signature counter
|
||||||
|
CloneWarning bool `json:"clone_warning"` // True if cloning detected
|
||||||
|
Transports []string `json:"transports,omitempty"` // usb, nfc, ble, internal
|
||||||
|
BackupEligible bool `json:"backup_eligible"` // Credential can be backed up
|
||||||
|
BackupState bool `json:"backup_state"` // Credential is currently backed up
|
||||||
|
Name string `json:"name,omitempty"` // User-friendly name
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
LastUsedAt time.Time `json:"last_used_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyRegistrationOptions contains options for beginning passkey registration
|
||||||
|
type PasskeyRegistrationOptions struct {
|
||||||
|
Challenge []byte `json:"challenge"`
|
||||||
|
RelyingParty PasskeyRelyingParty `json:"rp"`
|
||||||
|
User PasskeyUser `json:"user"`
|
||||||
|
PubKeyCredParams []PasskeyCredentialParam `json:"pubKeyCredParams"`
|
||||||
|
Timeout int64 `json:"timeout,omitempty"` // Milliseconds
|
||||||
|
ExcludeCredentials []PasskeyCredentialDescriptor `json:"excludeCredentials,omitempty"`
|
||||||
|
AuthenticatorSelection *PasskeyAuthenticatorSelection `json:"authenticatorSelection,omitempty"`
|
||||||
|
Attestation string `json:"attestation,omitempty"` // none, indirect, direct, enterprise
|
||||||
|
Extensions map[string]any `json:"extensions,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyAuthenticationOptions contains options for beginning passkey authentication
|
||||||
|
type PasskeyAuthenticationOptions struct {
|
||||||
|
Challenge []byte `json:"challenge"`
|
||||||
|
Timeout int64 `json:"timeout,omitempty"`
|
||||||
|
RelyingPartyID string `json:"rpId,omitempty"`
|
||||||
|
AllowCredentials []PasskeyCredentialDescriptor `json:"allowCredentials,omitempty"`
|
||||||
|
UserVerification string `json:"userVerification,omitempty"` // required, preferred, discouraged
|
||||||
|
Extensions map[string]any `json:"extensions,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyRelyingParty identifies the relying party
|
||||||
|
type PasskeyRelyingParty struct {
|
||||||
|
ID string `json:"id"` // Domain (e.g., "example.com")
|
||||||
|
Name string `json:"name"` // Display name
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyUser identifies the user
|
||||||
|
type PasskeyUser struct {
|
||||||
|
ID []byte `json:"id"` // User handle (unique, persistent)
|
||||||
|
Name string `json:"name"` // Username
|
||||||
|
DisplayName string `json:"displayName"` // Display name
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyCredentialParam specifies supported public key algorithm
|
||||||
|
type PasskeyCredentialParam struct {
|
||||||
|
Type string `json:"type"` // "public-key"
|
||||||
|
Alg int `json:"alg"` // COSE algorithm identifier (e.g., -7 for ES256, -257 for RS256)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyCredentialDescriptor describes a credential
|
||||||
|
type PasskeyCredentialDescriptor struct {
|
||||||
|
Type string `json:"type"` // "public-key"
|
||||||
|
ID []byte `json:"id"` // Credential ID
|
||||||
|
Transports []string `json:"transports,omitempty"` // usb, nfc, ble, internal
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyAuthenticatorSelection specifies authenticator requirements
|
||||||
|
type PasskeyAuthenticatorSelection struct {
|
||||||
|
AuthenticatorAttachment string `json:"authenticatorAttachment,omitempty"` // platform, cross-platform
|
||||||
|
RequireResidentKey bool `json:"requireResidentKey,omitempty"`
|
||||||
|
ResidentKey string `json:"residentKey,omitempty"` // discouraged, preferred, required
|
||||||
|
UserVerification string `json:"userVerification,omitempty"` // required, preferred, discouraged
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyRegistrationResponse contains the client's registration response
|
||||||
|
type PasskeyRegistrationResponse struct {
|
||||||
|
ID string `json:"id"` // Base64URL encoded credential ID
|
||||||
|
RawID []byte `json:"rawId"` // Raw credential ID
|
||||||
|
Type string `json:"type"` // "public-key"
|
||||||
|
Response PasskeyAuthenticatorAttestationResponse `json:"response"`
|
||||||
|
ClientExtensionResults map[string]any `json:"clientExtensionResults,omitempty"`
|
||||||
|
Transports []string `json:"transports,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyAuthenticatorAttestationResponse contains attestation data
|
||||||
|
type PasskeyAuthenticatorAttestationResponse struct {
|
||||||
|
ClientDataJSON []byte `json:"clientDataJSON"`
|
||||||
|
AttestationObject []byte `json:"attestationObject"`
|
||||||
|
Transports []string `json:"transports,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyAuthenticationResponse contains the client's authentication response
|
||||||
|
type PasskeyAuthenticationResponse struct {
|
||||||
|
ID string `json:"id"` // Base64URL encoded credential ID
|
||||||
|
RawID []byte `json:"rawId"` // Raw credential ID
|
||||||
|
Type string `json:"type"` // "public-key"
|
||||||
|
Response PasskeyAuthenticatorAssertionResponse `json:"response"`
|
||||||
|
ClientExtensionResults map[string]any `json:"clientExtensionResults,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyAuthenticatorAssertionResponse contains assertion data
|
||||||
|
type PasskeyAuthenticatorAssertionResponse struct {
|
||||||
|
ClientDataJSON []byte `json:"clientDataJSON"`
|
||||||
|
AuthenticatorData []byte `json:"authenticatorData"`
|
||||||
|
Signature []byte `json:"signature"`
|
||||||
|
UserHandle []byte `json:"userHandle,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyProvider handles passkey registration and authentication
|
||||||
|
type PasskeyProvider interface {
|
||||||
|
// BeginRegistration creates registration options for a new passkey
|
||||||
|
BeginRegistration(ctx context.Context, userID int, username, displayName string) (*PasskeyRegistrationOptions, error)
|
||||||
|
|
||||||
|
// CompleteRegistration verifies and stores a new passkey credential
|
||||||
|
CompleteRegistration(ctx context.Context, userID int, response PasskeyRegistrationResponse, expectedChallenge []byte) (*PasskeyCredential, error)
|
||||||
|
|
||||||
|
// BeginAuthentication creates authentication options for passkey login
|
||||||
|
BeginAuthentication(ctx context.Context, username string) (*PasskeyAuthenticationOptions, error)
|
||||||
|
|
||||||
|
// CompleteAuthentication verifies a passkey assertion and returns the user
|
||||||
|
CompleteAuthentication(ctx context.Context, response PasskeyAuthenticationResponse, expectedChallenge []byte) (int, error)
|
||||||
|
|
||||||
|
// GetCredentials returns all passkey credentials for a user
|
||||||
|
GetCredentials(ctx context.Context, userID int) ([]PasskeyCredential, error)
|
||||||
|
|
||||||
|
// DeleteCredential removes a passkey credential
|
||||||
|
DeleteCredential(ctx context.Context, userID int, credentialID string) error
|
||||||
|
|
||||||
|
// UpdateCredentialName updates the friendly name of a credential
|
||||||
|
UpdateCredentialName(ctx context.Context, userID int, credentialID string, name string) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyLoginRequest contains passkey authentication data
|
||||||
|
type PasskeyLoginRequest struct {
|
||||||
|
Response PasskeyAuthenticationResponse `json:"response"`
|
||||||
|
ExpectedChallenge []byte `json:"expected_challenge"`
|
||||||
|
Claims map[string]any `json:"claims"` // Additional login data
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyRegisterRequest contains passkey registration data
|
||||||
|
type PasskeyRegisterRequest struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
Response PasskeyRegistrationResponse `json:"response"`
|
||||||
|
ExpectedChallenge []byte `json:"expected_challenge"`
|
||||||
|
CredentialName string `json:"credential_name,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyBeginRegistrationRequest contains options for starting passkey registration
|
||||||
|
type PasskeyBeginRegistrationRequest struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
Username string `json:"username"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyBeginAuthenticationRequest contains options for starting passkey authentication
|
||||||
|
type PasskeyBeginAuthenticationRequest struct {
|
||||||
|
Username string `json:"username,omitempty"` // Optional for resident key flow
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParsePasskeyRegistrationResponse parses a JSON passkey registration response
|
||||||
|
func ParsePasskeyRegistrationResponse(data []byte) (*PasskeyRegistrationResponse, error) {
|
||||||
|
var response PasskeyRegistrationResponse
|
||||||
|
if err := json.Unmarshal(data, &response); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParsePasskeyAuthenticationResponse parses a JSON passkey authentication response
|
||||||
|
func ParsePasskeyAuthenticationResponse(data []byte) (*PasskeyAuthenticationResponse, error) {
|
||||||
|
var response PasskeyAuthenticationResponse
|
||||||
|
if err := json.Unmarshal(data, &response); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &response, nil
|
||||||
|
}
|
||||||
+432
@@ -0,0 +1,432 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PasskeyAuthenticationExample demonstrates passkey (WebAuthn/FIDO2) authentication
|
||||||
|
func PasskeyAuthenticationExample() {
|
||||||
|
// Setup database connection
|
||||||
|
db, _ := sql.Open("postgres", "postgres://user:pass@localhost/db")
|
||||||
|
|
||||||
|
// Create passkey provider
|
||||||
|
passkeyProvider := NewDatabasePasskeyProvider(db, DatabasePasskeyProviderOptions{
|
||||||
|
RPID: "example.com", // Your domain
|
||||||
|
RPName: "Example Application", // Display name
|
||||||
|
RPOrigin: "https://example.com", // Expected origin
|
||||||
|
Timeout: 60000, // 60 seconds
|
||||||
|
})
|
||||||
|
|
||||||
|
// Create authenticator with passkey support
|
||||||
|
// Option 1: Pass during creation
|
||||||
|
_ = NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{
|
||||||
|
PasskeyProvider: passkeyProvider,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Option 2: Use WithPasskey method
|
||||||
|
auth := NewDatabaseAuthenticator(db).WithPasskey(passkeyProvider)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// === REGISTRATION FLOW ===
|
||||||
|
|
||||||
|
// Step 1: Begin registration
|
||||||
|
regOptions, _ := auth.BeginPasskeyRegistration(ctx, PasskeyBeginRegistrationRequest{
|
||||||
|
UserID: 1,
|
||||||
|
Username: "alice",
|
||||||
|
DisplayName: "Alice Smith",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Send regOptions to client as JSON
|
||||||
|
// Client will call navigator.credentials.create() with these options
|
||||||
|
_ = regOptions
|
||||||
|
|
||||||
|
// Step 2: Complete registration (after client returns credential)
|
||||||
|
// This would come from the client's navigator.credentials.create() response
|
||||||
|
clientResponse := PasskeyRegistrationResponse{
|
||||||
|
ID: "base64-credential-id",
|
||||||
|
RawID: []byte("raw-credential-id"),
|
||||||
|
Type: "public-key",
|
||||||
|
Response: PasskeyAuthenticatorAttestationResponse{
|
||||||
|
ClientDataJSON: []byte("..."),
|
||||||
|
AttestationObject: []byte("..."),
|
||||||
|
},
|
||||||
|
Transports: []string{"internal"},
|
||||||
|
}
|
||||||
|
|
||||||
|
credential, _ := auth.CompletePasskeyRegistration(ctx, PasskeyRegisterRequest{
|
||||||
|
UserID: 1,
|
||||||
|
Response: clientResponse,
|
||||||
|
ExpectedChallenge: regOptions.Challenge,
|
||||||
|
CredentialName: "My iPhone",
|
||||||
|
})
|
||||||
|
|
||||||
|
fmt.Printf("Registered credential: %s\n", credential.ID)
|
||||||
|
|
||||||
|
// === AUTHENTICATION FLOW ===
|
||||||
|
|
||||||
|
// Step 1: Begin authentication
|
||||||
|
authOptions, _ := auth.BeginPasskeyAuthentication(ctx, PasskeyBeginAuthenticationRequest{
|
||||||
|
Username: "alice", // Optional - omit for resident key flow
|
||||||
|
})
|
||||||
|
|
||||||
|
// Send authOptions to client as JSON
|
||||||
|
// Client will call navigator.credentials.get() with these options
|
||||||
|
_ = authOptions
|
||||||
|
|
||||||
|
// Step 2: Complete authentication (after client returns assertion)
|
||||||
|
// This would come from the client's navigator.credentials.get() response
|
||||||
|
clientAssertion := PasskeyAuthenticationResponse{
|
||||||
|
ID: "base64-credential-id",
|
||||||
|
RawID: []byte("raw-credential-id"),
|
||||||
|
Type: "public-key",
|
||||||
|
Response: PasskeyAuthenticatorAssertionResponse{
|
||||||
|
ClientDataJSON: []byte("..."),
|
||||||
|
AuthenticatorData: []byte("..."),
|
||||||
|
Signature: []byte("..."),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
loginResponse, _ := auth.LoginWithPasskey(ctx, PasskeyLoginRequest{
|
||||||
|
Response: clientAssertion,
|
||||||
|
ExpectedChallenge: authOptions.Challenge,
|
||||||
|
Claims: map[string]any{
|
||||||
|
"ip_address": "192.168.1.1",
|
||||||
|
"user_agent": "Mozilla/5.0...",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
fmt.Printf("Logged in user: %s with token: %s\n",
|
||||||
|
loginResponse.User.UserName, loginResponse.Token)
|
||||||
|
|
||||||
|
// === CREDENTIAL MANAGEMENT ===
|
||||||
|
|
||||||
|
// Get all credentials for a user
|
||||||
|
credentials, _ := auth.GetPasskeyCredentials(ctx, 1)
|
||||||
|
for i := range credentials {
|
||||||
|
fmt.Printf("Credential: %s (created: %s, last used: %s)\n",
|
||||||
|
credentials[i].Name, credentials[i].CreatedAt, credentials[i].LastUsedAt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update credential name
|
||||||
|
_ = auth.UpdatePasskeyCredentialName(ctx, 1, credential.ID, "My New iPhone")
|
||||||
|
|
||||||
|
// Delete credential
|
||||||
|
_ = auth.DeletePasskeyCredential(ctx, 1, credential.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyHTTPHandlersExample shows HTTP handlers for passkey authentication
|
||||||
|
func PasskeyHTTPHandlersExample(auth *DatabaseAuthenticator) {
|
||||||
|
// Store challenges in session/cache in production
|
||||||
|
challenges := make(map[string][]byte)
|
||||||
|
|
||||||
|
// Begin registration endpoint
|
||||||
|
http.HandleFunc("/api/passkey/register/begin", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
Username string `json:"username"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
}
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
options, err := auth.BeginPasskeyRegistration(r.Context(), PasskeyBeginRegistrationRequest{
|
||||||
|
UserID: req.UserID,
|
||||||
|
Username: req.Username,
|
||||||
|
DisplayName: req.DisplayName,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store challenge for verification (use session ID as key in production)
|
||||||
|
sessionID := "session-123"
|
||||||
|
challenges[sessionID] = options.Challenge
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(options)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Complete registration endpoint
|
||||||
|
http.HandleFunc("/api/passkey/register/complete", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
Response PasskeyRegistrationResponse `json:"response"`
|
||||||
|
CredentialName string `json:"credential_name"`
|
||||||
|
}
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Get stored challenge (from session in production)
|
||||||
|
sessionID := "session-123"
|
||||||
|
challenge := challenges[sessionID]
|
||||||
|
delete(challenges, sessionID)
|
||||||
|
|
||||||
|
credential, err := auth.CompletePasskeyRegistration(r.Context(), PasskeyRegisterRequest{
|
||||||
|
UserID: req.UserID,
|
||||||
|
Response: req.Response,
|
||||||
|
ExpectedChallenge: challenge,
|
||||||
|
CredentialName: req.CredentialName,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(credential)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Begin authentication endpoint
|
||||||
|
http.HandleFunc("/api/passkey/login/begin", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
Username string `json:"username"` // Optional
|
||||||
|
}
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
options, err := auth.BeginPasskeyAuthentication(r.Context(), PasskeyBeginAuthenticationRequest{
|
||||||
|
Username: req.Username,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store challenge for verification (use session ID as key in production)
|
||||||
|
sessionID := "session-456"
|
||||||
|
challenges[sessionID] = options.Challenge
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(options)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Complete authentication endpoint
|
||||||
|
http.HandleFunc("/api/passkey/login/complete", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
var req struct {
|
||||||
|
Response PasskeyAuthenticationResponse `json:"response"`
|
||||||
|
}
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
// Get stored challenge (from session in production)
|
||||||
|
sessionID := "session-456"
|
||||||
|
challenge := challenges[sessionID]
|
||||||
|
delete(challenges, sessionID)
|
||||||
|
|
||||||
|
loginResponse, err := auth.LoginWithPasskey(r.Context(), PasskeyLoginRequest{
|
||||||
|
Response: req.Response,
|
||||||
|
ExpectedChallenge: challenge,
|
||||||
|
Claims: map[string]any{
|
||||||
|
"ip_address": r.RemoteAddr,
|
||||||
|
"user_agent": r.UserAgent(),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set session cookie
|
||||||
|
http.SetCookie(w, &http.Cookie{
|
||||||
|
Name: "session_token",
|
||||||
|
Value: loginResponse.Token,
|
||||||
|
Path: "/",
|
||||||
|
HttpOnly: true,
|
||||||
|
Secure: true,
|
||||||
|
SameSite: http.SameSiteLaxMode,
|
||||||
|
})
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(loginResponse)
|
||||||
|
})
|
||||||
|
|
||||||
|
// List credentials endpoint
|
||||||
|
http.HandleFunc("/api/passkey/credentials", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Get user from authenticated session
|
||||||
|
userCtx, err := auth.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Unauthorized", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
credentials, err := auth.GetPasskeyCredentials(r.Context(), userCtx.UserID)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(credentials)
|
||||||
|
})
|
||||||
|
|
||||||
|
// Delete credential endpoint
|
||||||
|
http.HandleFunc("/api/passkey/credentials/delete", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
userCtx, err := auth.Authenticate(r)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, "Unauthorized", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var req struct {
|
||||||
|
CredentialID string `json:"credential_id"`
|
||||||
|
}
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&req)
|
||||||
|
|
||||||
|
err = auth.DeletePasskeyCredential(r.Context(), userCtx.UserID, req.CredentialID)
|
||||||
|
if err != nil {
|
||||||
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// PasskeyClientSideExample shows the client-side JavaScript code needed
|
||||||
|
func PasskeyClientSideExample() string {
|
||||||
|
return `
|
||||||
|
// === CLIENT-SIDE JAVASCRIPT FOR PASSKEY AUTHENTICATION ===
|
||||||
|
|
||||||
|
// Helper function to convert base64 to ArrayBuffer
|
||||||
|
function base64ToArrayBuffer(base64) {
|
||||||
|
const binary = atob(base64);
|
||||||
|
const bytes = new Uint8Array(binary.length);
|
||||||
|
for (let i = 0; i < binary.length; i++) {
|
||||||
|
bytes[i] = binary.charCodeAt(i);
|
||||||
|
}
|
||||||
|
return bytes.buffer;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper function to convert ArrayBuffer to base64
|
||||||
|
function arrayBufferToBase64(buffer) {
|
||||||
|
const bytes = new Uint8Array(buffer);
|
||||||
|
let binary = '';
|
||||||
|
for (let i = 0; i < bytes.length; i++) {
|
||||||
|
binary += String.fromCharCode(bytes[i]);
|
||||||
|
}
|
||||||
|
return btoa(binary);
|
||||||
|
}
|
||||||
|
|
||||||
|
// === REGISTRATION ===
|
||||||
|
|
||||||
|
async function registerPasskey(userId, username, displayName) {
|
||||||
|
// Step 1: Get registration options from server
|
||||||
|
const optionsResponse = await fetch('/api/passkey/register/begin', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({ user_id: userId, username, display_name: displayName })
|
||||||
|
});
|
||||||
|
const options = await optionsResponse.json();
|
||||||
|
|
||||||
|
// Convert base64 strings to ArrayBuffers
|
||||||
|
options.challenge = base64ToArrayBuffer(options.challenge);
|
||||||
|
options.user.id = base64ToArrayBuffer(options.user.id);
|
||||||
|
if (options.excludeCredentials) {
|
||||||
|
options.excludeCredentials = options.excludeCredentials.map(cred => ({
|
||||||
|
...cred,
|
||||||
|
id: base64ToArrayBuffer(cred.id)
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 2: Create credential using WebAuthn API
|
||||||
|
const credential = await navigator.credentials.create({
|
||||||
|
publicKey: options
|
||||||
|
});
|
||||||
|
|
||||||
|
// Step 3: Send credential to server
|
||||||
|
const credentialResponse = {
|
||||||
|
id: credential.id,
|
||||||
|
rawId: arrayBufferToBase64(credential.rawId),
|
||||||
|
type: credential.type,
|
||||||
|
response: {
|
||||||
|
clientDataJSON: arrayBufferToBase64(credential.response.clientDataJSON),
|
||||||
|
attestationObject: arrayBufferToBase64(credential.response.attestationObject)
|
||||||
|
},
|
||||||
|
transports: credential.response.getTransports ? credential.response.getTransports() : []
|
||||||
|
};
|
||||||
|
|
||||||
|
const completeResponse = await fetch('/api/passkey/register/complete', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({
|
||||||
|
user_id: userId,
|
||||||
|
response: credentialResponse,
|
||||||
|
credential_name: 'My Device'
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
return await completeResponse.json();
|
||||||
|
}
|
||||||
|
|
||||||
|
// === AUTHENTICATION ===
|
||||||
|
|
||||||
|
async function loginWithPasskey(username) {
|
||||||
|
// Step 1: Get authentication options from server
|
||||||
|
const optionsResponse = await fetch('/api/passkey/login/begin', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({ username })
|
||||||
|
});
|
||||||
|
const options = await optionsResponse.json();
|
||||||
|
|
||||||
|
// Convert base64 strings to ArrayBuffers
|
||||||
|
options.challenge = base64ToArrayBuffer(options.challenge);
|
||||||
|
if (options.allowCredentials) {
|
||||||
|
options.allowCredentials = options.allowCredentials.map(cred => ({
|
||||||
|
...cred,
|
||||||
|
id: base64ToArrayBuffer(cred.id)
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 2: Get credential using WebAuthn API
|
||||||
|
const credential = await navigator.credentials.get({
|
||||||
|
publicKey: options
|
||||||
|
});
|
||||||
|
|
||||||
|
// Step 3: Send assertion to server
|
||||||
|
const assertionResponse = {
|
||||||
|
id: credential.id,
|
||||||
|
rawId: arrayBufferToBase64(credential.rawId),
|
||||||
|
type: credential.type,
|
||||||
|
response: {
|
||||||
|
clientDataJSON: arrayBufferToBase64(credential.response.clientDataJSON),
|
||||||
|
authenticatorData: arrayBufferToBase64(credential.response.authenticatorData),
|
||||||
|
signature: arrayBufferToBase64(credential.response.signature),
|
||||||
|
userHandle: credential.response.userHandle ? arrayBufferToBase64(credential.response.userHandle) : null
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const loginResponse = await fetch('/api/passkey/login/complete', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({ response: assertionResponse })
|
||||||
|
});
|
||||||
|
|
||||||
|
return await loginResponse.json();
|
||||||
|
}
|
||||||
|
|
||||||
|
// === USAGE ===
|
||||||
|
|
||||||
|
// Register a new passkey
|
||||||
|
document.getElementById('register-btn').addEventListener('click', async () => {
|
||||||
|
try {
|
||||||
|
const result = await registerPasskey(1, 'alice', 'Alice Smith');
|
||||||
|
console.log('Passkey registered:', result);
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Registration failed:', error);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// Login with passkey
|
||||||
|
document.getElementById('login-btn').addEventListener('click', async () => {
|
||||||
|
try {
|
||||||
|
const result = await loginWithPasskey('alice');
|
||||||
|
console.log('Logged in:', result);
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Login failed:', error);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
`
|
||||||
|
}
|
||||||
+447
@@ -0,0 +1,447 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DatabasePasskeyProvider implements PasskeyProvider using database storage
|
||||||
|
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
||||||
|
type DatabasePasskeyProvider struct {
|
||||||
|
db *sql.DB
|
||||||
|
dbMu sync.RWMutex
|
||||||
|
dbFactory func() (*sql.DB, error)
|
||||||
|
rpID string // Relying Party ID (domain)
|
||||||
|
rpName string // Relying Party display name
|
||||||
|
rpOrigin string // Expected origin for WebAuthn
|
||||||
|
timeout int64 // Timeout in milliseconds (default: 60000)
|
||||||
|
sqlNames *SQLNames
|
||||||
|
}
|
||||||
|
|
||||||
|
// DatabasePasskeyProviderOptions configures the passkey provider
|
||||||
|
type DatabasePasskeyProviderOptions struct {
|
||||||
|
// RPID is the Relying Party ID (typically your domain, e.g., "example.com")
|
||||||
|
RPID string
|
||||||
|
// RPName is the display name for your relying party
|
||||||
|
RPName string
|
||||||
|
// RPOrigin is the expected origin (e.g., "https://example.com")
|
||||||
|
RPOrigin string
|
||||||
|
// Timeout is the timeout for operations in milliseconds (default: 60000)
|
||||||
|
Timeout int64
|
||||||
|
// SQLNames provides custom SQL procedure/function names. If nil, uses DefaultSQLNames().
|
||||||
|
SQLNames *SQLNames
|
||||||
|
// DBFactory is called to obtain a fresh *sql.DB when the existing connection is closed.
|
||||||
|
// If nil, reconnection is disabled.
|
||||||
|
DBFactory func() (*sql.DB, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDatabasePasskeyProvider creates a new database-backed passkey provider
|
||||||
|
func NewDatabasePasskeyProvider(db *sql.DB, opts DatabasePasskeyProviderOptions) *DatabasePasskeyProvider {
|
||||||
|
if opts.Timeout == 0 {
|
||||||
|
opts.Timeout = 60000 // 60 seconds default
|
||||||
|
}
|
||||||
|
|
||||||
|
sqlNames := MergeSQLNames(DefaultSQLNames(), opts.SQLNames)
|
||||||
|
|
||||||
|
return &DatabasePasskeyProvider{
|
||||||
|
db: db,
|
||||||
|
dbFactory: opts.DBFactory,
|
||||||
|
rpID: opts.RPID,
|
||||||
|
rpName: opts.RPName,
|
||||||
|
rpOrigin: opts.RPOrigin,
|
||||||
|
timeout: opts.Timeout,
|
||||||
|
sqlNames: sqlNames,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *DatabasePasskeyProvider) getDB() *sql.DB {
|
||||||
|
p.dbMu.RLock()
|
||||||
|
defer p.dbMu.RUnlock()
|
||||||
|
return p.db
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *DatabasePasskeyProvider) reconnectDB() error {
|
||||||
|
if p.dbFactory == nil {
|
||||||
|
return fmt.Errorf("no db factory configured for reconnect")
|
||||||
|
}
|
||||||
|
newDB, err := p.dbFactory()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
p.dbMu.Lock()
|
||||||
|
p.db = newDB
|
||||||
|
p.dbMu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeginRegistration creates registration options for a new passkey
|
||||||
|
func (p *DatabasePasskeyProvider) BeginRegistration(ctx context.Context, userID int, username, displayName string) (*PasskeyRegistrationOptions, error) {
|
||||||
|
// Generate challenge
|
||||||
|
challenge := make([]byte, 32)
|
||||||
|
if _, err := rand.Read(challenge); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate challenge: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get existing credentials to exclude
|
||||||
|
credentials, err := p.GetCredentials(ctx, userID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get existing credentials: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
excludeCredentials := make([]PasskeyCredentialDescriptor, 0, len(credentials))
|
||||||
|
for i := range credentials {
|
||||||
|
excludeCredentials = append(excludeCredentials, PasskeyCredentialDescriptor{
|
||||||
|
Type: "public-key",
|
||||||
|
ID: credentials[i].CredentialID,
|
||||||
|
Transports: credentials[i].Transports,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create user handle (persistent user ID)
|
||||||
|
userHandle := []byte(fmt.Sprintf("user_%d", userID))
|
||||||
|
|
||||||
|
return &PasskeyRegistrationOptions{
|
||||||
|
Challenge: challenge,
|
||||||
|
RelyingParty: PasskeyRelyingParty{
|
||||||
|
ID: p.rpID,
|
||||||
|
Name: p.rpName,
|
||||||
|
},
|
||||||
|
User: PasskeyUser{
|
||||||
|
ID: userHandle,
|
||||||
|
Name: username,
|
||||||
|
DisplayName: displayName,
|
||||||
|
},
|
||||||
|
PubKeyCredParams: []PasskeyCredentialParam{
|
||||||
|
{Type: "public-key", Alg: -7}, // ES256 (ECDSA with SHA-256)
|
||||||
|
{Type: "public-key", Alg: -257}, // RS256 (RSASSA-PKCS1-v1_5 with SHA-256)
|
||||||
|
},
|
||||||
|
Timeout: p.timeout,
|
||||||
|
ExcludeCredentials: excludeCredentials,
|
||||||
|
AuthenticatorSelection: &PasskeyAuthenticatorSelection{
|
||||||
|
RequireResidentKey: false,
|
||||||
|
ResidentKey: "preferred",
|
||||||
|
UserVerification: "preferred",
|
||||||
|
},
|
||||||
|
Attestation: "none",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompleteRegistration verifies and stores a new passkey credential
|
||||||
|
// NOTE: This is a simplified implementation. In production, you should use a WebAuthn library
|
||||||
|
// like github.com/go-webauthn/webauthn to properly verify attestation and parse credentials.
|
||||||
|
func (p *DatabasePasskeyProvider) CompleteRegistration(ctx context.Context, userID int, response PasskeyRegistrationResponse, expectedChallenge []byte) (*PasskeyCredential, error) {
|
||||||
|
// TODO: Implement full WebAuthn verification
|
||||||
|
// 1. Verify clientDataJSON contains correct challenge and origin
|
||||||
|
// 2. Parse and verify attestationObject
|
||||||
|
// 3. Extract public key and credential ID
|
||||||
|
// 4. Verify attestation signature (if not "none")
|
||||||
|
|
||||||
|
// For now, this is a placeholder that stores the credential data
|
||||||
|
// In production, you MUST use a proper WebAuthn library
|
||||||
|
|
||||||
|
credData := map[string]any{
|
||||||
|
"user_id": userID,
|
||||||
|
"credential_id": base64.StdEncoding.EncodeToString(response.RawID),
|
||||||
|
"public_key": base64.StdEncoding.EncodeToString(response.Response.AttestationObject),
|
||||||
|
"attestation_type": "none",
|
||||||
|
"sign_count": 0,
|
||||||
|
"transports": response.Transports,
|
||||||
|
"backup_eligible": false,
|
||||||
|
"backup_state": false,
|
||||||
|
"name": "Passkey",
|
||||||
|
}
|
||||||
|
|
||||||
|
credJSON, err := json.Marshal(credData)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal credential data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var credentialID sql.NullInt64
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential_id FROM %s($1::jsonb)`, p.sqlNames.PasskeyStoreCredential)
|
||||||
|
err = p.getDB().QueryRowContext(ctx, query, string(credJSON)).Scan(&success, &errorMsg, &credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to store credential: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to store credential")
|
||||||
|
}
|
||||||
|
|
||||||
|
return &PasskeyCredential{
|
||||||
|
ID: fmt.Sprintf("%d", credentialID.Int64),
|
||||||
|
UserID: userID,
|
||||||
|
CredentialID: response.RawID,
|
||||||
|
PublicKey: response.Response.AttestationObject,
|
||||||
|
AttestationType: "none",
|
||||||
|
Transports: response.Transports,
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
LastUsedAt: time.Now(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeginAuthentication creates authentication options for passkey login
|
||||||
|
func (p *DatabasePasskeyProvider) BeginAuthentication(ctx context.Context, username string) (*PasskeyAuthenticationOptions, error) {
|
||||||
|
// Generate challenge
|
||||||
|
challenge := make([]byte, 32)
|
||||||
|
if _, err := rand.Read(challenge); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate challenge: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If username is provided, get user's credentials
|
||||||
|
var allowCredentials []PasskeyCredentialDescriptor
|
||||||
|
if username != "" {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var userID sql.NullInt64
|
||||||
|
var credentialsJSON sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.sqlNames.PasskeyGetCredsByUsername)
|
||||||
|
err := p.getDB().QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to get credentials")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse credentials
|
||||||
|
var creds []struct {
|
||||||
|
ID string `json:"credential_id"`
|
||||||
|
Transports []string `json:"transports"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse credentials: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
allowCredentials = make([]PasskeyCredentialDescriptor, 0, len(creds))
|
||||||
|
for _, cred := range creds {
|
||||||
|
credID, err := base64.StdEncoding.DecodeString(cred.ID)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
allowCredentials = append(allowCredentials, PasskeyCredentialDescriptor{
|
||||||
|
Type: "public-key",
|
||||||
|
ID: credID,
|
||||||
|
Transports: cred.Transports,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &PasskeyAuthenticationOptions{
|
||||||
|
Challenge: challenge,
|
||||||
|
Timeout: p.timeout,
|
||||||
|
RelyingPartyID: p.rpID,
|
||||||
|
AllowCredentials: allowCredentials,
|
||||||
|
UserVerification: "preferred",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompleteAuthentication verifies a passkey assertion and returns the user ID
|
||||||
|
// NOTE: This is a simplified implementation. In production, you should use a WebAuthn library
|
||||||
|
// like github.com/go-webauthn/webauthn to properly verify the assertion signature.
|
||||||
|
func (p *DatabasePasskeyProvider) CompleteAuthentication(ctx context.Context, response PasskeyAuthenticationResponse, expectedChallenge []byte) (int, error) {
|
||||||
|
// TODO: Implement full WebAuthn verification
|
||||||
|
// 1. Verify clientDataJSON contains correct challenge and origin
|
||||||
|
// 2. Verify authenticatorData
|
||||||
|
// 3. Verify signature using stored public key
|
||||||
|
// 4. Update sign counter and check for cloning
|
||||||
|
|
||||||
|
// Get credential from database
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var credentialJSON sql.NullString
|
||||||
|
|
||||||
|
runQuery := func() error {
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential::text FROM %s($1)`, p.sqlNames.PasskeyGetCredential)
|
||||||
|
return p.getDB().QueryRowContext(ctx, query, response.RawID).Scan(&success, &errorMsg, &credentialJSON)
|
||||||
|
}
|
||||||
|
err := runQuery()
|
||||||
|
if isDBClosed(err) {
|
||||||
|
if reconnErr := p.reconnectDB(); reconnErr == nil {
|
||||||
|
err = runQuery()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to get credential: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return 0, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return 0, fmt.Errorf("credential not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse credential
|
||||||
|
var cred struct {
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
SignCount uint32 `json:"sign_count"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(credentialJSON.String), &cred); err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to parse credential: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: Verify signature here
|
||||||
|
// For now, we'll just update the counter as a placeholder
|
||||||
|
|
||||||
|
// Update counter (in production, this should be done after successful verification)
|
||||||
|
newCounter := cred.SignCount + 1
|
||||||
|
var updateSuccess bool
|
||||||
|
var updateError sql.NullString
|
||||||
|
var cloneWarning sql.NullBool
|
||||||
|
|
||||||
|
updateQuery := fmt.Sprintf(`SELECT p_success, p_error, p_clone_warning FROM %s($1, $2)`, p.sqlNames.PasskeyUpdateCounter)
|
||||||
|
err = p.getDB().QueryRowContext(ctx, updateQuery, response.RawID, newCounter).Scan(&updateSuccess, &updateError, &cloneWarning)
|
||||||
|
if err != nil {
|
||||||
|
return 0, fmt.Errorf("failed to update counter: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cloneWarning.Valid && cloneWarning.Bool {
|
||||||
|
return 0, fmt.Errorf("credential cloning detected")
|
||||||
|
}
|
||||||
|
|
||||||
|
return cred.UserID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCredentials returns all passkey credentials for a user
|
||||||
|
func (p *DatabasePasskeyProvider) GetCredentials(ctx context.Context, userID int) ([]PasskeyCredential, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var credentialsJSON sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_credentials::text FROM %s($1)`, p.sqlNames.PasskeyGetUserCredentials)
|
||||||
|
err := p.getDB().QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &credentialsJSON)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to get credentials")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse credentials
|
||||||
|
var rawCreds []struct {
|
||||||
|
ID int `json:"id"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
CredentialID string `json:"credential_id"`
|
||||||
|
PublicKey string `json:"public_key"`
|
||||||
|
AttestationType string `json:"attestation_type"`
|
||||||
|
AAGUID string `json:"aaguid"`
|
||||||
|
SignCount uint32 `json:"sign_count"`
|
||||||
|
CloneWarning bool `json:"clone_warning"`
|
||||||
|
Transports []string `json:"transports"`
|
||||||
|
BackupEligible bool `json:"backup_eligible"`
|
||||||
|
BackupState bool `json:"backup_state"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
LastUsedAt time.Time `json:"last_used_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal([]byte(credentialsJSON.String), &rawCreds); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to parse credentials: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
credentials := make([]PasskeyCredential, 0, len(rawCreds))
|
||||||
|
for i := range rawCreds {
|
||||||
|
raw := rawCreds[i]
|
||||||
|
credID, err := base64.StdEncoding.DecodeString(raw.CredentialID)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
pubKey, err := base64.StdEncoding.DecodeString(raw.PublicKey)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
aaguid, _ := base64.StdEncoding.DecodeString(raw.AAGUID)
|
||||||
|
|
||||||
|
credentials = append(credentials, PasskeyCredential{
|
||||||
|
ID: fmt.Sprintf("%d", raw.ID),
|
||||||
|
UserID: raw.UserID,
|
||||||
|
CredentialID: credID,
|
||||||
|
PublicKey: pubKey,
|
||||||
|
AttestationType: raw.AttestationType,
|
||||||
|
AAGUID: aaguid,
|
||||||
|
SignCount: raw.SignCount,
|
||||||
|
CloneWarning: raw.CloneWarning,
|
||||||
|
Transports: raw.Transports,
|
||||||
|
BackupEligible: raw.BackupEligible,
|
||||||
|
BackupState: raw.BackupState,
|
||||||
|
Name: raw.Name,
|
||||||
|
CreatedAt: raw.CreatedAt,
|
||||||
|
LastUsedAt: raw.LastUsedAt,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return credentials, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteCredential removes a passkey credential
|
||||||
|
func (p *DatabasePasskeyProvider) DeleteCredential(ctx context.Context, userID int, credentialID string) error {
|
||||||
|
credID, err := base64.StdEncoding.DecodeString(credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid credential ID: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, p.sqlNames.PasskeyDeleteCredential)
|
||||||
|
err = p.getDB().QueryRowContext(ctx, query, userID, credID).Scan(&success, &errorMsg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to delete credential: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to delete credential")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateCredentialName updates the friendly name of a credential
|
||||||
|
func (p *DatabasePasskeyProvider) UpdateCredentialName(ctx context.Context, userID int, credentialID string, name string) error {
|
||||||
|
credID, err := base64.StdEncoding.DecodeString(credentialID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid credential ID: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3)`, p.sqlNames.PasskeyUpdateName)
|
||||||
|
err = p.getDB().QueryRowContext(ctx, query, userID, credID, name).Scan(&success, &errorMsg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to update credential name: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to update credential name")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+455
@@ -0,0 +1,455 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
|
||||||
|
"github.com/tidwall/gjson"
|
||||||
|
"github.com/tidwall/sjson"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ColumnSecurity struct {
|
||||||
|
Schema string `json:"schema"`
|
||||||
|
Tablename string `json:"tablename"`
|
||||||
|
Path []string `json:"path"`
|
||||||
|
ExtraFilters map[string]string `json:"extra_filters"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
Accesstype string `json:"accesstype"`
|
||||||
|
MaskStart int `json:"mask_start"`
|
||||||
|
MaskEnd int `json:"mask_end"`
|
||||||
|
MaskInvert bool `json:"mask_invert"`
|
||||||
|
MaskChar string `json:"mask_char"`
|
||||||
|
Control string `json:"control"`
|
||||||
|
ID int `json:"id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type RowSecurity struct {
|
||||||
|
Schema string `json:"schema"`
|
||||||
|
Tablename string `json:"tablename"`
|
||||||
|
Template string `json:"template"`
|
||||||
|
HasBlock bool `json:"has_block"`
|
||||||
|
UserID int `json:"user_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) string {
|
||||||
|
str := m.Template
|
||||||
|
str = strings.ReplaceAll(str, "{PrimaryKeyName}", pPrimaryKeyName)
|
||||||
|
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
|
||||||
|
str = strings.ReplaceAll(str, "{SchemaName}", m.Schema)
|
||||||
|
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%d", m.UserID))
|
||||||
|
return str
|
||||||
|
}
|
||||||
|
|
||||||
|
// SecurityList manages security state and caching
|
||||||
|
// It wraps a SecurityProvider and provides caching and utility methods
|
||||||
|
type SecurityList struct {
|
||||||
|
provider SecurityProvider
|
||||||
|
|
||||||
|
ColumnSecurityMutex sync.RWMutex
|
||||||
|
ColumnSecurity map[string][]ColumnSecurity
|
||||||
|
RowSecurityMutex sync.RWMutex
|
||||||
|
RowSecurity map[string]RowSecurity
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSecurityList creates a new security list with the given provider
|
||||||
|
func NewSecurityList(provider SecurityProvider) (*SecurityList, error) {
|
||||||
|
if provider == nil {
|
||||||
|
return nil, fmt.Errorf("security provider cannot be nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
return &SecurityList{
|
||||||
|
provider: provider,
|
||||||
|
ColumnSecurity: make(map[string][]ColumnSecurity),
|
||||||
|
RowSecurity: make(map[string]RowSecurity),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Provider returns the underlying security provider
|
||||||
|
func (m *SecurityList) Provider() SecurityProvider {
|
||||||
|
return m.provider
|
||||||
|
}
|
||||||
|
|
||||||
|
type CONTEXT_KEY string
|
||||||
|
|
||||||
|
const SECURITY_CONTEXT_KEY CONTEXT_KEY = "SecurityList"
|
||||||
|
|
||||||
|
func maskString(pString string, maskStart, maskEnd int, maskChar string, invert bool) string {
|
||||||
|
strLen := len(pString)
|
||||||
|
middleIndex := (strLen / 2)
|
||||||
|
newStr := ""
|
||||||
|
if maskStart == 0 && maskEnd == 0 {
|
||||||
|
maskStart = strLen
|
||||||
|
maskEnd = strLen
|
||||||
|
}
|
||||||
|
if maskEnd > strLen {
|
||||||
|
maskEnd = strLen
|
||||||
|
}
|
||||||
|
if maskStart > strLen {
|
||||||
|
maskStart = strLen
|
||||||
|
}
|
||||||
|
if maskChar == "" {
|
||||||
|
maskChar = "*"
|
||||||
|
}
|
||||||
|
for index, char := range pString {
|
||||||
|
if invert && index >= middleIndex-maskStart && index <= middleIndex {
|
||||||
|
newStr += maskChar
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if invert && index <= middleIndex+maskEnd && index >= middleIndex {
|
||||||
|
newStr += maskChar
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !invert && index <= maskStart {
|
||||||
|
newStr += maskChar
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !invert && index >= strLen-1-maskEnd {
|
||||||
|
newStr += maskChar
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
newStr += string(char)
|
||||||
|
}
|
||||||
|
|
||||||
|
return newStr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *SecurityList) ColumSecurityApplyOnRecord(prevRecord reflect.Value, newRecord reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) ([]string, error) {
|
||||||
|
cols := make([]string, 0)
|
||||||
|
if m.ColumnSecurity == nil {
|
||||||
|
return cols, fmt.Errorf("security not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
if prevRecord.Type() != newRecord.Type() {
|
||||||
|
logger.Error("prev:%s and new:%s record type mismatch", prevRecord.Type(), newRecord.Type())
|
||||||
|
return cols, fmt.Errorf("prev and new record type mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
m.ColumnSecurityMutex.RLock()
|
||||||
|
defer m.ColumnSecurityMutex.RUnlock()
|
||||||
|
|
||||||
|
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
|
||||||
|
if !ok || colsecList == nil {
|
||||||
|
return cols, fmt.Errorf("no column security data")
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range colsecList {
|
||||||
|
colsec := &colsecList[i]
|
||||||
|
if !strings.EqualFold(colsec.Accesstype, "mask") && !strings.EqualFold(colsec.Accesstype, "hide") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
lastRecords := interateStruct(prevRecord)
|
||||||
|
newRecords := interateStruct(newRecord)
|
||||||
|
var lastLoopField, lastLoopNewField reflect.Value
|
||||||
|
pathLen := len(colsec.Path)
|
||||||
|
for i, path := range colsec.Path {
|
||||||
|
var nameType, fieldName string
|
||||||
|
if len(newRecords) == 0 {
|
||||||
|
if lastLoopNewField.IsValid() && lastLoopField.IsValid() && i < pathLen-1 {
|
||||||
|
lastLoopNewField.Set(lastLoopField)
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
for ri := range newRecords {
|
||||||
|
if !newRecords[ri].IsValid() || !lastRecords[ri].IsValid() {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
var field, oldField reflect.Value
|
||||||
|
|
||||||
|
columnData := reflection.GetModelColumnDetail(newRecords[ri])
|
||||||
|
lastColumnData := reflection.GetModelColumnDetail(lastRecords[ri])
|
||||||
|
for i, cols := range columnData {
|
||||||
|
if cols.SQLName != "" && strings.EqualFold(cols.SQLName, path) {
|
||||||
|
nameType = "sql"
|
||||||
|
fieldName = cols.SQLName
|
||||||
|
field = cols.FieldValue
|
||||||
|
oldField = lastColumnData[i].FieldValue
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if cols.Name != "" && strings.EqualFold(cols.Name, path) {
|
||||||
|
nameType = "struct"
|
||||||
|
fieldName = cols.Name
|
||||||
|
field = cols.FieldValue
|
||||||
|
oldField = lastColumnData[i].FieldValue
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !field.IsValid() || !oldField.IsValid() {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
lastLoopField = oldField
|
||||||
|
lastLoopNewField = field
|
||||||
|
|
||||||
|
if i == pathLen-1 {
|
||||||
|
if strings.Contains(strings.ToLower(fieldName), "json") {
|
||||||
|
prevSrc := oldField.Bytes()
|
||||||
|
newSrc := field.Bytes()
|
||||||
|
pathstr := strings.Join(colsec.Path, ".")
|
||||||
|
prevPathValue := gjson.GetBytes(prevSrc, pathstr)
|
||||||
|
newBytes, err := sjson.SetBytes(newSrc, pathstr, prevPathValue.Str)
|
||||||
|
if err == nil {
|
||||||
|
if field.CanSet() {
|
||||||
|
field.SetBytes(newBytes)
|
||||||
|
} else {
|
||||||
|
logger.Warn("Value not settable: %v", field)
|
||||||
|
cols = append(cols, pathstr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
if nameType == "sql" {
|
||||||
|
if strings.EqualFold(colsec.Accesstype, "mask") || strings.EqualFold(colsec.Accesstype, "hide") {
|
||||||
|
field.Set(oldField)
|
||||||
|
cols = append(cols, strings.Join(colsec.Path, "."))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
lastRecords = interateStruct(field)
|
||||||
|
newRecords = interateStruct(oldField)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return cols, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func interateStruct(val reflect.Value) []reflect.Value {
|
||||||
|
list := make([]reflect.Value, 0)
|
||||||
|
|
||||||
|
switch val.Kind() {
|
||||||
|
case reflect.Pointer, reflect.Interface:
|
||||||
|
elem := val.Elem()
|
||||||
|
if elem.IsValid() {
|
||||||
|
list = append(list, interateStruct(elem)...)
|
||||||
|
}
|
||||||
|
return list
|
||||||
|
case reflect.Array, reflect.Slice:
|
||||||
|
for i := 0; i < val.Len(); i++ {
|
||||||
|
elem := val.Index(i)
|
||||||
|
if !elem.IsValid() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
list = append(list, interateStruct(elem)...)
|
||||||
|
}
|
||||||
|
return list
|
||||||
|
case reflect.Struct:
|
||||||
|
list = append(list, val)
|
||||||
|
return list
|
||||||
|
default:
|
||||||
|
return list
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func setColSecValue(fieldsrc reflect.Value, colsec ColumnSecurity, fieldTypeName string) (int, reflect.Value) {
|
||||||
|
fieldval := fieldsrc
|
||||||
|
if fieldsrc.Kind() == reflect.Pointer || fieldsrc.Kind() == reflect.Interface {
|
||||||
|
fieldval = fieldval.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
|
fieldKindLower := strings.ToLower(fieldval.Kind().String())
|
||||||
|
switch {
|
||||||
|
case strings.Contains(fieldKindLower, "int") &&
|
||||||
|
(strings.EqualFold(colsec.Accesstype, "mask") || strings.EqualFold(colsec.Accesstype, "hide")):
|
||||||
|
if fieldval.CanInt() && fieldval.CanSet() {
|
||||||
|
fieldval.SetInt(0)
|
||||||
|
}
|
||||||
|
case (strings.Contains(fieldKindLower, "time") || strings.Contains(fieldKindLower, "date")) &&
|
||||||
|
(strings.EqualFold(colsec.Accesstype, "mask") || strings.EqualFold(colsec.Accesstype, "hide")):
|
||||||
|
fieldval.SetZero()
|
||||||
|
case strings.Contains(fieldKindLower, "string"):
|
||||||
|
strVal := fieldval.String()
|
||||||
|
if strings.EqualFold(colsec.Accesstype, "mask") {
|
||||||
|
fieldval.SetString(maskString(strVal, colsec.MaskStart, colsec.MaskEnd, colsec.MaskChar, colsec.MaskInvert))
|
||||||
|
} else if strings.EqualFold(colsec.Accesstype, "hide") {
|
||||||
|
fieldval.SetString("")
|
||||||
|
}
|
||||||
|
case strings.Contains(fieldTypeName, "json") &&
|
||||||
|
(strings.EqualFold(colsec.Accesstype, "mask") || strings.EqualFold(colsec.Accesstype, "hide")):
|
||||||
|
if len(colsec.Path) < 2 {
|
||||||
|
return 1, fieldval
|
||||||
|
}
|
||||||
|
pathstr := strings.Join(colsec.Path, ".")
|
||||||
|
src := fieldval.Bytes()
|
||||||
|
pathValue := gjson.GetBytes(src, pathstr)
|
||||||
|
strValue := pathValue.String()
|
||||||
|
if strings.EqualFold(colsec.Accesstype, "mask") {
|
||||||
|
strValue = maskString(strValue, colsec.MaskStart, colsec.MaskEnd, colsec.MaskChar, colsec.MaskInvert)
|
||||||
|
} else if strings.EqualFold(colsec.Accesstype, "hide") {
|
||||||
|
strValue = ""
|
||||||
|
}
|
||||||
|
newBytes, err := sjson.SetBytes(src, pathstr, strValue)
|
||||||
|
if err == nil {
|
||||||
|
fieldval.SetBytes(newBytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0, fieldsrc
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) (reflect.Value, error) {
|
||||||
|
defer logger.CatchPanic("ApplyColumnSecurity")()
|
||||||
|
|
||||||
|
if m.ColumnSecurity == nil {
|
||||||
|
return records, fmt.Errorf("security not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
m.ColumnSecurityMutex.RLock()
|
||||||
|
defer m.ColumnSecurityMutex.RUnlock()
|
||||||
|
|
||||||
|
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
|
||||||
|
if !ok || colsecList == nil {
|
||||||
|
return records, fmt.Errorf("nocolumn security data")
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range colsecList {
|
||||||
|
colsec := &colsecList[i]
|
||||||
|
if !strings.EqualFold(colsec.Accesstype, "mask") && !strings.EqualFold(colsec.Accesstype, "hide") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if records.Kind() == reflect.Array || records.Kind() == reflect.Slice {
|
||||||
|
for i := 0; i < records.Len(); i++ {
|
||||||
|
record := records.Index(i)
|
||||||
|
if !record.IsValid() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
lastRecord := interateStruct(record)
|
||||||
|
pathLen := len(colsec.Path)
|
||||||
|
for i, path := range colsec.Path {
|
||||||
|
var field reflect.Value
|
||||||
|
var nameType, fieldName string
|
||||||
|
if len(lastRecord) == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
columnData := reflection.GetModelColumnDetail(lastRecord[0])
|
||||||
|
for _, cols := range columnData {
|
||||||
|
if cols.SQLName != "" && strings.EqualFold(cols.SQLName, path) {
|
||||||
|
nameType = "sql"
|
||||||
|
fieldName = cols.SQLName
|
||||||
|
field = cols.FieldValue
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if cols.Name != "" && strings.EqualFold(cols.Name, path) {
|
||||||
|
nameType = "struct"
|
||||||
|
fieldName = cols.Name
|
||||||
|
field = cols.FieldValue
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if i == pathLen-1 {
|
||||||
|
if nameType == "sql" || nameType == "struct" {
|
||||||
|
setColSecValue(field, *colsec, fieldName)
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if field.IsValid() {
|
||||||
|
lastRecord = interateStruct(field)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return records, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *SecurityList) LoadColumnSecurity(ctx context.Context, pUserID int, pSchema, pTablename string, pOverwrite bool) error {
|
||||||
|
if m.provider == nil {
|
||||||
|
return fmt.Errorf("security provider not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
m.ColumnSecurityMutex.Lock()
|
||||||
|
defer m.ColumnSecurityMutex.Unlock()
|
||||||
|
|
||||||
|
if m.ColumnSecurity == nil {
|
||||||
|
m.ColumnSecurity = make(map[string][]ColumnSecurity, 0)
|
||||||
|
}
|
||||||
|
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
|
||||||
|
|
||||||
|
if pOverwrite || m.ColumnSecurity[secKey] == nil {
|
||||||
|
m.ColumnSecurity[secKey] = make([]ColumnSecurity, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call the provider to load security rules
|
||||||
|
colSecList, err := m.provider.GetColumnSecurity(ctx, pUserID, pSchema, pTablename)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("GetColumnSecurity failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
m.ColumnSecurity[secKey] = colSecList
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *SecurityList) ClearSecurity(pUserID int, pSchema, pTablename string) error {
|
||||||
|
var filtered []ColumnSecurity
|
||||||
|
m.ColumnSecurityMutex.Lock()
|
||||||
|
defer m.ColumnSecurityMutex.Unlock()
|
||||||
|
|
||||||
|
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
|
||||||
|
list, ok := m.ColumnSecurity[secKey]
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range list {
|
||||||
|
cs := &list[i]
|
||||||
|
if cs.Schema != pSchema && cs.Tablename != pTablename && cs.UserID != pUserID {
|
||||||
|
filtered = append(filtered, *cs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
m.ColumnSecurity[secKey] = filtered
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserID int, pSchema, pTablename string, pOverwrite bool) (RowSecurity, error) {
|
||||||
|
if m.provider == nil {
|
||||||
|
return RowSecurity{}, fmt.Errorf("security provider not set")
|
||||||
|
}
|
||||||
|
|
||||||
|
m.RowSecurityMutex.Lock()
|
||||||
|
defer m.RowSecurityMutex.Unlock()
|
||||||
|
|
||||||
|
if m.RowSecurity == nil {
|
||||||
|
m.RowSecurity = make(map[string]RowSecurity, 0)
|
||||||
|
}
|
||||||
|
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
|
||||||
|
|
||||||
|
// Call the provider to load security rules
|
||||||
|
record, err := m.provider.GetRowSecurity(ctx, pUserID, pSchema, pTablename)
|
||||||
|
if err != nil {
|
||||||
|
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
m.RowSecurity[secKey] = record
|
||||||
|
return record, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *SecurityList) GetRowSecurityTemplate(pUserID int, pSchema, pTablename string) (RowSecurity, error) {
|
||||||
|
defer logger.CatchPanic("GetRowSecurityTemplate")()
|
||||||
|
|
||||||
|
if m.RowSecurity == nil {
|
||||||
|
return RowSecurity{}, fmt.Errorf("security not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
m.RowSecurityMutex.RLock()
|
||||||
|
defer m.RowSecurityMutex.RUnlock()
|
||||||
|
|
||||||
|
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
|
||||||
|
if !ok {
|
||||||
|
return RowSecurity{}, fmt.Errorf("no row security data")
|
||||||
|
}
|
||||||
|
|
||||||
|
return rowSec, nil
|
||||||
|
}
|
||||||
+1145
File diff suppressed because it is too large
Load Diff
+267
@@ -0,0 +1,267 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"regexp"
|
||||||
|
)
|
||||||
|
|
||||||
|
var validSQLIdentifier = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`)
|
||||||
|
|
||||||
|
// SQLNames defines all configurable SQL stored procedure and table names
|
||||||
|
// used by the security package. Override individual fields to remap
|
||||||
|
// to custom database objects. Use DefaultSQLNames() for baseline defaults,
|
||||||
|
// and MergeSQLNames() to apply partial overrides.
|
||||||
|
type SQLNames struct {
|
||||||
|
// Auth procedures (DatabaseAuthenticator)
|
||||||
|
Login string // default: "resolvespec_login"
|
||||||
|
Register string // default: "resolvespec_register"
|
||||||
|
Logout string // default: "resolvespec_logout"
|
||||||
|
Session string // default: "resolvespec_session"
|
||||||
|
SessionUpdate string // default: "resolvespec_session_update"
|
||||||
|
RefreshToken string // default: "resolvespec_refresh_token"
|
||||||
|
|
||||||
|
// JWT procedures (JWTAuthenticator)
|
||||||
|
JWTLogin string // default: "resolvespec_jwt_login"
|
||||||
|
JWTLogout string // default: "resolvespec_jwt_logout"
|
||||||
|
|
||||||
|
// Security policy procedures
|
||||||
|
ColumnSecurity string // default: "resolvespec_column_security"
|
||||||
|
RowSecurity string // default: "resolvespec_row_security"
|
||||||
|
|
||||||
|
// TOTP procedures (DatabaseTwoFactorProvider)
|
||||||
|
TOTPEnable string // default: "resolvespec_totp_enable"
|
||||||
|
TOTPDisable string // default: "resolvespec_totp_disable"
|
||||||
|
TOTPGetStatus string // default: "resolvespec_totp_get_status"
|
||||||
|
TOTPGetSecret string // default: "resolvespec_totp_get_secret"
|
||||||
|
TOTPRegenerateBackup string // default: "resolvespec_totp_regenerate_backup_codes"
|
||||||
|
TOTPValidateBackupCode string // default: "resolvespec_totp_validate_backup_code"
|
||||||
|
|
||||||
|
// Passkey procedures (DatabasePasskeyProvider)
|
||||||
|
PasskeyStoreCredential string // default: "resolvespec_passkey_store_credential"
|
||||||
|
PasskeyGetCredsByUsername string // default: "resolvespec_passkey_get_credentials_by_username"
|
||||||
|
PasskeyGetCredential string // default: "resolvespec_passkey_get_credential"
|
||||||
|
PasskeyUpdateCounter string // default: "resolvespec_passkey_update_counter"
|
||||||
|
PasskeyGetUserCredentials string // default: "resolvespec_passkey_get_user_credentials"
|
||||||
|
PasskeyDeleteCredential string // default: "resolvespec_passkey_delete_credential"
|
||||||
|
PasskeyUpdateName string // default: "resolvespec_passkey_update_name"
|
||||||
|
PasskeyLogin string // default: "resolvespec_passkey_login"
|
||||||
|
|
||||||
|
// Password reset procedures (DatabaseAuthenticator)
|
||||||
|
PasswordResetRequest string // default: "resolvespec_password_reset_request"
|
||||||
|
PasswordResetComplete string // default: "resolvespec_password_reset"
|
||||||
|
|
||||||
|
// OAuth2 procedures (DatabaseAuthenticator OAuth2 methods)
|
||||||
|
OAuthGetOrCreateUser string // default: "resolvespec_oauth_getorcreateuser"
|
||||||
|
OAuthCreateSession string // default: "resolvespec_oauth_createsession"
|
||||||
|
OAuthGetRefreshToken string // default: "resolvespec_oauth_getrefreshtoken"
|
||||||
|
OAuthUpdateRefreshToken string // default: "resolvespec_oauth_updaterefreshtoken"
|
||||||
|
OAuthGetUser string // default: "resolvespec_oauth_getuser"
|
||||||
|
|
||||||
|
// OAuth2 server procedures (OAuthServer persistence)
|
||||||
|
OAuthRegisterClient string // default: "resolvespec_oauth_register_client"
|
||||||
|
OAuthGetClient string // default: "resolvespec_oauth_get_client"
|
||||||
|
OAuthSaveCode string // default: "resolvespec_oauth_save_code"
|
||||||
|
OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code"
|
||||||
|
OAuthIntrospect string // default: "resolvespec_oauth_introspect"
|
||||||
|
OAuthRevoke string // default: "resolvespec_oauth_revoke"
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultSQLNames returns an SQLNames with all default resolvespec_* values.
|
||||||
|
func DefaultSQLNames() *SQLNames {
|
||||||
|
return &SQLNames{
|
||||||
|
Login: "resolvespec_login",
|
||||||
|
Register: "resolvespec_register",
|
||||||
|
Logout: "resolvespec_logout",
|
||||||
|
Session: "resolvespec_session",
|
||||||
|
SessionUpdate: "resolvespec_session_update",
|
||||||
|
RefreshToken: "resolvespec_refresh_token",
|
||||||
|
|
||||||
|
JWTLogin: "resolvespec_jwt_login",
|
||||||
|
JWTLogout: "resolvespec_jwt_logout",
|
||||||
|
|
||||||
|
ColumnSecurity: "resolvespec_column_security",
|
||||||
|
RowSecurity: "resolvespec_row_security",
|
||||||
|
|
||||||
|
TOTPEnable: "resolvespec_totp_enable",
|
||||||
|
TOTPDisable: "resolvespec_totp_disable",
|
||||||
|
TOTPGetStatus: "resolvespec_totp_get_status",
|
||||||
|
TOTPGetSecret: "resolvespec_totp_get_secret",
|
||||||
|
TOTPRegenerateBackup: "resolvespec_totp_regenerate_backup_codes",
|
||||||
|
TOTPValidateBackupCode: "resolvespec_totp_validate_backup_code",
|
||||||
|
|
||||||
|
PasskeyStoreCredential: "resolvespec_passkey_store_credential",
|
||||||
|
PasskeyGetCredsByUsername: "resolvespec_passkey_get_credentials_by_username",
|
||||||
|
PasskeyGetCredential: "resolvespec_passkey_get_credential",
|
||||||
|
PasskeyUpdateCounter: "resolvespec_passkey_update_counter",
|
||||||
|
PasskeyGetUserCredentials: "resolvespec_passkey_get_user_credentials",
|
||||||
|
PasskeyDeleteCredential: "resolvespec_passkey_delete_credential",
|
||||||
|
PasskeyUpdateName: "resolvespec_passkey_update_name",
|
||||||
|
PasskeyLogin: "resolvespec_passkey_login",
|
||||||
|
|
||||||
|
PasswordResetRequest: "resolvespec_password_reset_request",
|
||||||
|
PasswordResetComplete: "resolvespec_password_reset",
|
||||||
|
|
||||||
|
OAuthGetOrCreateUser: "resolvespec_oauth_getorcreateuser",
|
||||||
|
OAuthCreateSession: "resolvespec_oauth_createsession",
|
||||||
|
OAuthGetRefreshToken: "resolvespec_oauth_getrefreshtoken",
|
||||||
|
OAuthUpdateRefreshToken: "resolvespec_oauth_updaterefreshtoken",
|
||||||
|
OAuthGetUser: "resolvespec_oauth_getuser",
|
||||||
|
|
||||||
|
OAuthRegisterClient: "resolvespec_oauth_register_client",
|
||||||
|
OAuthGetClient: "resolvespec_oauth_get_client",
|
||||||
|
OAuthSaveCode: "resolvespec_oauth_save_code",
|
||||||
|
OAuthExchangeCode: "resolvespec_oauth_exchange_code",
|
||||||
|
OAuthIntrospect: "resolvespec_oauth_introspect",
|
||||||
|
OAuthRevoke: "resolvespec_oauth_revoke",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MergeSQLNames returns a copy of base with any non-empty fields from override applied.
|
||||||
|
// If override is nil, a copy of base is returned.
|
||||||
|
func MergeSQLNames(base, override *SQLNames) *SQLNames {
|
||||||
|
if override == nil {
|
||||||
|
copied := *base
|
||||||
|
return &copied
|
||||||
|
}
|
||||||
|
merged := *base
|
||||||
|
if override.Login != "" {
|
||||||
|
merged.Login = override.Login
|
||||||
|
}
|
||||||
|
if override.Register != "" {
|
||||||
|
merged.Register = override.Register
|
||||||
|
}
|
||||||
|
if override.Logout != "" {
|
||||||
|
merged.Logout = override.Logout
|
||||||
|
}
|
||||||
|
if override.Session != "" {
|
||||||
|
merged.Session = override.Session
|
||||||
|
}
|
||||||
|
if override.SessionUpdate != "" {
|
||||||
|
merged.SessionUpdate = override.SessionUpdate
|
||||||
|
}
|
||||||
|
if override.RefreshToken != "" {
|
||||||
|
merged.RefreshToken = override.RefreshToken
|
||||||
|
}
|
||||||
|
if override.JWTLogin != "" {
|
||||||
|
merged.JWTLogin = override.JWTLogin
|
||||||
|
}
|
||||||
|
if override.JWTLogout != "" {
|
||||||
|
merged.JWTLogout = override.JWTLogout
|
||||||
|
}
|
||||||
|
if override.ColumnSecurity != "" {
|
||||||
|
merged.ColumnSecurity = override.ColumnSecurity
|
||||||
|
}
|
||||||
|
if override.RowSecurity != "" {
|
||||||
|
merged.RowSecurity = override.RowSecurity
|
||||||
|
}
|
||||||
|
if override.TOTPEnable != "" {
|
||||||
|
merged.TOTPEnable = override.TOTPEnable
|
||||||
|
}
|
||||||
|
if override.TOTPDisable != "" {
|
||||||
|
merged.TOTPDisable = override.TOTPDisable
|
||||||
|
}
|
||||||
|
if override.TOTPGetStatus != "" {
|
||||||
|
merged.TOTPGetStatus = override.TOTPGetStatus
|
||||||
|
}
|
||||||
|
if override.TOTPGetSecret != "" {
|
||||||
|
merged.TOTPGetSecret = override.TOTPGetSecret
|
||||||
|
}
|
||||||
|
if override.TOTPRegenerateBackup != "" {
|
||||||
|
merged.TOTPRegenerateBackup = override.TOTPRegenerateBackup
|
||||||
|
}
|
||||||
|
if override.TOTPValidateBackupCode != "" {
|
||||||
|
merged.TOTPValidateBackupCode = override.TOTPValidateBackupCode
|
||||||
|
}
|
||||||
|
if override.PasskeyStoreCredential != "" {
|
||||||
|
merged.PasskeyStoreCredential = override.PasskeyStoreCredential
|
||||||
|
}
|
||||||
|
if override.PasskeyGetCredsByUsername != "" {
|
||||||
|
merged.PasskeyGetCredsByUsername = override.PasskeyGetCredsByUsername
|
||||||
|
}
|
||||||
|
if override.PasskeyGetCredential != "" {
|
||||||
|
merged.PasskeyGetCredential = override.PasskeyGetCredential
|
||||||
|
}
|
||||||
|
if override.PasskeyUpdateCounter != "" {
|
||||||
|
merged.PasskeyUpdateCounter = override.PasskeyUpdateCounter
|
||||||
|
}
|
||||||
|
if override.PasskeyGetUserCredentials != "" {
|
||||||
|
merged.PasskeyGetUserCredentials = override.PasskeyGetUserCredentials
|
||||||
|
}
|
||||||
|
if override.PasskeyDeleteCredential != "" {
|
||||||
|
merged.PasskeyDeleteCredential = override.PasskeyDeleteCredential
|
||||||
|
}
|
||||||
|
if override.PasskeyUpdateName != "" {
|
||||||
|
merged.PasskeyUpdateName = override.PasskeyUpdateName
|
||||||
|
}
|
||||||
|
if override.PasskeyLogin != "" {
|
||||||
|
merged.PasskeyLogin = override.PasskeyLogin
|
||||||
|
}
|
||||||
|
if override.PasswordResetRequest != "" {
|
||||||
|
merged.PasswordResetRequest = override.PasswordResetRequest
|
||||||
|
}
|
||||||
|
if override.PasswordResetComplete != "" {
|
||||||
|
merged.PasswordResetComplete = override.PasswordResetComplete
|
||||||
|
}
|
||||||
|
if override.OAuthGetOrCreateUser != "" {
|
||||||
|
merged.OAuthGetOrCreateUser = override.OAuthGetOrCreateUser
|
||||||
|
}
|
||||||
|
if override.OAuthCreateSession != "" {
|
||||||
|
merged.OAuthCreateSession = override.OAuthCreateSession
|
||||||
|
}
|
||||||
|
if override.OAuthGetRefreshToken != "" {
|
||||||
|
merged.OAuthGetRefreshToken = override.OAuthGetRefreshToken
|
||||||
|
}
|
||||||
|
if override.OAuthUpdateRefreshToken != "" {
|
||||||
|
merged.OAuthUpdateRefreshToken = override.OAuthUpdateRefreshToken
|
||||||
|
}
|
||||||
|
if override.OAuthGetUser != "" {
|
||||||
|
merged.OAuthGetUser = override.OAuthGetUser
|
||||||
|
}
|
||||||
|
if override.OAuthRegisterClient != "" {
|
||||||
|
merged.OAuthRegisterClient = override.OAuthRegisterClient
|
||||||
|
}
|
||||||
|
if override.OAuthGetClient != "" {
|
||||||
|
merged.OAuthGetClient = override.OAuthGetClient
|
||||||
|
}
|
||||||
|
if override.OAuthSaveCode != "" {
|
||||||
|
merged.OAuthSaveCode = override.OAuthSaveCode
|
||||||
|
}
|
||||||
|
if override.OAuthExchangeCode != "" {
|
||||||
|
merged.OAuthExchangeCode = override.OAuthExchangeCode
|
||||||
|
}
|
||||||
|
if override.OAuthIntrospect != "" {
|
||||||
|
merged.OAuthIntrospect = override.OAuthIntrospect
|
||||||
|
}
|
||||||
|
if override.OAuthRevoke != "" {
|
||||||
|
merged.OAuthRevoke = override.OAuthRevoke
|
||||||
|
}
|
||||||
|
return &merged
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateSQLNames checks that all non-empty fields in names are valid SQL identifiers.
|
||||||
|
// Returns an error if any field contains invalid characters.
|
||||||
|
func ValidateSQLNames(names *SQLNames) error {
|
||||||
|
v := reflect.ValueOf(names).Elem()
|
||||||
|
typ := v.Type()
|
||||||
|
for i := 0; i < v.NumField(); i++ {
|
||||||
|
field := v.Field(i)
|
||||||
|
if field.Kind() != reflect.String {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
val := field.String()
|
||||||
|
if val != "" && !validSQLIdentifier.MatchString(val) {
|
||||||
|
return fmt.Errorf("SQLNames.%s contains invalid characters: %q", typ.Field(i).Name, val)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveSQLNames merges an optional override with defaults.
|
||||||
|
// Used by constructors that accept variadic *SQLNames.
|
||||||
|
func resolveSQLNames(override ...*SQLNames) *SQLNames {
|
||||||
|
if len(override) > 0 && override[0] != nil {
|
||||||
|
return MergeSQLNames(DefaultSQLNames(), override[0])
|
||||||
|
}
|
||||||
|
return DefaultSQLNames()
|
||||||
|
}
|
||||||
+188
@@ -0,0 +1,188 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/hmac"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha1"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/sha512"
|
||||||
|
"encoding/base32"
|
||||||
|
"encoding/binary"
|
||||||
|
"fmt"
|
||||||
|
"hash"
|
||||||
|
"math"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TwoFactorAuthProvider defines interface for 2FA operations
|
||||||
|
type TwoFactorAuthProvider interface {
|
||||||
|
// Generate2FASecret creates a new secret for a user
|
||||||
|
Generate2FASecret(userID int, issuer, accountName string) (*TwoFactorSecret, error)
|
||||||
|
|
||||||
|
// Validate2FACode verifies a TOTP code
|
||||||
|
Validate2FACode(secret string, code string) (bool, error)
|
||||||
|
|
||||||
|
// Enable2FA activates 2FA for a user (store secret in your database)
|
||||||
|
Enable2FA(userID int, secret string, backupCodes []string) error
|
||||||
|
|
||||||
|
// Disable2FA deactivates 2FA for a user
|
||||||
|
Disable2FA(userID int) error
|
||||||
|
|
||||||
|
// Get2FAStatus checks if user has 2FA enabled
|
||||||
|
Get2FAStatus(userID int) (bool, error)
|
||||||
|
|
||||||
|
// Get2FASecret retrieves the user's 2FA secret
|
||||||
|
Get2FASecret(userID int) (string, error)
|
||||||
|
|
||||||
|
// GenerateBackupCodes creates backup codes for 2FA
|
||||||
|
GenerateBackupCodes(userID int, count int) ([]string, error)
|
||||||
|
|
||||||
|
// ValidateBackupCode checks and consumes a backup code
|
||||||
|
ValidateBackupCode(userID int, code string) (bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TwoFactorSecret contains 2FA setup information
|
||||||
|
type TwoFactorSecret struct {
|
||||||
|
Secret string `json:"secret"` // Base32 encoded secret
|
||||||
|
QRCodeURL string `json:"qr_code_url"` // URL for QR code generation
|
||||||
|
BackupCodes []string `json:"backup_codes"` // One-time backup codes
|
||||||
|
Issuer string `json:"issuer"` // Application name
|
||||||
|
AccountName string `json:"account_name"` // User identifier (email/username)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TwoFactorConfig holds TOTP configuration
|
||||||
|
type TwoFactorConfig struct {
|
||||||
|
Algorithm string // SHA1, SHA256, SHA512
|
||||||
|
Digits int // Number of digits in code (6 or 8)
|
||||||
|
Period int // Time step in seconds (default 30)
|
||||||
|
SkewWindow int // Number of time steps to check before/after (default 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultTwoFactorConfig returns standard TOTP configuration
|
||||||
|
func DefaultTwoFactorConfig() *TwoFactorConfig {
|
||||||
|
return &TwoFactorConfig{
|
||||||
|
Algorithm: "SHA1",
|
||||||
|
Digits: 6,
|
||||||
|
Period: 30,
|
||||||
|
SkewWindow: 1,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TOTPGenerator handles TOTP code generation and validation
|
||||||
|
type TOTPGenerator struct {
|
||||||
|
config *TwoFactorConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTOTPGenerator creates a new TOTP generator with config
|
||||||
|
func NewTOTPGenerator(config *TwoFactorConfig) *TOTPGenerator {
|
||||||
|
if config == nil {
|
||||||
|
config = DefaultTwoFactorConfig()
|
||||||
|
}
|
||||||
|
return &TOTPGenerator{
|
||||||
|
config: config,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateSecret creates a random base32-encoded secret
|
||||||
|
func (t *TOTPGenerator) GenerateSecret() (string, error) {
|
||||||
|
secret := make([]byte, 20)
|
||||||
|
_, err := rand.Read(secret)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to generate random secret: %w", err)
|
||||||
|
}
|
||||||
|
return base32.StdEncoding.WithPadding(base32.NoPadding).EncodeToString(secret), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateQRCodeURL creates a URL for QR code generation
|
||||||
|
func (t *TOTPGenerator) GenerateQRCodeURL(secret, issuer, accountName string) string {
|
||||||
|
params := url.Values{}
|
||||||
|
params.Set("secret", secret)
|
||||||
|
params.Set("issuer", issuer)
|
||||||
|
params.Set("algorithm", t.config.Algorithm)
|
||||||
|
params.Set("digits", fmt.Sprintf("%d", t.config.Digits))
|
||||||
|
params.Set("period", fmt.Sprintf("%d", t.config.Period))
|
||||||
|
|
||||||
|
label := url.PathEscape(fmt.Sprintf("%s:%s", issuer, accountName))
|
||||||
|
return fmt.Sprintf("otpauth://totp/%s?%s", label, params.Encode())
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateCode creates a TOTP code for a given time
|
||||||
|
func (t *TOTPGenerator) GenerateCode(secret string, timestamp time.Time) (string, error) {
|
||||||
|
// Decode secret
|
||||||
|
key, err := base32.StdEncoding.WithPadding(base32.NoPadding).DecodeString(strings.ToUpper(secret))
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("invalid secret: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Calculate counter (time steps since Unix epoch)
|
||||||
|
counter := uint64(timestamp.Unix()) / uint64(t.config.Period)
|
||||||
|
|
||||||
|
// Generate HMAC
|
||||||
|
h := t.getHashFunc()
|
||||||
|
mac := hmac.New(h, key)
|
||||||
|
|
||||||
|
// Convert counter to 8-byte array
|
||||||
|
buf := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint64(buf, counter)
|
||||||
|
mac.Write(buf)
|
||||||
|
|
||||||
|
sum := mac.Sum(nil)
|
||||||
|
|
||||||
|
// Dynamic truncation
|
||||||
|
offset := sum[len(sum)-1] & 0x0f
|
||||||
|
truncated := binary.BigEndian.Uint32(sum[offset:]) & 0x7fffffff
|
||||||
|
|
||||||
|
// Generate code with specified digits
|
||||||
|
code := truncated % uint32(math.Pow10(t.config.Digits))
|
||||||
|
|
||||||
|
format := fmt.Sprintf("%%0%dd", t.config.Digits)
|
||||||
|
return fmt.Sprintf(format, code), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateCode checks if a code is valid for the secret
|
||||||
|
func (t *TOTPGenerator) ValidateCode(secret, code string) (bool, error) {
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
// Check current time and skew window
|
||||||
|
for i := -t.config.SkewWindow; i <= t.config.SkewWindow; i++ {
|
||||||
|
timestamp := now.Add(time.Duration(i*t.config.Period) * time.Second)
|
||||||
|
expected, err := t.GenerateCode(secret, timestamp)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if code == expected {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// getHashFunc returns the hash function based on algorithm
|
||||||
|
func (t *TOTPGenerator) getHashFunc() func() hash.Hash {
|
||||||
|
switch strings.ToUpper(t.config.Algorithm) {
|
||||||
|
case "SHA256":
|
||||||
|
return sha256.New
|
||||||
|
case "SHA512":
|
||||||
|
return sha512.New
|
||||||
|
default:
|
||||||
|
return sha1.New
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateBackupCodes creates random backup codes
|
||||||
|
func GenerateBackupCodes(count int) ([]string, error) {
|
||||||
|
codes := make([]string, count)
|
||||||
|
for i := 0; i < count; i++ {
|
||||||
|
code := make([]byte, 4)
|
||||||
|
_, err := rand.Read(code)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate backup code: %w", err)
|
||||||
|
}
|
||||||
|
codes[i] = fmt.Sprintf("%08X", binary.BigEndian.Uint32(code))
|
||||||
|
}
|
||||||
|
return codes, nil
|
||||||
|
}
|
||||||
+134
@@ -0,0 +1,134 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TwoFactorAuthenticator wraps an Authenticator and adds 2FA support
|
||||||
|
type TwoFactorAuthenticator struct {
|
||||||
|
baseAuth Authenticator
|
||||||
|
totp *TOTPGenerator
|
||||||
|
provider TwoFactorAuthProvider
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTwoFactorAuthenticator creates a new 2FA-enabled authenticator
|
||||||
|
func NewTwoFactorAuthenticator(baseAuth Authenticator, provider TwoFactorAuthProvider, config *TwoFactorConfig) *TwoFactorAuthenticator {
|
||||||
|
if config == nil {
|
||||||
|
config = DefaultTwoFactorConfig()
|
||||||
|
}
|
||||||
|
return &TwoFactorAuthenticator{
|
||||||
|
baseAuth: baseAuth,
|
||||||
|
totp: NewTOTPGenerator(config),
|
||||||
|
provider: provider,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Login authenticates with 2FA support
|
||||||
|
func (t *TwoFactorAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||||
|
// First, perform standard authentication
|
||||||
|
resp, err := t.baseAuth.Login(ctx, req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if user has 2FA enabled
|
||||||
|
if resp.User == nil {
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
has2FA, err := t.provider.Get2FAStatus(resp.User.UserID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to check 2FA status: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !has2FA {
|
||||||
|
// User doesn't have 2FA enabled, return normal response
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// User has 2FA enabled
|
||||||
|
if req.TwoFactorCode == "" {
|
||||||
|
// No 2FA code provided, require it
|
||||||
|
resp.Requires2FA = true
|
||||||
|
resp.Token = "" // Don't return token until 2FA is verified
|
||||||
|
resp.RefreshToken = ""
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 2FA code
|
||||||
|
secret, err := t.provider.Get2FASecret(resp.User.UserID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get 2FA secret: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try TOTP code first
|
||||||
|
valid, err := t.totp.ValidateCode(secret, req.TwoFactorCode)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to validate 2FA code: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !valid {
|
||||||
|
// Try backup code
|
||||||
|
valid, err = t.provider.ValidateBackupCode(resp.User.UserID, req.TwoFactorCode)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to validate backup code: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !valid {
|
||||||
|
return nil, fmt.Errorf("invalid 2FA code")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2FA verified, return full response with token
|
||||||
|
resp.User.TwoFactorEnabled = true
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Logout delegates to base authenticator
|
||||||
|
func (t *TwoFactorAuthenticator) Logout(ctx context.Context, req LogoutRequest) error {
|
||||||
|
return t.baseAuth.Logout(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticate delegates to base authenticator
|
||||||
|
func (t *TwoFactorAuthenticator) Authenticate(r *http.Request) (*UserContext, error) {
|
||||||
|
return t.baseAuth.Authenticate(r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Setup2FA initiates 2FA setup for a user
|
||||||
|
func (t *TwoFactorAuthenticator) Setup2FA(userID int, issuer, accountName string) (*TwoFactorSecret, error) {
|
||||||
|
return t.provider.Generate2FASecret(userID, issuer, accountName)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enable2FA completes 2FA setup after user confirms with a valid code
|
||||||
|
func (t *TwoFactorAuthenticator) Enable2FA(userID int, secret, verificationCode string) error {
|
||||||
|
// Verify the code before enabling
|
||||||
|
valid, err := t.totp.ValidateCode(secret, verificationCode)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to validate code: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !valid {
|
||||||
|
return fmt.Errorf("invalid verification code")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate backup codes
|
||||||
|
backupCodes, err := t.provider.GenerateBackupCodes(userID, 10)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to generate backup codes: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enable 2FA
|
||||||
|
return t.provider.Enable2FA(userID, secret, backupCodes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Disable2FA removes 2FA from a user account
|
||||||
|
func (t *TwoFactorAuthenticator) Disable2FA(userID int) error {
|
||||||
|
return t.provider.Disable2FA(userID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegenerateBackupCodes creates new backup codes for a user
|
||||||
|
func (t *TwoFactorAuthenticator) RegenerateBackupCodes(userID int, count int) ([]string, error) {
|
||||||
|
return t.provider.GenerateBackupCodes(userID, count)
|
||||||
|
}
|
||||||
+229
@@ -0,0 +1,229 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DatabaseTwoFactorProvider implements TwoFactorAuthProvider using PostgreSQL stored procedures
|
||||||
|
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
||||||
|
// See totp_database_schema.sql for procedure definitions
|
||||||
|
type DatabaseTwoFactorProvider struct {
|
||||||
|
db *sql.DB
|
||||||
|
totpGen *TOTPGenerator
|
||||||
|
sqlNames *SQLNames
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDatabaseTwoFactorProvider creates a new database-backed 2FA provider
|
||||||
|
func NewDatabaseTwoFactorProvider(db *sql.DB, config *TwoFactorConfig, names ...*SQLNames) *DatabaseTwoFactorProvider {
|
||||||
|
if config == nil {
|
||||||
|
config = DefaultTwoFactorConfig()
|
||||||
|
}
|
||||||
|
return &DatabaseTwoFactorProvider{
|
||||||
|
db: db,
|
||||||
|
totpGen: NewTOTPGenerator(config),
|
||||||
|
sqlNames: resolveSQLNames(names...),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate2FASecret creates a new secret for a user
|
||||||
|
func (p *DatabaseTwoFactorProvider) Generate2FASecret(userID int, issuer, accountName string) (*TwoFactorSecret, error) {
|
||||||
|
secret, err := p.totpGen.GenerateSecret()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate secret: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
qrURL := p.totpGen.GenerateQRCodeURL(secret, issuer, accountName)
|
||||||
|
|
||||||
|
backupCodes, err := GenerateBackupCodes(10)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate backup codes: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &TwoFactorSecret{
|
||||||
|
Secret: secret,
|
||||||
|
QRCodeURL: qrURL,
|
||||||
|
BackupCodes: backupCodes,
|
||||||
|
Issuer: issuer,
|
||||||
|
AccountName: accountName,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate2FACode verifies a TOTP code
|
||||||
|
func (p *DatabaseTwoFactorProvider) Validate2FACode(secret string, code string) (bool, error) {
|
||||||
|
return p.totpGen.ValidateCode(secret, code)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enable2FA activates 2FA for a user
|
||||||
|
func (p *DatabaseTwoFactorProvider) Enable2FA(userID int, secret string, backupCodes []string) error {
|
||||||
|
// Hash backup codes for secure storage
|
||||||
|
hashedCodes := make([]string, len(backupCodes))
|
||||||
|
for i, code := range backupCodes {
|
||||||
|
hash := sha256.Sum256([]byte(code))
|
||||||
|
hashedCodes[i] = hex.EncodeToString(hash[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to JSON array
|
||||||
|
codesJSON, err := json.Marshal(hashedCodes)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal backup codes: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call stored procedure
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3::jsonb)`, p.sqlNames.TOTPEnable)
|
||||||
|
err = p.db.QueryRow(query, userID, secret, string(codesJSON)).Scan(&success, &errorMsg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("enable 2FA query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to enable 2FA")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Disable2FA deactivates 2FA for a user
|
||||||
|
func (p *DatabaseTwoFactorProvider) Disable2FA(userID int) error {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1)`, p.sqlNames.TOTPDisable)
|
||||||
|
err := p.db.QueryRow(query, userID).Scan(&success, &errorMsg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("disable 2FA query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("failed to disable 2FA")
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get2FAStatus checks if user has 2FA enabled
|
||||||
|
func (p *DatabaseTwoFactorProvider) Get2FAStatus(userID int) (bool, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var enabled bool
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_enabled FROM %s($1)`, p.sqlNames.TOTPGetStatus)
|
||||||
|
err := p.db.QueryRow(query, userID).Scan(&success, &errorMsg, &enabled)
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("get 2FA status query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return false, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return false, fmt.Errorf("failed to get 2FA status")
|
||||||
|
}
|
||||||
|
|
||||||
|
return enabled, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get2FASecret retrieves the user's 2FA secret
|
||||||
|
func (p *DatabaseTwoFactorProvider) Get2FASecret(userID int) (string, error) {
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var secret sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_secret FROM %s($1)`, p.sqlNames.TOTPGetSecret)
|
||||||
|
err := p.db.QueryRow(query, userID).Scan(&success, &errorMsg, &secret)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return "", fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("failed to get 2FA secret")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !secret.Valid {
|
||||||
|
return "", fmt.Errorf("2FA secret not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
return secret.String, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateBackupCodes creates backup codes for 2FA
|
||||||
|
func (p *DatabaseTwoFactorProvider) GenerateBackupCodes(userID int, count int) ([]string, error) {
|
||||||
|
codes, err := GenerateBackupCodes(count)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate backup codes: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hash backup codes for storage
|
||||||
|
hashedCodes := make([]string, len(codes))
|
||||||
|
for i, code := range codes {
|
||||||
|
hash := sha256.Sum256([]byte(code))
|
||||||
|
hashedCodes[i] = hex.EncodeToString(hash[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to JSON array
|
||||||
|
codesJSON, err := json.Marshal(hashedCodes)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal backup codes: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call stored procedure
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2::jsonb)`, p.sqlNames.TOTPRegenerateBackup)
|
||||||
|
err = p.db.QueryRow(query, userID, string(codesJSON)).Scan(&success, &errorMsg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("regenerate backup codes query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return nil, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed to regenerate backup codes")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return unhashed codes to user (only time they see them)
|
||||||
|
return codes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateBackupCode checks and consumes a backup code
|
||||||
|
func (p *DatabaseTwoFactorProvider) ValidateBackupCode(userID int, code string) (bool, error) {
|
||||||
|
// Hash the code
|
||||||
|
hash := sha256.Sum256([]byte(code))
|
||||||
|
codeHash := hex.EncodeToString(hash[:])
|
||||||
|
|
||||||
|
var success bool
|
||||||
|
var errorMsg sql.NullString
|
||||||
|
var valid bool
|
||||||
|
|
||||||
|
query := fmt.Sprintf(`SELECT p_success, p_error, p_valid FROM %s($1, $2)`, p.sqlNames.TOTPValidateBackupCode)
|
||||||
|
err := p.db.QueryRow(query, userID, codeHash).Scan(&success, &errorMsg, &valid)
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("validate backup code query failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !success {
|
||||||
|
if errorMsg.Valid {
|
||||||
|
return false, fmt.Errorf("%s", errorMsg.String)
|
||||||
|
}
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return valid, nil
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user