chore(deps): bump go modules in backend and example
Backend (backend/go.mod): - github.com/go-pkgz/auth/v2 v2.1.2 → v2.1.4 - github.com/klauspost/compress v1.18.5 → v1.18.6 - github.com/redis/go-redis/v9 v9.18.0 → v9.19.0 - github.com/slack-go/slack v0.21.1 → v0.23.1 - golang.org/x/crypto v0.50.0 → v0.51.0 - golang.org/x/image v0.39.0 → v0.40.0 - golang.org/x/net v0.53.0 → v0.54.0 - golang.org/x/sys v0.43.0 → v0.44.0 - golang.org/x/text v0.36.0 → v0.37.0 Example (backend/_example/memory_store/go.mod): - golang.org/x/crypto v0.50.0 → v0.51.0 - golang.org/x/image v0.39.0 → v0.40.0 - golang.org/x/net v0.53.0 → v0.54.0 - golang.org/x/sys v0.43.0 → v0.44.0 Transitive cleanup: github.com/dgryski/go-rendezvous is no longer required after redis/go-redis bump and gets pruned by `go mod tidy`. `go mod tidy` + `go mod vendor` run on both modules. Both build with -race and full test suites pass.
This commit is contained in:
committed by
Umputun
parent
f3a7dea1f1
commit
45c17a913f
@@ -30,10 +30,10 @@ require (
|
||||
github.com/rs/xid v1.6.0 // indirect
|
||||
github.com/russross/blackfriday/v2 v2.1.0 // indirect
|
||||
go.etcd.io/bbolt v1.4.3 // indirect
|
||||
golang.org/x/crypto v0.50.0 // indirect
|
||||
golang.org/x/image v0.39.0 // indirect
|
||||
golang.org/x/net v0.53.0 // indirect
|
||||
golang.org/x/sys v0.43.0 // indirect
|
||||
golang.org/x/crypto v0.51.0 // indirect
|
||||
golang.org/x/image v0.40.0 // indirect
|
||||
golang.org/x/net v0.54.0 // indirect
|
||||
golang.org/x/sys v0.44.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
)
|
||||
|
||||
|
||||
@@ -62,10 +62,10 @@ golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliY
|
||||
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
|
||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
|
||||
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
||||
golang.org/x/image v0.39.0 h1:skVYidAEVKgn8lZ602XO75asgXBgLj9G/FE3RbuPFww=
|
||||
golang.org/x/image v0.39.0/go.mod h1:sIbmppfU+xFLPIG0FoVUTvyBMmgng1/XAMhQ2ft0hpA=
|
||||
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI=
|
||||
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8=
|
||||
golang.org/x/image v0.40.0 h1:Tw4GyDXMo+daZN1znreBRC3VayR1aLFUyUEOLUdW1a8=
|
||||
golang.org/x/image v0.40.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA=
|
||||
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.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
@@ -80,8 +80,8 @@ golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
|
||||
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
|
||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||
golang.org/x/net v0.33.0/go.mod h1:HXLR5J+9DxmrqMwG9qjGCxZ+zKXxBru04zlTvWlWuN4=
|
||||
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
||||
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
|
||||
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w=
|
||||
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ=
|
||||
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.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
@@ -102,8 +102,8 @@ golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.17.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.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
||||
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ=
|
||||
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/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=
|
||||
|
||||
+9
-10
@@ -9,7 +9,7 @@ require (
|
||||
github.com/didip/tollbooth/v8 v8.0.1
|
||||
github.com/go-chi/chi/v5 v5.2.5
|
||||
github.com/go-chi/cors v1.2.2
|
||||
github.com/go-pkgz/auth/v2 v2.1.2
|
||||
github.com/go-pkgz/auth/v2 v2.1.4
|
||||
github.com/go-pkgz/jrpc v0.4.0
|
||||
github.com/go-pkgz/lcw/v2 v2.0.0
|
||||
github.com/go-pkgz/lgr v0.12.3
|
||||
@@ -30,9 +30,9 @@ require (
|
||||
github.com/stretchr/testify v1.11.1
|
||||
go.etcd.io/bbolt v1.4.3
|
||||
go.uber.org/goleak v1.3.0
|
||||
golang.org/x/crypto v0.50.0
|
||||
golang.org/x/image v0.39.0
|
||||
golang.org/x/net v0.53.0
|
||||
golang.org/x/crypto v0.51.0
|
||||
golang.org/x/image v0.40.0
|
||||
golang.org/x/net v0.54.0
|
||||
golang.org/x/oauth2 v0.36.0
|
||||
)
|
||||
|
||||
@@ -43,7 +43,6 @@ require (
|
||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/dghubble/oauth1 v0.7.3 // indirect
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
|
||||
github.com/dlclark/regexp2 v1.12.0 // indirect
|
||||
github.com/go-oauth2/oauth2/v4 v4.5.4 // indirect
|
||||
github.com/go-pkgz/email v0.6.0 // indirect
|
||||
@@ -55,12 +54,12 @@ require (
|
||||
github.com/gorilla/websocket v1.5.3 // indirect
|
||||
github.com/hashicorp/errwrap v1.1.0 // indirect
|
||||
github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect
|
||||
github.com/klauspost/compress v1.18.5 // indirect
|
||||
github.com/klauspost/compress v1.18.6 // indirect
|
||||
github.com/montanaflynn/stats v0.9.0 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/redis/go-redis/v9 v9.18.0 // indirect
|
||||
github.com/redis/go-redis/v9 v9.19.0 // indirect
|
||||
github.com/rrivera/identicon v0.0.0-20240116195454-d5ba35832c0d // indirect
|
||||
github.com/slack-go/slack v0.21.1 // indirect
|
||||
github.com/slack-go/slack v0.23.1 // indirect
|
||||
github.com/xdg-go/pbkdf2 v1.0.0 // indirect
|
||||
github.com/xdg-go/scram v1.2.0 // indirect
|
||||
github.com/xdg-go/stringprep v1.0.4 // indirect
|
||||
@@ -68,7 +67,7 @@ require (
|
||||
go.mongodb.org/mongo-driver v1.17.9 // indirect
|
||||
go.uber.org/atomic v1.11.0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
golang.org/x/sys v0.43.0 // indirect
|
||||
golang.org/x/text v0.36.0 // indirect
|
||||
golang.org/x/sys v0.44.0 // indirect
|
||||
golang.org/x/text v0.37.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
)
|
||||
|
||||
+22
-24
@@ -32,8 +32,6 @@ 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/dghubble/oauth1 v0.7.3 h1:EkEM/zMDMp3zOsX2DC/ZQ2vnEX3ELK0/l9kb+vs4ptE=
|
||||
github.com/dghubble/oauth1 v0.7.3/go.mod h1:oxTe+az9NSMIucDPDCCtzJGsPhciJV33xocHfcR2sVY=
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
|
||||
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
|
||||
github.com/didip/tollbooth/v8 v8.0.1 h1:VAAapTo1t4Bn6bbpcHjuovwoa9u3JH++wgjbpWv+rB8=
|
||||
github.com/didip/tollbooth/v8 v8.0.1/go.mod h1:oEd9l+ep373d7DmvKLc0a5gasPOev2mTewi6KPQBGJ4=
|
||||
github.com/dlclark/regexp2 v1.12.0 h1:0j4c5qQmnC6XOWNjP3PIXURXN2gWx76rd3KvgdPkCz8=
|
||||
@@ -48,8 +46,8 @@ github.com/go-chi/cors v1.2.2 h1:Jmey33TE+b+rB7fT8MUy1u0I4L+NARQlK6LhzKPSyQE=
|
||||
github.com/go-chi/cors v1.2.2/go.mod h1:sSbTewc+6wYHBBCW7ytsFSn836hqM7JxpglAy2Vzc58=
|
||||
github.com/go-oauth2/oauth2/v4 v4.5.4 h1:YjI0tmGW8oxVhn9QSBIxlr641QugWrJY5UWa6XmLcW0=
|
||||
github.com/go-oauth2/oauth2/v4 v4.5.4/go.mod h1:BXiOY+QZtZy2ewbsGk2B5P8TWmtz/Rf7ES5ZttQFxfQ=
|
||||
github.com/go-pkgz/auth/v2 v2.1.2 h1:tB2TxXR3oiUpqU45K6Zqw5TpB6AUeVQkTAuTCGCI01g=
|
||||
github.com/go-pkgz/auth/v2 v2.1.2/go.mod h1:IvxxhJIrwd1hKqFwQgBF9i+sMTmGfzAw66wmhw1zfJc=
|
||||
github.com/go-pkgz/auth/v2 v2.1.4 h1:bCF0vMscOrShF2gelcvKPgskpwQNGCk6AQcoXOf2kbE=
|
||||
github.com/go-pkgz/auth/v2 v2.1.4/go.mod h1:IvxxhJIrwd1hKqFwQgBF9i+sMTmGfzAw66wmhw1zfJc=
|
||||
github.com/go-pkgz/email v0.6.0 h1:snZnXldjeF4PgKSjnx9Fa25mtOgFpAOEeWvnQvrxjLE=
|
||||
github.com/go-pkgz/email v0.6.0/go.mod h1:+wgi4x7S33IuCzfcCM5euN0GwQG6XvO/PBLxrNffYLI=
|
||||
github.com/go-pkgz/expirable-cache/v3 v3.1.0 h1:s05P851/O6QJ6Mc+7o2bh9aGtD3romB1SxDTXifdoqc=
|
||||
@@ -108,10 +106,10 @@ github.com/jessevdk/go-flags v1.6.1 h1:Cvu5U8UGrLay1rZfv/zP7iLpSHGUZ/Ou68T0iX1bB
|
||||
github.com/jessevdk/go-flags v1.6.1/go.mod h1:Mk8T1hIAWpOiJiHa9rJASDK2UGWji0EuPGBnNLMooyc=
|
||||
github.com/jtolds/gls v4.20.0+incompatible h1:xdiiI2gbIgH/gLH7ADydsJ1uDOEzR8yvV7C0MuV77Wo=
|
||||
github.com/jtolds/gls v4.20.0+incompatible/go.mod h1:QJZ7F/aHp+rZTRtaJ1ow/lLfFfVYBRgL+9YlvaHOwJU=
|
||||
github.com/klauspost/compress v1.18.5 h1:/h1gH5Ce+VWNLSWqPzOVn6XBO+vJbCNGvjoaGBFW2IE=
|
||||
github.com/klauspost/compress v1.18.5/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||
github.com/klauspost/cpuid/v2 v2.0.9 h1:lgaqFMSdTdQYdZ04uHyN2d/eKdOMyi2YLSvlQIBFYa4=
|
||||
github.com/klauspost/cpuid/v2 v2.0.9/go.mod h1:FInQzS24/EEf25PyTYn52gqo7WaD8xa0213Md/qVLRg=
|
||||
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
|
||||
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.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||
@@ -126,8 +124,8 @@ github.com/moul/http2curl v1.0.0 h1:dRMWoAtb+ePxMlLkrCbAqh4TlPHXvoGUSQ323/9Zahs=
|
||||
github.com/moul/http2curl v1.0.0/go.mod h1:8UbvGypXm98wA/IqH45anm5Y2Z6ep6O31QGOAZ3H0fQ=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/redis/go-redis/v9 v9.18.0 h1:pMkxYPkEbMPwRdenAzUNyFNrDgHx9U+DrBabWNfSRQs=
|
||||
github.com/redis/go-redis/v9 v9.18.0/go.mod h1:k3ufPphLU5YXwNTUcCRXGxUoF1fqxnhFQmscfkCoDA0=
|
||||
github.com/redis/go-redis/v9 v9.19.0 h1:XPVaaPSnG6RhYf7p+rmSa9zZfeVAnWsH5h3lxthOm/k=
|
||||
github.com/redis/go-redis/v9 v9.19.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
||||
github.com/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZVejAe8=
|
||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||
github.com/rrivera/identicon v0.0.0-20240116195454-d5ba35832c0d h1:l3+2LWCbVxn5itfvXAfH9n4YL9jh8l1g5zcncbIc1cs=
|
||||
@@ -140,8 +138,8 @@ github.com/sergi/go-diff v1.1.0 h1:we8PVUC3FE2uYfodKH/nBHMSetSfHDR6scGdBi+erh0=
|
||||
github.com/sergi/go-diff v1.1.0/go.mod h1:STckp+ISIX8hZLjrqAeVduY0gWCT9IjLuqbuNXdaHfM=
|
||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e h1:MRM5ITcdelLK2j1vwZ3Je0FKVCfqOLp5zO6trqMLYs0=
|
||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e/go.mod h1:XV66xRDqSt+GTGFMVlhk3ULuV0y9ZmzeVGR4mloJI3M=
|
||||
github.com/slack-go/slack v0.21.1 h1:vBHR+IkaXbv9RLY6w/RiN82D+5/OTI06CGqrlZ3Vyas=
|
||||
github.com/slack-go/slack v0.21.1/go.mod h1:K81UmCivcYd/5Jmz8vLBfuyoZ3B4rQC2GHVXHteXiAE=
|
||||
github.com/slack-go/slack v0.23.1 h1:ZS5B96wxxYQRwvJ3/vJFtqtUZi3tXhsZCyT44Nv7M80=
|
||||
github.com/slack-go/slack v0.23.1/go.mod h1:H0yR/YBuRJ39RkE+JpV/d/oEsbanzTRowR82bCN0cEs=
|
||||
github.com/smartystreets/assertions v1.1.0 h1:MkTeG1DMwsrdH7QtLXy5W+fUxWq+vmb6cLmyJ7aRtF0=
|
||||
github.com/smartystreets/assertions v1.1.0/go.mod h1:tcbTF8ujkAEcZ8TElKY+i30BzYlVhC/LOxJk7iOWnoo=
|
||||
github.com/smartystreets/goconvey v1.6.4 h1:fv0U8FUIMPNf1L9lnHLvLhgicrIVChEkdzIKYqbNC9s=
|
||||
@@ -191,8 +189,8 @@ github.com/yudai/golcs v0.0.0-20170316035057-ecda9a501e82/go.mod h1:lgjkn3NuSvDf
|
||||
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
|
||||
github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M=
|
||||
github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw=
|
||||
github.com/zeebo/xxh3 v1.0.2 h1:xZmwmqxHZA8AI603jOQ0tMqmBr9lPeFwGg6d+xy9DC0=
|
||||
github.com/zeebo/xxh3 v1.0.2/go.mod h1:5NWz9Sef7zIDm2JHfFlcQvNekmcEl9ekUZQQKCYaDcA=
|
||||
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
||||
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
||||
go.etcd.io/bbolt v1.4.3 h1:dEadXpI6G79deX5prL3QRNP6JB8UxVkqo4UPnHaNXJo=
|
||||
go.etcd.io/bbolt v1.4.3/go.mod h1:tKQlpPaYCVFctUIgFKFnAlvbmB3tpy1vkTnDWohtc0E=
|
||||
go.mongodb.org/mongo-driver v1.17.9 h1:IexDdCuuNJ3BHrELgBlyaH9p60JXAvdzWR128q+U5tU=
|
||||
@@ -207,10 +205,10 @@ golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliY
|
||||
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
|
||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
||||
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
|
||||
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
|
||||
golang.org/x/image v0.39.0 h1:skVYidAEVKgn8lZ602XO75asgXBgLj9G/FE3RbuPFww=
|
||||
golang.org/x/image v0.39.0/go.mod h1:sIbmppfU+xFLPIG0FoVUTvyBMmgng1/XAMhQ2ft0hpA=
|
||||
golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI=
|
||||
golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8=
|
||||
golang.org/x/image v0.40.0 h1:Tw4GyDXMo+daZN1znreBRC3VayR1aLFUyUEOLUdW1a8=
|
||||
golang.org/x/image v0.40.0/go.mod h1:uIc348UZMSvS5Z65CVZ7iDPaNobNFEPeJ4kbqTOszmA=
|
||||
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.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
|
||||
@@ -225,8 +223,8 @@ golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
|
||||
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
|
||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||
golang.org/x/net v0.33.0/go.mod h1:HXLR5J+9DxmrqMwG9qjGCxZ+zKXxBru04zlTvWlWuN4=
|
||||
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
||||
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
|
||||
golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w=
|
||||
golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ=
|
||||
golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs=
|
||||
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=
|
||||
@@ -249,8 +247,8 @@ golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.17.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.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
||||
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ=
|
||||
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/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=
|
||||
@@ -270,8 +268,8 @@ golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
|
||||
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.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
|
||||
golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg=
|
||||
golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164=
|
||||
golang.org/x/text v0.37.0 h1:Cqjiwd9eSg8e0QAkyCaQTNHFIIzWtidPahFWR83rTrc=
|
||||
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-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
||||
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
|
||||
|
||||
-21
@@ -1,21 +0,0 @@
|
||||
The MIT License (MIT)
|
||||
|
||||
Copyright (c) 2017-2020 Damian Gryski <damian@gryski.com>
|
||||
|
||||
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.
|
||||
-79
@@ -1,79 +0,0 @@
|
||||
package rendezvous
|
||||
|
||||
type Rendezvous struct {
|
||||
nodes map[string]int
|
||||
nstr []string
|
||||
nhash []uint64
|
||||
hash Hasher
|
||||
}
|
||||
|
||||
type Hasher func(s string) uint64
|
||||
|
||||
func New(nodes []string, hash Hasher) *Rendezvous {
|
||||
r := &Rendezvous{
|
||||
nodes: make(map[string]int, len(nodes)),
|
||||
nstr: make([]string, len(nodes)),
|
||||
nhash: make([]uint64, len(nodes)),
|
||||
hash: hash,
|
||||
}
|
||||
|
||||
for i, n := range nodes {
|
||||
r.nodes[n] = i
|
||||
r.nstr[i] = n
|
||||
r.nhash[i] = hash(n)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *Rendezvous) Lookup(k string) string {
|
||||
// short-circuit if we're empty
|
||||
if len(r.nodes) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
khash := r.hash(k)
|
||||
|
||||
var midx int
|
||||
var mhash = xorshiftMult64(khash ^ r.nhash[0])
|
||||
|
||||
for i, nhash := range r.nhash[1:] {
|
||||
if h := xorshiftMult64(khash ^ nhash); h > mhash {
|
||||
midx = i + 1
|
||||
mhash = h
|
||||
}
|
||||
}
|
||||
|
||||
return r.nstr[midx]
|
||||
}
|
||||
|
||||
func (r *Rendezvous) Add(node string) {
|
||||
r.nodes[node] = len(r.nstr)
|
||||
r.nstr = append(r.nstr, node)
|
||||
r.nhash = append(r.nhash, r.hash(node))
|
||||
}
|
||||
|
||||
func (r *Rendezvous) Remove(node string) {
|
||||
// find index of node to remove
|
||||
nidx := r.nodes[node]
|
||||
|
||||
// remove from the slices
|
||||
l := len(r.nstr)
|
||||
r.nstr[nidx] = r.nstr[l]
|
||||
r.nstr = r.nstr[:l]
|
||||
|
||||
r.nhash[nidx] = r.nhash[l]
|
||||
r.nhash = r.nhash[:l]
|
||||
|
||||
// update the map
|
||||
delete(r.nodes, node)
|
||||
moved := r.nstr[nidx]
|
||||
r.nodes[moved] = nidx
|
||||
}
|
||||
|
||||
func xorshiftMult64(x uint64) uint64 {
|
||||
x ^= x >> 12 // a
|
||||
x ^= x << 25 // b
|
||||
x ^= x >> 27 // c
|
||||
return x * 2685821657736338717
|
||||
}
|
||||
+69
-9
@@ -7,6 +7,7 @@ import (
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/go-pkgz/rest"
|
||||
@@ -26,14 +27,16 @@ type Client struct {
|
||||
|
||||
// Service provides higher level wrapper allowing to construct everything and get back token middleware
|
||||
type Service struct {
|
||||
logger logger.L
|
||||
opts Opts
|
||||
jwtService *token.Service
|
||||
providers []provider.Service
|
||||
authMiddleware middleware.Authenticator
|
||||
avatarProxy *avatar.Proxy
|
||||
issuer string
|
||||
useGravatar bool
|
||||
logger logger.L
|
||||
opts Opts
|
||||
jwtService *token.Service
|
||||
providers []provider.Service
|
||||
authMiddleware middleware.Authenticator
|
||||
avatarProxy *avatar.Proxy
|
||||
issuer string
|
||||
useGravatar bool
|
||||
verifConfirmStore provider.VerifConfirmationStore
|
||||
verifConfirmStoreOnce sync.Once
|
||||
}
|
||||
|
||||
// Opts is a full set of all parameters to initialize Service
|
||||
@@ -84,6 +87,15 @@ type Opts struct {
|
||||
Logger logger.L // logger interface, default is no logging at all
|
||||
RefreshCache middleware.RefreshCache // optional cache to keep refreshed tokens
|
||||
ErrorHandler middleware.ErrorHandlerFunc // custom error handler for auth failures
|
||||
|
||||
// VerifConfirmationStore enforces one-shot consumption of email
|
||||
// confirmation tokens issued by the verify provider. The default
|
||||
// (nil) installs an in-memory store on first use of AddVerifProvider —
|
||||
// fine for single-instance deployments. Multi-instance deployments
|
||||
// MUST supply a shared backend (e.g. Redis) implementing
|
||||
// provider.VerifConfirmationStore, otherwise replay rejection works
|
||||
// only on the instance that consumed the token.
|
||||
VerifConfirmationStore provider.VerifConfirmationStore
|
||||
}
|
||||
|
||||
// NewService initializes everything
|
||||
@@ -225,7 +237,39 @@ func (s *Service) Handlers() (authHandler, avatarHandler http.Handler) {
|
||||
p.Handler(w, r)
|
||||
}
|
||||
|
||||
return http.HandlerFunc(ah), http.HandlerFunc(s.avatarProxy.Handler)
|
||||
return withSecurityHeaders(http.HandlerFunc(ah)), withSecurityHeaders(http.HandlerFunc(s.avatarProxy.Handler))
|
||||
}
|
||||
|
||||
// withSecurityHeaders wraps an auth response handler to apply strict CSP and nosniff
|
||||
// on every response. The go-pkgz/auth package's own response surface is JSON-only
|
||||
// for auth routes and images for the avatar route — no built-in HTML rendering
|
||||
// anywhere — so this CSP is unconditionally safe and gives the auth origin
|
||||
// defense-in-depth against any future trust-boundary regression that might emit a
|
||||
// renderable body.
|
||||
//
|
||||
// - Content-Security-Policy: default-src 'none'; sandbox — blocks inline scripts
|
||||
// and event handlers even if a body is ever served as HTML by mistake; the
|
||||
// sandbox directive additionally isolates any rendered document from this origin.
|
||||
// - X-Content-Type-Options: nosniff — prevents browsers from MIME-overriding the
|
||||
// declared Content-Type to a more dangerous one.
|
||||
//
|
||||
// The avatar Handler additionally sets Content-Disposition: inline; filename="avatar"
|
||||
// inside itself, so direct callers (tests, custom mounts) still get the full header
|
||||
// set without going through this wrapper.
|
||||
//
|
||||
// CONSUMER NOTE: custom providers added via Service.AddCustomHandler / AddProvider
|
||||
// are also wrapped. If a custom provider renders HTML (login forms, JS-based flows,
|
||||
// the dev_provider's login page, etc.), the strict CSP will block inline scripts and
|
||||
// event handlers on those pages. Such providers should either (a) override the CSP
|
||||
// for their own response by calling w.Header().Set("Content-Security-Policy", ...)
|
||||
// before writing — Set replaces the wrapper's value — or (b) move any required
|
||||
// scripts/styles to external files served from 'self'.
|
||||
func withSecurityHeaders(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Security-Policy", "default-src 'none'; sandbox; frame-ancestors 'none'")
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// Middleware returns auth middleware
|
||||
@@ -435,6 +479,21 @@ func (s *Service) AddDirectProviderWithUserIDFunc(name string, credChecker provi
|
||||
|
||||
// AddVerifProvider adds provider user's verification sent by sender
|
||||
func (s *Service) AddVerifProvider(name, msgTmpl string, sender provider.Sender) {
|
||||
s.verifConfirmStoreOnce.Do(func() {
|
||||
store := s.opts.VerifConfirmationStore
|
||||
// guard against a typed-nil VerifConfirmationStoreFunc: a non-nil
|
||||
// interface wrapping a nil func would survive the != nil check below
|
||||
// and silently disable replay protection (the handler-level guard
|
||||
// at LoginHandler then normalizes it to nil).
|
||||
if fn, ok := store.(provider.VerifConfirmationStoreFunc); ok && fn == nil {
|
||||
store = nil
|
||||
}
|
||||
if store != nil {
|
||||
s.verifConfirmStore = store
|
||||
return
|
||||
}
|
||||
s.verifConfirmStore = provider.NewInMemoryVerifStore()
|
||||
})
|
||||
dh := provider.VerifyHandler{
|
||||
L: s.logger,
|
||||
ProviderName: name,
|
||||
@@ -446,6 +505,7 @@ func (s *Service) AddVerifProvider(name, msgTmpl string, sender provider.Sender)
|
||||
UseGravatar: s.useGravatar,
|
||||
URL: s.opts.URL,
|
||||
AllowedRedirectHosts: s.opts.AllowedRedirectHosts,
|
||||
ConfirmationStore: s.verifConfirmStore,
|
||||
}
|
||||
s.addProvider(dh)
|
||||
}
|
||||
|
||||
+217
-54
@@ -11,6 +11,7 @@ import (
|
||||
"image/png"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -18,6 +19,7 @@ import (
|
||||
"github.com/go-pkgz/rest"
|
||||
"github.com/rrivera/identicon"
|
||||
"golang.org/x/image/draw"
|
||||
_ "golang.org/x/image/webp" // register WebP decoder so Discord-style .webp avatars validate
|
||||
|
||||
"github.com/go-pkgz/auth/v2/logger"
|
||||
"github.com/go-pkgz/auth/v2/token"
|
||||
@@ -26,6 +28,19 @@ import (
|
||||
// http.sniffLen is 512 bytes which is how much we need to read to detect content type
|
||||
const sniffLen = 512
|
||||
|
||||
// maxAvatarFetchSize bounds the bytes read from a remote avatar URL. 10 MiB is
|
||||
// generous for any reasonable avatar (Telegram caps photo at 5 MiB; Gravatar is
|
||||
// much smaller); the cap protects Proxy.Put against an upstream sending an
|
||||
// unbounded body that would exhaust process memory inside resize.
|
||||
const maxAvatarFetchSize = 10 << 20
|
||||
|
||||
// maxAvatarPixels caps the declared pixel count of an avatar before any raster
|
||||
// decode is allowed. Without this, a tiny compressed "decompression bomb" image
|
||||
// declaring e.g. 65535x65535 px would force image.Decode to allocate gigabytes
|
||||
// of pixel memory and OOM the auth service on a single login attempt. 16 MP
|
||||
// covers any realistic avatar (~4096x4096) while keeping peak allocation bounded.
|
||||
const maxAvatarPixels = 16 * 1024 * 1024
|
||||
|
||||
// Proxy provides http handler for avatars from avatar.Store
|
||||
// On user login token will call Put and it will retrieve and save picture locally.
|
||||
type Proxy struct {
|
||||
@@ -45,7 +60,7 @@ func (p *Proxy) Put(u token.User, client *http.Client) (avatarURL string, err er
|
||||
return "", fmt.Errorf("no picture for %s: %w", userID, e)
|
||||
}
|
||||
// put returns avatar base name, like 123456.image
|
||||
avatarID, e := p.Store.Put(userID, p.resize(bytes.NewBuffer(b), p.ResizeLimit))
|
||||
avatarID, e := p.Store.Put(userID, p.resize(b, p.ResizeLimit))
|
||||
if e != nil {
|
||||
return "", e
|
||||
}
|
||||
@@ -61,30 +76,73 @@ func (p *Proxy) Put(u token.User, client *http.Client) (avatarURL string, err er
|
||||
|
||||
body, err := p.load(u.Picture, client)
|
||||
if err != nil {
|
||||
p.Logf("[DEBUG] failed to fetch avatar from the orig %s, %v", u.Picture, err)
|
||||
p.Logf("[DEBUG] failed to fetch avatar from the orig %s, %v", redactAvatarURL(u.Picture), err)
|
||||
return genIdenticon(u.ID)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
if e := body.Close(); e != nil {
|
||||
p.Logf("[WARN] can't close response body, %s", e)
|
||||
}
|
||||
}()
|
||||
resized := p.resize(body, p.ResizeLimit)
|
||||
if resized == nil {
|
||||
// non-image upstream — refuse to store attacker-controlled bytes under
|
||||
// the user's avatar id and fall back to a generated identicon instead.
|
||||
p.Logf("[WARN] upstream avatar from %s is not a valid image, using identicon", redactAvatarURL(u.Picture))
|
||||
return genIdenticon(u.ID)
|
||||
}
|
||||
|
||||
avatarID, err := p.Store.Put(u.ID, p.resize(body, p.ResizeLimit)) // put returns avatar base name, like 123456.image
|
||||
avatarID, err := p.Store.Put(u.ID, resized) // put returns avatar base name, like 123456.image
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
p.Logf("[DEBUG] saved avatar from %s to %s, user %q", u.Picture, avatarID, u.Name)
|
||||
p.Logf("[DEBUG] saved avatar from %s to %s, user %q", redactAvatarURL(u.Picture), avatarID, u.Name)
|
||||
return p.URL + p.RoutePath + "/" + avatarID, nil
|
||||
}
|
||||
|
||||
// load avatar from remote url and return body. Caller has to close the reader
|
||||
func (p *Proxy) load(url string, client *http.Client) (rc io.ReadCloser, err error) {
|
||||
// load avatar from remote location
|
||||
// PutContent stores already-fetched avatar bytes via the underlying Store and returns
|
||||
// the proxied URL. It exists so providers that authenticate with credentials embedded
|
||||
// in the upstream URL (e.g. Telegram bot file API: /file/bot{TOKEN}/...) can fetch the
|
||||
// content themselves and avoid exposing the credential to Put's URL-fetching path —
|
||||
// where it would land in u.Picture, debug logs, and the user JSON returned to clients.
|
||||
//
|
||||
// Bytes are read into memory bounded by maxAvatarFetchSize so an unbounded caller
|
||||
// (e.g. a streaming HTTP body) cannot exhaust process memory.
|
||||
func (p *Proxy) PutContent(userID string, content io.Reader) (avatarURL string, err error) {
|
||||
body, err := io.ReadAll(io.LimitReader(content, maxAvatarFetchSize+1))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read avatar content for %s: %w", userID, err)
|
||||
}
|
||||
if int64(len(body)) > maxAvatarFetchSize {
|
||||
return "", fmt.Errorf("avatar content for %s exceeds %d bytes", userID, maxAvatarFetchSize)
|
||||
}
|
||||
resized := p.resize(body, p.ResizeLimit)
|
||||
if resized == nil {
|
||||
return "", fmt.Errorf("avatar content for %s is not a valid image", userID)
|
||||
}
|
||||
avatarID, err := p.Store.Put(userID, resized)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
p.Logf("[DEBUG] saved avatar bytes to %s, user %q", avatarID, userID)
|
||||
return p.URL + p.RoutePath + "/" + avatarID, nil
|
||||
}
|
||||
|
||||
// redactAvatarURL returns the hostname only, dropping scheme, userinfo, path,
|
||||
// query and fragment. This is enough to keep avatar URLs identifiable in logs
|
||||
// while ensuring credentials carried in any of those parts (e.g. Telegram bot
|
||||
// tokens, time-limited signed-URL tokens, basic-auth in userinfo) don't reach
|
||||
// log destinations. On parse failure a sentinel is returned.
|
||||
func redactAvatarURL(raw string) string {
|
||||
if u, err := url.Parse(raw); err == nil && u.Hostname() != "" {
|
||||
return u.Hostname()
|
||||
}
|
||||
return "<unparseable>"
|
||||
}
|
||||
|
||||
// load fetches an avatar from a remote url and returns the body bytes, capped at
|
||||
// maxAvatarFetchSize. The bytes are passed straight to resize without an
|
||||
// intermediate Reader wrapper so we don't pay for buffering twice.
|
||||
func (p *Proxy) load(url string, client *http.Client) ([]byte, error) {
|
||||
var resp *http.Response
|
||||
err = retry(5, time.Second, func() error {
|
||||
err := retry(5, time.Second, func() error {
|
||||
var e error
|
||||
resp, e = client.Get(url)
|
||||
return e
|
||||
@@ -92,19 +150,31 @@ func (p *Proxy) load(url string, client *http.Client) (rc io.ReadCloser, err err
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch avatar from the orig: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
_ = resp.Body.Close() // caller won't close on error
|
||||
return nil, fmt.Errorf("failed to get avatar from the orig, status %s", resp.Status)
|
||||
}
|
||||
|
||||
return resp.Body, nil
|
||||
// buffer the body up to the cap to fail fast on oversized inputs.
|
||||
// Reading +1 byte beyond the cap distinguishes "exactly cap" from "too big".
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxAvatarFetchSize+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read avatar body: %w", err)
|
||||
}
|
||||
if int64(len(body)) > maxAvatarFetchSize {
|
||||
return nil, fmt.Errorf("avatar body exceeds %d bytes", maxAvatarFetchSize)
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// Handler returns token routes for given provider
|
||||
func (p *Proxy) Handler(w http.ResponseWriter, r *http.Request) {
|
||||
setAvatarDefenseHeaders(w)
|
||||
|
||||
if r.Method != "GET" {
|
||||
w.WriteHeader(http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
elems := strings.Split(r.URL.Path, "/")
|
||||
avatarID := elems[len(elems)-1]
|
||||
@@ -113,19 +183,6 @@ func (p *Proxy) Handler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// enforce client-side caching
|
||||
etag := `"` + p.Store.ID(avatarID) + `"`
|
||||
w.Header().Set("Etag", etag)
|
||||
w.Header().Set("Cache-Control", "max-age=604800") // 7 days
|
||||
if match := r.Header.Get("If-None-Match"); match != "" {
|
||||
etag = strings.TrimPrefix(etag, `"`)
|
||||
etag = strings.TrimSuffix(etag, `"`)
|
||||
if match == etag {
|
||||
w.WriteHeader(http.StatusNotModified)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
avReader, size, err := p.Store.Get(avatarID)
|
||||
if err != nil {
|
||||
rest.SendErrorJSON(w, r, p.L, http.StatusBadRequest, err, "can't load avatar")
|
||||
@@ -138,18 +195,47 @@ func (p *Proxy) Handler(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}()
|
||||
|
||||
// io.ReadFull keeps reading until the buffer is full or EOF, so a Store
|
||||
// implementation that returns a buffered reader with a small first-Read size
|
||||
// won't cause DetectContentType to misclassify a real image. Short bodies
|
||||
// (avatars under sniffLen) return ErrUnexpectedEOF — that's expected, we sniff
|
||||
// what we got.
|
||||
buf := make([]byte, sniffLen)
|
||||
n, err := avReader.Read(buf)
|
||||
if err != nil && err != io.EOF {
|
||||
n, err := io.ReadFull(avReader, buf)
|
||||
if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF {
|
||||
p.Logf("[WARN] can't read from avatar reader for %s, %s", avatarID, err)
|
||||
rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "can't read avatar")
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Length", strconv.Itoa(size))
|
||||
contentType := http.DetectContentType(buf)
|
||||
if contentType == "application/octet-stream" {
|
||||
|
||||
// validate the bytes really are an image before declaring a content type. Even
|
||||
// though Put() now refuses to store non-image content, this catches stores poisoned
|
||||
// before the fix and any future regression — never trust the bytes at the store
|
||||
// boundary alone, validate again at serve time. An empty body (e.g. NoOp store)
|
||||
// is treated as a benign no-content case: nothing to render, no XSS surface.
|
||||
var contentType string
|
||||
if n > 0 {
|
||||
var ctErr error
|
||||
contentType, ctErr = safeImgContentType(buf[:n])
|
||||
if ctErr != nil {
|
||||
p.Logf("[WARN] rejecting non-image avatar %s: %v", avatarID, ctErr)
|
||||
rest.SendErrorJSON(w, r, p.L, http.StatusUnsupportedMediaType, ctErr, "invalid avatar content")
|
||||
return
|
||||
}
|
||||
} else {
|
||||
contentType = "image/*"
|
||||
}
|
||||
|
||||
// caching headers only after validation so error responses aren't cached
|
||||
etag := `"` + p.Store.ID(avatarID) + `"`
|
||||
w.Header().Set("Etag", etag)
|
||||
w.Header().Set("Cache-Control", "max-age=604800") // 7 days
|
||||
if match := r.Header.Get("If-None-Match"); match != "" && etagMatches(match, etag) {
|
||||
w.WriteHeader(http.StatusNotModified)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Length", strconv.Itoa(size))
|
||||
w.Header().Set("Content-Type", contentType)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
if _, err = w.Write(buf[:n]); err != nil {
|
||||
@@ -162,33 +248,57 @@ func (p *Proxy) Handler(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
// resize an image of supported format (PNG, JPG, GIF) to the size of "limit" px of the biggest side
|
||||
// (width or height) preserving aspect ratio.
|
||||
// Returns original reader if resizing is not needed or failed.
|
||||
func (p *Proxy) resize(reader io.Reader, limit int) io.Reader {
|
||||
if reader == nil {
|
||||
p.Logf("[WARN] avatar resize(): reader is nil")
|
||||
// resize validates that the input is a real image and, if needed, re-encodes it to
|
||||
// fit within "limit" px on the larger side preserving aspect ratio. Returns nil for
|
||||
// non-image content or for declared dimensions exceeding maxAvatarPixels so attacker
|
||||
// payloads (HTML/SVG/decompression bombs) never reach the store. With limit <= 0 or
|
||||
// when the image already fits, the original bytes are returned verbatim so animated
|
||||
// GIFs and other multi-frame formats round-trip without being flattened to one frame.
|
||||
//
|
||||
// Validation uses image.DecodeConfig (cheap — declares dimensions, allocates nothing)
|
||||
// before any full image.Decode, so a 100 KB compressed image declaring 65535x65535 px
|
||||
// is rejected without ever materializing the raster.
|
||||
//
|
||||
// Callers (load, PutContent, identicon generation) must ensure body is bounded by
|
||||
// maxAvatarFetchSize before calling; resize trusts the size invariant rather than
|
||||
// re-buffering. An empty body or a body over the cap is refused defensively.
|
||||
func (p *Proxy) resize(body []byte, limit int) io.Reader {
|
||||
if len(body) == 0 || int64(len(body)) > maxAvatarFetchSize {
|
||||
p.Logf("[WARN] avatar resize(): refusing body of size %d (cap %d)", len(body), maxAvatarFetchSize)
|
||||
return nil
|
||||
}
|
||||
if limit <= 0 {
|
||||
p.Logf("[DEBUG] avatar resize(): limit should be greater than 0")
|
||||
return reader
|
||||
}
|
||||
|
||||
var teeBuf bytes.Buffer
|
||||
tee := io.TeeReader(reader, &teeBuf)
|
||||
src, _, err := image.Decode(tee)
|
||||
// validate format and dimensions without allocating pixel memory.
|
||||
cfg, _, err := image.DecodeConfig(bytes.NewReader(body))
|
||||
if err != nil {
|
||||
// non-image input must never reach the store: refuse and let the caller
|
||||
// fall back to an identicon. Returning the raw bytes here previously let
|
||||
// an attacker who controlled u.Picture poison the store with HTML/SVG that
|
||||
// the Handler would later serve back with text/html content type.
|
||||
p.Logf("[WARN] avatar resize(): can't decode avatar image, %s", err)
|
||||
return &teeBuf
|
||||
return nil
|
||||
}
|
||||
// multiply in int64 — on 32-bit builds (GOARCH=386, 32-bit arm) the int
|
||||
// product of two 16-bit-or-larger dimensions can overflow and wrap below
|
||||
// maxAvatarPixels, bypassing the cap. GIF's 16-bit logical screen and
|
||||
// JPEG's 16-bit SOF dimensions both hit this if multiplied as int32.
|
||||
if cfg.Width <= 0 || cfg.Height <= 0 || int64(cfg.Width)*int64(cfg.Height) > int64(maxAvatarPixels) {
|
||||
p.Logf("[WARN] avatar resize(): declared dimensions %dx%d exceed safe limit", cfg.Width, cfg.Height)
|
||||
return nil
|
||||
}
|
||||
|
||||
bounds := src.Bounds()
|
||||
w, h := bounds.Dx(), bounds.Dy()
|
||||
if w <= limit && h <= limit || w <= 0 || h <= 0 {
|
||||
p.Logf("[DEBUG] resizing image is smaller that the limit or has 0 size")
|
||||
return &teeBuf
|
||||
if limit <= 0 || (cfg.Width <= limit && cfg.Height <= limit) {
|
||||
p.Logf("[DEBUG] avatar resize(): no resize needed (dim %dx%d, limit %d)", cfg.Width, cfg.Height, limit)
|
||||
return bytes.NewReader(body)
|
||||
}
|
||||
|
||||
// dimensions are bounded — full decode is now safe to allocate.
|
||||
src, _, err := image.Decode(bytes.NewReader(body))
|
||||
if err != nil {
|
||||
p.Logf("[WARN] avatar resize(): decode after dim-check failed, %s", err)
|
||||
return nil
|
||||
}
|
||||
w, h := src.Bounds().Dx(), src.Bounds().Dy()
|
||||
newW, newH := w*limit/h, limit
|
||||
if w > h {
|
||||
newW, newH = limit, h*limit/w
|
||||
@@ -200,11 +310,64 @@ func (p *Proxy) resize(reader io.Reader, limit int) io.Reader {
|
||||
var out bytes.Buffer
|
||||
if err = png.Encode(&out, m); err != nil {
|
||||
p.Logf("[WARN] avatar resize(): can't encode resized avatar to PNG, %s", err)
|
||||
return &teeBuf
|
||||
return bytes.NewReader(body) // fall back to the validated original
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
// safeImgContentType returns the sniffed content type if the bytes look like a safe
|
||||
// raster image format. The set is an explicit allowlist (PNG, JPEG, GIF, WebP, BMP,
|
||||
// ICO) — no HasPrefix("image/") catch-all — so future scriptable image/* MIME types
|
||||
// added by http.DetectContentType cannot silently pass. Returns an error otherwise.
|
||||
// image/svg+xml is excluded because SVG can execute scripts when navigated to top
|
||||
// level; image/* coverage of icon files uses both spellings http.DetectContentType
|
||||
// is known to return.
|
||||
func safeImgContentType(img []byte) (string, error) {
|
||||
ct := http.DetectContentType(img)
|
||||
base := ct
|
||||
if idx := strings.Index(base, ";"); idx >= 0 {
|
||||
base = strings.TrimSpace(base[:idx])
|
||||
}
|
||||
switch base {
|
||||
case "image/png", "image/jpeg", "image/gif", "image/webp", "image/bmp",
|
||||
"image/x-icon", "image/vnd.microsoft.icon":
|
||||
return base, nil
|
||||
}
|
||||
return "", fmt.Errorf("non-image content type %q", ct)
|
||||
}
|
||||
|
||||
// etagMatches reports whether the If-None-Match header value matches the response
|
||||
// ETag per RFC 7232: the header is a comma-separated list of opaque-tags (each in
|
||||
// double quotes), optionally weak-prefixed with W/. The wildcard "*" matches anything.
|
||||
// We deliberately ignore weak/strong distinction because avatar responses are static
|
||||
// per id — both forms identify the same resource.
|
||||
func etagMatches(header, etag string) bool {
|
||||
header = strings.TrimSpace(header)
|
||||
if header == "*" {
|
||||
return true
|
||||
}
|
||||
for tag := range strings.SplitSeq(header, ",") {
|
||||
tag = strings.TrimSpace(tag)
|
||||
tag = strings.TrimPrefix(tag, "W/")
|
||||
if tag == etag {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// setAvatarDefenseHeaders applies layered defense headers on every avatar response
|
||||
// (success, 304, or error). Each header survives content-type validation regressions,
|
||||
// browser sniffing, and top-level navigation:
|
||||
// - Content-Security-Policy: strict, with sandbox — blocks inline scripts/handlers
|
||||
// - X-Content-Type-Options: nosniff — prevents MIME-overriding the declared type
|
||||
// - Content-Disposition: inline; filename="avatar" — frames the response as a file
|
||||
func setAvatarDefenseHeaders(w http.ResponseWriter) {
|
||||
w.Header().Set("Content-Security-Policy", "default-src 'none'; sandbox; frame-ancestors 'none'")
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
w.Header().Set("Content-Disposition", `inline; filename="avatar"`)
|
||||
}
|
||||
|
||||
// GenerateAvatar for give user with identicon
|
||||
func GenerateAvatar(user string) ([]byte, error) {
|
||||
|
||||
@@ -243,7 +406,7 @@ func GetGravatarURL(email string) (res string, err error) {
|
||||
}
|
||||
|
||||
func retry(retries int, delay time.Duration, fn func() error) (err error) {
|
||||
for i := 0; i < retries; i++ {
|
||||
for range retries {
|
||||
if err = fn(); err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
+5
-5
@@ -6,17 +6,17 @@ import "log"
|
||||
|
||||
// L defined logger interface used everywhere in the package
|
||||
type L interface {
|
||||
Logf(format string, args ...interface{})
|
||||
Logf(format string, args ...any)
|
||||
}
|
||||
|
||||
// Func type is an adapter to allow the use of ordinary functions as Logger.
|
||||
type Func func(format string, args ...interface{})
|
||||
type Func func(format string, args ...any)
|
||||
|
||||
// Logf calls f(id)
|
||||
func (f Func) Logf(format string, args ...interface{}) { f(format, args...) }
|
||||
func (f Func) Logf(format string, args ...any) { f(format, args...) }
|
||||
|
||||
// NoOp logger
|
||||
var NoOp = Func(func(string, ...interface{}) {})
|
||||
var NoOp = Func(func(string, ...any) {})
|
||||
|
||||
// Std logger sends to std default logger directly
|
||||
var Std = Func(func(format string, args ...interface{}) { log.Printf(format, args...) })
|
||||
var Std = Func(func(format string, args ...any) { log.Printf(format, args...) })
|
||||
|
||||
+2
-2
@@ -57,7 +57,7 @@ type ErrorHandlerFunc func(w http.ResponseWriter, r *http.Request, statusCode in
|
||||
var adminUser = token.User{
|
||||
ID: "admin",
|
||||
Name: "admin",
|
||||
Attributes: map[string]interface{}{
|
||||
Attributes: map[string]any{
|
||||
"admin": true,
|
||||
},
|
||||
}
|
||||
@@ -243,7 +243,7 @@ func (a *Authenticator) basicAdminUser(r *http.Request) bool {
|
||||
|
||||
// using ConstantTimeCompare to avoid timing attack
|
||||
if user != "admin" || subtle.ConstantTimeCompare([]byte(passwd), []byte(a.AdminPasswd)) != 1 {
|
||||
a.Logf("[WARN] admin basic auth failed, user/passwd mismatch, %s:%s", user, passwd)
|
||||
a.Logf("[WARN] admin basic auth failed for user %q", user)
|
||||
return false
|
||||
}
|
||||
|
||||
|
||||
+41
-5
@@ -13,6 +13,7 @@ import (
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -76,7 +77,7 @@ type AppleConfig struct {
|
||||
ResponseMode string // changes method of receiving data in callback. Default value "form_post" (https://developer.apple.com/documentation/sign_in_with_apple/request_an_authorization_to_the_sign_in_with_apple_server?changes=_1_2#4066168)
|
||||
|
||||
scopes []string // for this package allow only username scope and UID in token claims. Apple service API provide only "email" and "name" scope values (https://developer.apple.com/documentation/sign_in_with_apple/clientconfigi/3230955-scope)
|
||||
privateKey interface{} // private key from Apple obtained in developer account (the keys section). Required for create the Client Secret (https://developer.apple.com/documentation/sign_in_with_apple/generate_and_validate_tokens#3262048)
|
||||
privateKey any // private key from Apple obtained in developer account (the keys section). Required for create the Client Secret (https://developer.apple.com/documentation/sign_in_with_apple/generate_and_validate_tokens#3262048)
|
||||
publicKey crypto.PublicKey // need for validate sign of token
|
||||
clientSecret string // is the JWT client secret will create after first call and then used until expired
|
||||
jwkURL string // URL for fetch JWK Apple keys, need redefine for tests
|
||||
@@ -228,7 +229,7 @@ func (ah *AppleHandler) initPrivateKey() error {
|
||||
}
|
||||
|
||||
// tokenKeyFunc use for verify JWT sign, it receives the parsed token and should return the key for validating.
|
||||
func (ah *AppleHandler) tokenKeyFunc(jwtToken *jwt.Token) (interface{}, error) {
|
||||
func (ah *AppleHandler) tokenKeyFunc(jwtToken *jwt.Token) (any, error) {
|
||||
if jwtToken == nil {
|
||||
return nil, fmt.Errorf("failed to call token keyFunc, because token is nil")
|
||||
}
|
||||
@@ -331,7 +332,7 @@ func (ah AppleHandler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
rest.SendErrorJSON(w, r, ah.L, http.StatusInternalServerError, err, "exchange failed")
|
||||
return
|
||||
}
|
||||
ah.Logf("[DEBUG] response data %+v", resp)
|
||||
ah.Logf("[DEBUG] apple exchange response: %s", appleVerificationResponseLogSummary(resp))
|
||||
if resp.Error != "" {
|
||||
rest.SendErrorJSON(w, r, ah.L, http.StatusInternalServerError, nil, fmt.Sprintf("fetch IDtoken response error: %s", resp.Error))
|
||||
return
|
||||
@@ -345,10 +346,22 @@ func (ah AppleHandler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// get token claims for extract uid (and email or name if they exist in scope)
|
||||
// get token claims for extract uid (and email or name if they exist in scope).
|
||||
// jwt v5 parser options enforce iss == https://appleid.apple.com and
|
||||
// aud == ClientID inline so we don't need a separate validate pass.
|
||||
tokenClaims := jwt.MapClaims{}
|
||||
_, err = jwt.ParseWithClaims(resp.IDToken, tokenClaims, keySet.keyFunc)
|
||||
_, err = jwt.ParseWithClaims(resp.IDToken, tokenClaims, keySet.keyFunc,
|
||||
jwt.WithIssuer(appleIDTokenIssuer),
|
||||
jwt.WithAudience(ah.conf.ClientID))
|
||||
if err != nil {
|
||||
// distinguish a confused-deputy reject (iss/aud) from a server-side
|
||||
// parse/sig failure so the handler returns the same 403 + body as
|
||||
// before for the security-relevant case.
|
||||
if errors.Is(err, jwt.ErrTokenInvalidIssuer) || errors.Is(err, jwt.ErrTokenInvalidAudience) {
|
||||
ah.Logf("[WARN] apple id_token rejected: %s", err.Error())
|
||||
rest.SendErrorJSON(w, r, ah.L, http.StatusForbidden, nil, "invalid id_token")
|
||||
return
|
||||
}
|
||||
ah.Logf("[ERROR] failed to get claims: " + err.Error())
|
||||
rest.SendErrorJSON(w, r, ah.L, http.StatusInternalServerError, nil, fmt.Sprintf("failed to token validation, key is invalid: %s", resp.Error))
|
||||
return
|
||||
@@ -549,3 +562,26 @@ func (ah AppleHandler) makeRedirURL(path string) string {
|
||||
|
||||
return strings.TrimRight(ah.URL, "/") + strings.TrimSuffix(newPath, "/") + urlCallbackSuffix
|
||||
}
|
||||
|
||||
// appleVerificationResponseLogSummary formats appleVerificationResponse for safe
|
||||
// logging. The struct's AccessToken, RefreshToken and IDToken fields are
|
||||
// credentials and must never appear verbatim in logs that may be shipped to
|
||||
// centralized logging or third-party observability systems; this helper logs
|
||||
// only their presence (present|missing), the non-secret token type and
|
||||
// expiry, plus any provider-side error string.
|
||||
func appleVerificationResponseLogSummary(r appleVerificationResponse) string {
|
||||
return fmt.Sprintf("type=%s expires_in=%d access_token=%s refresh_token=%s id_token=%s error=%q",
|
||||
r.TokenType, r.ExpiresIn,
|
||||
presence(r.AccessToken), presence(r.RefreshToken), presence(r.IDToken), r.Error)
|
||||
}
|
||||
|
||||
func presence(s string) string {
|
||||
if s == "" {
|
||||
return "missing"
|
||||
}
|
||||
return "present"
|
||||
}
|
||||
|
||||
// appleIDTokenIssuer is the issuer Apple sets on every id_token issued by Sign in with Apple.
|
||||
// see https://developer.apple.com/documentation/sign_in_with_apple/sign_in_with_apple_rest_api/verifying_a_user
|
||||
const appleIDTokenIssuer = "https://appleid.apple.com" // #nosec G101 -- public Apple issuer URL, not a credential
|
||||
|
||||
+1
-1
@@ -162,7 +162,7 @@ func (aks *appleKeySet) get(kid string) (keys *applePublicKey, err error) {
|
||||
}
|
||||
|
||||
// keyFunc use for JWT verify with specific public key
|
||||
func (aks *appleKeySet) keyFunc(token *jwt.Token) (interface{}, error) {
|
||||
func (aks *appleKeySet) keyFunc(token *jwt.Token) (any, error) {
|
||||
|
||||
keyID, ok := token.Header["kid"].(string)
|
||||
if !ok {
|
||||
|
||||
+10
-2
@@ -79,17 +79,25 @@ func (c *CustomServer) Run(ctx context.Context) {
|
||||
u, err := url.Parse(c.URL)
|
||||
if err != nil {
|
||||
c.Logf("[ERROR] failed to parse service base URL=%s", c.URL)
|
||||
c.lock.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
_, port, err := net.SplitHostPort(u.Host)
|
||||
host, port, err := net.SplitHostPort(u.Host)
|
||||
if err != nil {
|
||||
c.Logf("[ERROR] failed to get port from URL=%s", c.URL)
|
||||
c.lock.Unlock()
|
||||
return
|
||||
}
|
||||
// hostname from URL is honored; only an explicit non-loopback host
|
||||
// (e.g. "0.0.0.0" baked into c.URL) binds beyond loopback. Empty/
|
||||
// "localhost" falls through to localBindAddr's 127.0.0.1 default.
|
||||
if host == "localhost" {
|
||||
host = ""
|
||||
}
|
||||
|
||||
c.httpServer = &http.Server{
|
||||
Addr: fmt.Sprintf(":%s", port),
|
||||
Addr: localBindAddr(host, port),
|
||||
ReadHeaderTimeout: 5 * time.Second,
|
||||
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
|
||||
+1
-1
@@ -57,7 +57,7 @@ func (d *DevAuthServer) Run(ctx context.Context) { // nolint (gocyclo)
|
||||
}
|
||||
|
||||
d.httpServer = &http.Server{
|
||||
Addr: fmt.Sprintf(":%d", d.Provider.Port),
|
||||
Addr: localBindAddr(d.Provider.Host, fmt.Sprintf("%d", d.Provider.Port)),
|
||||
ReadHeaderTimeout: 5 * time.Second,
|
||||
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
d.Logf("[DEBUG] dev oauth request %s %s %+v", r.Method, r.URL, r.Header)
|
||||
|
||||
+1
-1
@@ -122,7 +122,7 @@ func (h Oauth1Handler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
jData := map[string]interface{}{}
|
||||
jData := map[string]any{}
|
||||
if e := json.Unmarshal(data, &jData); e != nil {
|
||||
rest.SendErrorJSON(w, r, h.L, http.StatusInternalServerError, err, "failed to unmarshal user info")
|
||||
return
|
||||
|
||||
+2
-2
@@ -57,7 +57,7 @@ type Params struct {
|
||||
}
|
||||
|
||||
// UserData is type for user information returned from oauth2 providers /info API method
|
||||
type UserData map[string]interface{}
|
||||
type UserData map[string]any
|
||||
|
||||
// Value returns value for key or empty string if not found
|
||||
func (u UserData) Value(key string) string {
|
||||
@@ -197,7 +197,7 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
jData := map[string]interface{}{}
|
||||
jData := map[string]any{}
|
||||
if e := json.Unmarshal(data, &jData); e != nil {
|
||||
rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to unmarshal user info")
|
||||
return
|
||||
|
||||
+6
-2
@@ -79,9 +79,13 @@ func NewEmailClient(emailParams EmailParams, l logger.L) *Email {
|
||||
return &Email{EmailParams: emailParams, L: l, sender: sender}
|
||||
}
|
||||
|
||||
// Send email with given text
|
||||
// Send email with given text. The body is not logged: confirmation emails
|
||||
// sent by the verify provider contain a one-shot magic-link token, and any
|
||||
// party with log access could redeem it before the user does. Logging only
|
||||
// the recipient and body length keeps the line useful for operators
|
||||
// without leaking the credential.
|
||||
func (e *Email) Send(to, text string) error {
|
||||
e.Logf("[DEBUG] send %q to %s", text, to)
|
||||
e.Logf("[DEBUG] send %d-byte message to %s", len(text), to)
|
||||
return e.sender.Send(text, email.Params{
|
||||
From: e.From,
|
||||
To: []string{to},
|
||||
|
||||
+15
@@ -4,6 +4,7 @@ import (
|
||||
"crypto/rand"
|
||||
"crypto/sha1"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
@@ -11,6 +12,20 @@ import (
|
||||
"github.com/go-pkgz/auth/v2/token"
|
||||
)
|
||||
|
||||
// localBindAddr returns the listen address for the dev oauth and custom-server
|
||||
// helpers. Both servers are intended for local development and embedded
|
||||
// flows; the historical default ":port" listened on every interface, which
|
||||
// silently exposed the dev OAuth UI to anyone on the LAN. Default the bind
|
||||
// to 127.0.0.1; callers that explicitly want a non-loopback bind can pass
|
||||
// a hostname (e.g. "0.0.0.0" or a specific IP) via Provider.Host (dev) or
|
||||
// the URL host (custom-server).
|
||||
func localBindAddr(host, port string) string {
|
||||
if host == "" {
|
||||
host = "127.0.0.1"
|
||||
}
|
||||
return net.JoinHostPort(host, port)
|
||||
}
|
||||
|
||||
const (
|
||||
urlLoginSuffix = "/login"
|
||||
urlCallbackSuffix = "/callback"
|
||||
|
||||
+126
-8
@@ -3,13 +3,16 @@ package provider
|
||||
//go:generate moq --out telegram_moq_test.go . TelegramAPI
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
neturl "net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -19,6 +22,7 @@ import (
|
||||
"github.com/go-pkgz/rest"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
|
||||
"github.com/go-pkgz/auth/v2/avatar"
|
||||
"github.com/go-pkgz/auth/v2/logger"
|
||||
authtoken "github.com/go-pkgz/auth/v2/token"
|
||||
)
|
||||
@@ -186,11 +190,18 @@ func (th *TelegramHandler) processUpdates(ctx context.Context, updates *telegram
|
||||
|
||||
id := th.ProviderName + "_" + authtoken.HashID(sha1.New(), fmt.Sprint(update.Message.Chat.ID))
|
||||
|
||||
// avatarURL embeds the bot token in its path
|
||||
// (https://api.telegram.org/file/bot{TOKEN}/...). Never store it in
|
||||
// User.Picture: it would leak through avatar.Proxy.Put logs and, when
|
||||
// no avatar saver is configured, into the JWT and on to the client.
|
||||
// Fetch the bytes here and hand them to the avatar store directly.
|
||||
picture := th.saveTelegramAvatar(ctx, id, avatarURL)
|
||||
|
||||
authRequest.confirmed = true
|
||||
authRequest.user = &authtoken.User{
|
||||
ID: id,
|
||||
Name: update.Message.Chat.Name,
|
||||
Picture: avatarURL,
|
||||
Picture: picture,
|
||||
}
|
||||
|
||||
th.requests.Lock()
|
||||
@@ -294,10 +305,19 @@ func (th *TelegramHandler) LoginHandler(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
u, err := setAvatar(th.AvatarSaver, *authUser, &http.Client{Timeout: 5 * time.Second})
|
||||
if err != nil {
|
||||
rest.SendErrorJSON(w, r, th.L, http.StatusInternalServerError, err, "failed to save avatar to proxy")
|
||||
return
|
||||
// when saveTelegramAvatar already populated Picture with a local proxy
|
||||
// URL, skip the URL-fetching avatar pipeline. Letting setAvatar run
|
||||
// here would have it call Proxy.Put which re-fetches Picture; in
|
||||
// split-DNS / unreachable-internal-Opts.URL deployments that fetch
|
||||
// fails and the identicon fallback would silently overwrite the
|
||||
// stored Telegram bytes with an identicon at the same store path.
|
||||
u := *authUser
|
||||
if u.Picture == "" {
|
||||
u, err = setAvatar(th.AvatarSaver, *authUser, &http.Client{Timeout: 5 * time.Second})
|
||||
if err != nil {
|
||||
rest.SendErrorJSON(w, r, th.L, http.StatusInternalServerError, err, "failed to save avatar to proxy")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
claims := authtoken.Claims{
|
||||
@@ -449,18 +469,18 @@ func (tg *tgAPI) BotInfo(ctx context.Context) (*botInfo, error) {
|
||||
return resp.Result, nil
|
||||
}
|
||||
|
||||
func (tg *tgAPI) request(ctx context.Context, method string, data interface{}) error {
|
||||
func (tg *tgAPI) request(ctx context.Context, method string, data any) error {
|
||||
return repeater.NewFixed(3, time.Millisecond*50).Do(ctx, func() error {
|
||||
url := fmt.Sprintf("https://api.telegram.org/bot%s/%s", tg.token, method)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", url, http.NoBody)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create request: %w", err)
|
||||
return fmt.Errorf("failed to create request: %w", redactBotURLInErr(err))
|
||||
}
|
||||
|
||||
resp, err := tg.client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to send request: %w", err)
|
||||
return fmt.Errorf("failed to send request: %w", redactBotURLInErr(err))
|
||||
}
|
||||
defer resp.Body.Close() //nolint gosec // we don't care about response body
|
||||
|
||||
@@ -485,3 +505,101 @@ func (tg *tgAPI) parseError(r io.Reader, statusCode int) error {
|
||||
}
|
||||
return fmt.Errorf("unexpected telegram API status code %d, error: %q", statusCode, tgErr.Description)
|
||||
}
|
||||
|
||||
// avatarContentSaver matches the optional method on AvatarSaver implementations
|
||||
// that can store already-fetched bytes (avatar.Proxy provides one). Used by the
|
||||
// Telegram provider to avoid passing a bot-token-bearing URL through the
|
||||
// URL-fetching avatar pipeline.
|
||||
type avatarContentSaver interface {
|
||||
PutContent(userID string, content io.Reader) (string, error)
|
||||
}
|
||||
|
||||
// saveTelegramAvatar fetches the avatar bytes from a bot-token-bearing Telegram
|
||||
// URL and stores them via th.AvatarSaver, returning a clean local proxy URL.
|
||||
// The bot URL is consumed entirely inside this function so it never reaches
|
||||
// User.Picture, JWT claims, or any debug log of the user object. Returns ""
|
||||
// when the avatar cannot be saved (no URL, no compatible saver, or fetch
|
||||
// failure) — the caller treats that as "no picture" and the avatar pipeline
|
||||
// falls back to identicon as usual.
|
||||
func (th *TelegramHandler) saveTelegramAvatar(ctx context.Context, userID, avatarURL string) string {
|
||||
if avatarURL == "" {
|
||||
return ""
|
||||
}
|
||||
// guard against typed-nil *avatar.Proxy. auth.go skips initializing
|
||||
// res.avatarProxy when Opts.AvatarStore is unset, so AvatarSaver can be
|
||||
// a non-nil interface wrapping a nil *avatar.Proxy. The type assertion
|
||||
// below would still succeed (interface satisfaction is structural), but
|
||||
// PutContent on a nil receiver panics on the first p.Store deref.
|
||||
if th.AvatarSaver == nil || th.AvatarSaver == (*avatar.Proxy)(nil) {
|
||||
th.Logf("[WARN] telegram avatar dropped: AvatarSaver is not configured")
|
||||
return ""
|
||||
}
|
||||
saver, ok := th.AvatarSaver.(avatarContentSaver)
|
||||
if !ok {
|
||||
// fallback intentionally drops the picture rather than expose the bot
|
||||
// token; warn so operators can wire a content-aware saver if they want
|
||||
// telegram avatars saved
|
||||
th.Logf("[WARN] telegram avatar dropped: configured AvatarSaver does not support direct content save")
|
||||
return ""
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, avatarURL, http.NoBody)
|
||||
if err != nil {
|
||||
th.Logf("[WARN] telegram avatar fetch request build failed: %v", redactBotURLInErr(err))
|
||||
return ""
|
||||
}
|
||||
client := &http.Client{Timeout: 5 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
th.Logf("[WARN] telegram avatar fetch failed: %v", redactBotURLInErr(err))
|
||||
return ""
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
th.Logf("[WARN] telegram avatar fetch returned status %d", resp.StatusCode)
|
||||
return ""
|
||||
}
|
||||
// cap body size to protect PutContent from an unbounded upstream response.
|
||||
// Telegram caps photos at 5 MiB; 10 MiB is generous headroom while still
|
||||
// bounding worst-case memory.
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxTelegramAvatarSize+1))
|
||||
if err != nil {
|
||||
th.Logf("[WARN] telegram avatar read failed: %v", err)
|
||||
return ""
|
||||
}
|
||||
if int64(len(body)) > maxTelegramAvatarSize {
|
||||
th.Logf("[WARN] telegram avatar dropped: body exceeds %d bytes", maxTelegramAvatarSize)
|
||||
return ""
|
||||
}
|
||||
picture, err := saver.PutContent(userID, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
th.Logf("[WARN] telegram avatar save failed: %v", err)
|
||||
return ""
|
||||
}
|
||||
return picture
|
||||
}
|
||||
|
||||
const maxTelegramAvatarSize = 10 << 20
|
||||
|
||||
// botTokenInURLPath matches the bot-token segment of a Telegram URL anchored
|
||||
// between path slashes ("/botTOKEN/..."). The leading and trailing slashes
|
||||
// avoid matching unrelated identifiers that happen to start with "bot" (e.g.
|
||||
// the username "botFather" appearing elsewhere in a log line). Replacement
|
||||
// preserves the slashes via "/bot<redacted>/" to keep surrounding URL
|
||||
// structure intact for diagnostics.
|
||||
var botTokenInURLPath = regexp.MustCompile(`/bot[A-Za-z0-9:_-]+/`)
|
||||
|
||||
// redactBotURLInErr returns the error with any embedded Telegram bot-token
|
||||
// segment in URL paths replaced by "bot<redacted>". net/http's *url.Error
|
||||
// stringifies as `Op "URL": Err`, so a transport failure on a URL like
|
||||
// https://api.telegram.org/file/bot<TOKEN>/... otherwise prints the token
|
||||
// verbatim.
|
||||
func redactBotURLInErr(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
redacted := botTokenInURLPath.ReplaceAllString(err.Error(), "/bot<redacted>/")
|
||||
if redacted == err.Error() {
|
||||
return err
|
||||
}
|
||||
return errors.New(redacted)
|
||||
}
|
||||
|
||||
+165
-1
@@ -3,10 +3,13 @@ package provider
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha1"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"html/template"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/go-pkgz/rest"
|
||||
@@ -19,6 +22,18 @@ import (
|
||||
|
||||
// VerifyHandler implements non-oauth2 provider authorizing users with some confirmation.
|
||||
// can be email, IM or anything else implementing Sender interface
|
||||
//
|
||||
// Identity caveat: the local user id returned to the application is derived
|
||||
// from the verified address (ProviderName + "_" + HashID(address)). The
|
||||
// confirmation round-trip proves current control of the address at login
|
||||
// time; it does not guarantee a stable+unique identity over time. The owner
|
||||
// of an address can change without the address changing — employer
|
||||
// offboarding, lapsed free-mail accounts, and recycled domains all hand
|
||||
// control of an address to the next person who claims it. Integrators that
|
||||
// need stable identity should map the verified address to a server-side
|
||||
// immutable user id at first successful verify and key their records on
|
||||
// that id, not on the value returned here. See the "Email-as-identity
|
||||
// caveat" section of the README for guidance.
|
||||
type VerifyHandler struct {
|
||||
logger.L
|
||||
ProviderName string
|
||||
@@ -38,6 +53,112 @@ type VerifyHandler struct {
|
||||
// here. Nil disables validation and preserves legacy permissive
|
||||
// behavior — any non-empty "from" value is honored.
|
||||
AllowedRedirectHosts token.AllowedHosts
|
||||
|
||||
// ConfirmationStore enforces one-shot consumption of confirmation tokens.
|
||||
// When non-nil, a token cannot be redeemed twice within its TTL window.
|
||||
// Leave nil to keep the legacy behavior (token replayable until expiry).
|
||||
ConfirmationStore VerifConfirmationStore
|
||||
}
|
||||
|
||||
// VerifConfirmationStore tracks consumed confirmation tokens to prevent replay.
|
||||
// Implementations must be safe for concurrent use.
|
||||
type VerifConfirmationStore interface {
|
||||
// MarkUsed records key as consumed and returns alreadyUsed=true if it was
|
||||
// already recorded. The implementation MUST retain the marker for at
|
||||
// least the supplied ttl, or return a non-nil err if it cannot --
|
||||
// dropping a marker before its ttl while the underlying JWT is still
|
||||
// valid reopens the replay window the store is meant to close. err
|
||||
// signals a backend failure (network, disk, capacity, etc.); callers
|
||||
// MUST treat a non-nil err as fail-closed (reject the redemption).
|
||||
//
|
||||
// Adapter authors: do NOT embed key (or any caller-supplied data) in
|
||||
// returned errors. The handler logs err on the fail-closed branch, and
|
||||
// although key is the SHA-256 of the raw token rather than the token
|
||||
// itself, it still uniquely identifies the live, unredeemed JWT in
|
||||
// log destinations. Wrap the underlying backend error with a generic
|
||||
// description (e.g. "redis SET failed: %w") instead.
|
||||
MarkUsed(key string, ttl time.Duration) (alreadyUsed bool, err error)
|
||||
}
|
||||
|
||||
// VerifConfirmationStoreFunc is an adapter to use ordinary functions as
|
||||
// VerifConfirmationStore, mirroring the SenderFunc / token.AllowedHostsFunc
|
||||
// house pattern for closure-based config.
|
||||
type VerifConfirmationStoreFunc func(key string, ttl time.Duration) (alreadyUsed bool, err error)
|
||||
|
||||
// MarkUsed calls f(key, ttl) to implement VerifConfirmationStore.
|
||||
func (f VerifConfirmationStoreFunc) MarkUsed(key string, ttl time.Duration) (bool, error) {
|
||||
return f(key, ttl)
|
||||
}
|
||||
|
||||
// NewInMemoryVerifStore returns a process-local default VerifConfirmationStore.
|
||||
// Suitable for single-instance deployments. Multi-instance deployments behind
|
||||
// a load balancer MUST supply a shared backend (e.g. Redis) -- otherwise an
|
||||
// attacker who lands on a different instance from the legitimate user can
|
||||
// replay the token there. The default's failure is silent: the request
|
||||
// completes normally and no log indicates the protection was bypassed.
|
||||
func NewInMemoryVerifStore() VerifConfirmationStore {
|
||||
return &inMemoryVerifStore{used: make(map[string]time.Time)}
|
||||
}
|
||||
|
||||
type inMemoryVerifStore struct {
|
||||
mu sync.Mutex
|
||||
used map[string]time.Time // key -> expiry
|
||||
insertCount int
|
||||
}
|
||||
|
||||
// inMemoryVerifStoreSweepEvery is the in-memory store's amortization cadence.
|
||||
// Walking the whole map on every redemption is O(n) under a single mutex,
|
||||
// which serializes the hot path. Sweeping every N inserts keeps the map size
|
||||
// bounded by ~N + (concurrent redemptions during the gap) without holding
|
||||
// the lock through a full walk on most calls. Declared as a var rather than
|
||||
// a const so tests can lower it to exercise the sweep branch.
|
||||
var inMemoryVerifStoreSweepEvery = 256
|
||||
|
||||
func (s *inMemoryVerifStore) MarkUsed(key string, ttl time.Duration) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
now := time.Now()
|
||||
if exp, ok := s.used[key]; ok && exp.After(now) {
|
||||
return true, nil
|
||||
}
|
||||
// amortized eviction: walk the map only every Nth insert, not on every
|
||||
// hot-path call. The lookup above already rejects unexpired duplicates,
|
||||
// so worst-case staleness is bounded by N inserts between sweeps.
|
||||
s.insertCount++
|
||||
if s.insertCount >= inMemoryVerifStoreSweepEvery {
|
||||
s.insertCount = 0
|
||||
for k, exp := range s.used {
|
||||
if !exp.After(now) {
|
||||
delete(s.used, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
s.used[key] = now.Add(ttl)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// confirmationKey hashes the raw token so the store key length is bounded
|
||||
// regardless of token size, and so the in-memory map doesn't retain the
|
||||
// signed token itself.
|
||||
func confirmationKey(rawToken string) string {
|
||||
sum := sha256.Sum256([]byte(rawToken))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// scrubTokenFromRequest returns a shallow clone of r with the "token" query
|
||||
// parameter replaced by "<redacted>". rest.SendErrorJSON logs r.URL, and the
|
||||
// fail-closed branches in LoginHandler fire while the confirmation JWT is
|
||||
// still live (store didn't record consumption) -- a single log line equals
|
||||
// an unredeemed magic link without this scrub.
|
||||
func scrubTokenFromRequest(r *http.Request) *http.Request {
|
||||
if r == nil || r.URL == nil || r.URL.Query().Get("token") == "" {
|
||||
return r
|
||||
}
|
||||
rc := r.Clone(r.Context())
|
||||
q := rc.URL.Query()
|
||||
q.Set("token", "<redacted>")
|
||||
rc.URL.RawQuery = q.Encode()
|
||||
return rc
|
||||
}
|
||||
|
||||
// Sender defines interface to send emails
|
||||
@@ -66,7 +187,13 @@ type VerifTokenService interface {
|
||||
func (e VerifyHandler) Name() string { return e.ProviderName }
|
||||
|
||||
// LoginHandler gets name and address from query, makes confirmation token and sends it to user.
|
||||
// In case if confirmation token presented in the query uses it to create auth token
|
||||
// In case if confirmation token presented in the query uses it to create auth token.
|
||||
//
|
||||
// Consumption is final when ConfirmationStore is configured: the token is
|
||||
// marked used before any further side effects (avatar fetch, token issuance),
|
||||
// so a transient downstream failure burns the token and the user must request
|
||||
// a new confirmation email rather than retry the same link. This trade-off
|
||||
// keeps the replay check atomic with the security boundary.
|
||||
func (e VerifyHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
// GET /login?site=site&user=name&address=someone@example.com
|
||||
@@ -89,6 +216,36 @@ func (e VerifyHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
store := e.ConfirmationStore
|
||||
// guard against a typed-nil VerifConfirmationStoreFunc: a non-nil
|
||||
// interface wrapping a nil func survives the != nil check above and
|
||||
// would panic at MarkUsed. Treat it as no store configured. Mirrors
|
||||
// the AllowedHostsFunc nil-guard in token/jwt.go.
|
||||
if fn, ok := store.(VerifConfirmationStoreFunc); ok && fn == nil {
|
||||
store = nil
|
||||
}
|
||||
if store != nil {
|
||||
ttl := time.Minute
|
||||
if confClaims.ExpiresAt != nil {
|
||||
if remaining := time.Until(confClaims.ExpiresAt.Time); remaining > 0 {
|
||||
ttl = remaining
|
||||
}
|
||||
}
|
||||
alreadyUsed, markErr := store.MarkUsed(confirmationKey(tkn), ttl)
|
||||
if markErr != nil {
|
||||
// fail-closed: a backend outage must not let attackers replay
|
||||
// tokens. Reject with the token scrubbed from the logged URL,
|
||||
// since on this branch the store did NOT record consumption so
|
||||
// the JWT in the URL is still live.
|
||||
rest.SendErrorJSON(w, scrubTokenFromRequest(r), e.L, http.StatusForbidden, markErr, "confirmation token store unavailable")
|
||||
return
|
||||
}
|
||||
if alreadyUsed {
|
||||
rest.SendErrorJSON(w, scrubTokenFromRequest(r), e.L, http.StatusForbidden, fmt.Errorf("token already used"), "confirmation token already consumed")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
elems := strings.Split(confClaims.Handshake.ID, "::")
|
||||
if len(elems) != 2 {
|
||||
rest.SendErrorJSON(w, r, e.L, http.StatusBadRequest, fmt.Errorf("%s", confClaims.Handshake.ID), "invalid handshake token")
|
||||
@@ -162,6 +319,13 @@ func (e VerifyHandler) sendConfirmation(w http.ResponseWriter, r *http.Request)
|
||||
Handshake: &token.Handshake{
|
||||
State: "",
|
||||
ID: user + "::" + address,
|
||||
// without copying "from" here the redirect validator at the
|
||||
// other end has nothing to validate or to redirect to. The
|
||||
// docs (and #275) advertise ?from=<url> on the verify login
|
||||
// path, but the original sendConfirmation never put it on
|
||||
// the handshake JWT, so production verify flows could never
|
||||
// honor from at all.
|
||||
From: r.URL.Query().Get("from"),
|
||||
},
|
||||
SessionOnly: r.URL.Query().Get("session") != "" && r.URL.Query().Get("session") != "0",
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
|
||||
+7
-7
@@ -28,16 +28,16 @@ type User struct {
|
||||
Audience string `json:"aud,omitempty"`
|
||||
|
||||
// set by client
|
||||
IP string `json:"ip,omitempty"`
|
||||
Email string `json:"email,omitempty"`
|
||||
Attributes map[string]interface{} `json:"attrs,omitempty"`
|
||||
Role string `json:"role,omitempty"`
|
||||
IP string `json:"ip,omitempty"`
|
||||
Email string `json:"email,omitempty"`
|
||||
Attributes map[string]any `json:"attrs,omitempty"`
|
||||
Role string `json:"role,omitempty"`
|
||||
}
|
||||
|
||||
// SetBoolAttr sets boolean attribute
|
||||
func (u *User) SetBoolAttr(key string, val bool) {
|
||||
if u.Attributes == nil {
|
||||
u.Attributes = map[string]interface{}{}
|
||||
u.Attributes = map[string]any{}
|
||||
}
|
||||
u.Attributes[key] = val
|
||||
}
|
||||
@@ -45,7 +45,7 @@ func (u *User) SetBoolAttr(key string, val bool) {
|
||||
// SetStrAttr sets string attribute
|
||||
func (u *User) SetStrAttr(key, val string) {
|
||||
if u.Attributes == nil {
|
||||
u.Attributes = map[string]interface{}{}
|
||||
u.Attributes = map[string]any{}
|
||||
}
|
||||
u.Attributes[key] = val
|
||||
}
|
||||
@@ -100,7 +100,7 @@ func (u *User) SliceAttr(key string) []string {
|
||||
// SetSliceAttr sets slice attribute for given key
|
||||
func (u *User) SetSliceAttr(key string, val []string) {
|
||||
if u.Attributes == nil {
|
||||
u.Attributes = map[string]interface{}{}
|
||||
u.Attributes = map[string]any{}
|
||||
}
|
||||
u.Attributes[key] = val
|
||||
}
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
* -text
|
||||
*.bin -text -diff
|
||||
*.md text eol=lf
|
||||
|
||||
+700
-700
File diff suppressed because it is too large
Load Diff
+78
-78
@@ -1,79 +1,79 @@
|
||||
# Finite State Entropy
|
||||
|
||||
This package provides Finite State Entropy encoding and decoding.
|
||||
|
||||
Finite State Entropy (also referenced as [tANS](https://en.wikipedia.org/wiki/Asymmetric_numeral_systems#tANS))
|
||||
encoding provides a fast near-optimal symbol encoding/decoding
|
||||
for byte blocks as implemented in [zstandard](https://github.com/facebook/zstd).
|
||||
|
||||
This can be used for compressing input with a lot of similar input values to the smallest number of bytes.
|
||||
This does not perform any multi-byte [dictionary coding](https://en.wikipedia.org/wiki/Dictionary_coder) as LZ coders,
|
||||
but it can be used as a secondary step to compressors (like Snappy) that does not do entropy encoding.
|
||||
|
||||
* [Godoc documentation](https://godoc.org/github.com/klauspost/compress/fse)
|
||||
|
||||
## News
|
||||
|
||||
* Feb 2018: First implementation released. Consider this beta software for now.
|
||||
|
||||
# Usage
|
||||
|
||||
This package provides a low level interface that allows to compress single independent blocks.
|
||||
|
||||
Each block is separate, and there is no built in integrity checks.
|
||||
This means that the caller should keep track of block sizes and also do checksums if needed.
|
||||
|
||||
Compressing a block is done via the [`Compress`](https://godoc.org/github.com/klauspost/compress/fse#Compress) function.
|
||||
You must provide input and will receive the output and maybe an error.
|
||||
|
||||
These error values can be returned:
|
||||
|
||||
| Error | Description |
|
||||
|---------------------|-----------------------------------------------------------------------------|
|
||||
| `<nil>` | Everything ok, output is returned |
|
||||
| `ErrIncompressible` | Returned when input is judged to be too hard to compress |
|
||||
| `ErrUseRLE` | Returned from the compressor when the input is a single byte value repeated |
|
||||
| `(error)` | An internal error occurred. |
|
||||
|
||||
As can be seen above there are errors that will be returned even under normal operation so it is important to handle these.
|
||||
|
||||
To reduce allocations you can provide a [`Scratch`](https://godoc.org/github.com/klauspost/compress/fse#Scratch) object
|
||||
that can be re-used for successive calls. Both compression and decompression accepts a `Scratch` object, and the same
|
||||
object can be used for both.
|
||||
|
||||
Be aware, that when re-using a `Scratch` object that the *output* buffer is also re-used, so if you are still using this
|
||||
you must set the `Out` field in the scratch to nil. The same buffer is used for compression and decompression output.
|
||||
|
||||
Decompressing is done by calling the [`Decompress`](https://godoc.org/github.com/klauspost/compress/fse#Decompress) function.
|
||||
You must provide the output from the compression stage, at exactly the size you got back. If you receive an error back
|
||||
your input was likely corrupted.
|
||||
|
||||
It is important to note that a successful decoding does *not* mean your output matches your original input.
|
||||
There are no integrity checks, so relying on errors from the decompressor does not assure your data is valid.
|
||||
|
||||
For more detailed usage, see examples in the [godoc documentation](https://godoc.org/github.com/klauspost/compress/fse#pkg-examples).
|
||||
|
||||
# Performance
|
||||
|
||||
A lot of factors are affecting speed. Block sizes and compressibility of the material are primary factors.
|
||||
All compression functions are currently only running on the calling goroutine so only one core will be used per block.
|
||||
|
||||
The compressor is significantly faster if symbols are kept as small as possible. The highest byte value of the input
|
||||
is used to reduce some of the processing, so if all your input is above byte value 64 for instance, it may be
|
||||
beneficial to transpose all your input values down by 64.
|
||||
|
||||
With moderate block sizes around 64k speed are typically 200MB/s per core for compression and
|
||||
around 300MB/s decompression speed.
|
||||
|
||||
The same hardware typically does Huffman (deflate) encoding at 125MB/s and decompression at 100MB/s.
|
||||
|
||||
# Plans
|
||||
|
||||
At one point, more internals will be exposed to facilitate more "expert" usage of the components.
|
||||
|
||||
A streaming interface is also likely to be implemented. Likely compatible with [FSE stream format](https://github.com/Cyan4973/FiniteStateEntropy/blob/dev/programs/fileio.c#L261).
|
||||
|
||||
# Contributing
|
||||
|
||||
Contributions are always welcome. Be aware that adding public functions will require good justification and breaking
|
||||
# Finite State Entropy
|
||||
|
||||
This package provides Finite State Entropy encoding and decoding.
|
||||
|
||||
Finite State Entropy (also referenced as [tANS](https://en.wikipedia.org/wiki/Asymmetric_numeral_systems#tANS))
|
||||
encoding provides a fast near-optimal symbol encoding/decoding
|
||||
for byte blocks as implemented in [zstandard](https://github.com/facebook/zstd).
|
||||
|
||||
This can be used for compressing input with a lot of similar input values to the smallest number of bytes.
|
||||
This does not perform any multi-byte [dictionary coding](https://en.wikipedia.org/wiki/Dictionary_coder) as LZ coders,
|
||||
but it can be used as a secondary step to compressors (like Snappy) that does not do entropy encoding.
|
||||
|
||||
* [Godoc documentation](https://godoc.org/github.com/klauspost/compress/fse)
|
||||
|
||||
## News
|
||||
|
||||
* Feb 2018: First implementation released. Consider this beta software for now.
|
||||
|
||||
# Usage
|
||||
|
||||
This package provides a low level interface that allows to compress single independent blocks.
|
||||
|
||||
Each block is separate, and there is no built in integrity checks.
|
||||
This means that the caller should keep track of block sizes and also do checksums if needed.
|
||||
|
||||
Compressing a block is done via the [`Compress`](https://godoc.org/github.com/klauspost/compress/fse#Compress) function.
|
||||
You must provide input and will receive the output and maybe an error.
|
||||
|
||||
These error values can be returned:
|
||||
|
||||
| Error | Description |
|
||||
|---------------------|-----------------------------------------------------------------------------|
|
||||
| `<nil>` | Everything ok, output is returned |
|
||||
| `ErrIncompressible` | Returned when input is judged to be too hard to compress |
|
||||
| `ErrUseRLE` | Returned from the compressor when the input is a single byte value repeated |
|
||||
| `(error)` | An internal error occurred. |
|
||||
|
||||
As can be seen above there are errors that will be returned even under normal operation so it is important to handle these.
|
||||
|
||||
To reduce allocations you can provide a [`Scratch`](https://godoc.org/github.com/klauspost/compress/fse#Scratch) object
|
||||
that can be re-used for successive calls. Both compression and decompression accepts a `Scratch` object, and the same
|
||||
object can be used for both.
|
||||
|
||||
Be aware, that when re-using a `Scratch` object that the *output* buffer is also re-used, so if you are still using this
|
||||
you must set the `Out` field in the scratch to nil. The same buffer is used for compression and decompression output.
|
||||
|
||||
Decompressing is done by calling the [`Decompress`](https://godoc.org/github.com/klauspost/compress/fse#Decompress) function.
|
||||
You must provide the output from the compression stage, at exactly the size you got back. If you receive an error back
|
||||
your input was likely corrupted.
|
||||
|
||||
It is important to note that a successful decoding does *not* mean your output matches your original input.
|
||||
There are no integrity checks, so relying on errors from the decompressor does not assure your data is valid.
|
||||
|
||||
For more detailed usage, see examples in the [godoc documentation](https://godoc.org/github.com/klauspost/compress/fse#pkg-examples).
|
||||
|
||||
# Performance
|
||||
|
||||
A lot of factors are affecting speed. Block sizes and compressibility of the material are primary factors.
|
||||
All compression functions are currently only running on the calling goroutine so only one core will be used per block.
|
||||
|
||||
The compressor is significantly faster if symbols are kept as small as possible. The highest byte value of the input
|
||||
is used to reduce some of the processing, so if all your input is above byte value 64 for instance, it may be
|
||||
beneficial to transpose all your input values down by 64.
|
||||
|
||||
With moderate block sizes around 64k speed are typically 200MB/s per core for compression and
|
||||
around 300MB/s decompression speed.
|
||||
|
||||
The same hardware typically does Huffman (deflate) encoding at 125MB/s and decompression at 100MB/s.
|
||||
|
||||
# Plans
|
||||
|
||||
At one point, more internals will be exposed to facilitate more "expert" usage of the components.
|
||||
|
||||
A streaming interface is also likely to be implemented. Likely compatible with [FSE stream format](https://github.com/Cyan4973/FiniteStateEntropy/blob/dev/programs/fileio.c#L261).
|
||||
|
||||
# Contributing
|
||||
|
||||
Contributions are always welcome. Be aware that adding public functions will require good justification and breaking
|
||||
changes will likely not be accepted. If in doubt open an issue before writing the PR.
|
||||
+89
-89
@@ -1,89 +1,89 @@
|
||||
# Huff0 entropy compression
|
||||
|
||||
This package provides Huff0 encoding and decoding as used in zstd.
|
||||
|
||||
[Huff0](https://github.com/Cyan4973/FiniteStateEntropy#new-generation-entropy-coders),
|
||||
a Huffman codec designed for modern CPU, featuring OoO (Out of Order) operations on multiple ALU
|
||||
(Arithmetic Logic Unit), achieving extremely fast compression and decompression speeds.
|
||||
|
||||
This can be used for compressing input with a lot of similar input values to the smallest number of bytes.
|
||||
This does not perform any multi-byte [dictionary coding](https://en.wikipedia.org/wiki/Dictionary_coder) as LZ coders,
|
||||
but it can be used as a secondary step to compressors (like Snappy) that does not do entropy encoding.
|
||||
|
||||
* [Godoc documentation](https://godoc.org/github.com/klauspost/compress/huff0)
|
||||
|
||||
## News
|
||||
|
||||
This is used as part of the [zstandard](https://github.com/klauspost/compress/tree/master/zstd#zstd) compression and decompression package.
|
||||
|
||||
This ensures that most functionality is well tested.
|
||||
|
||||
# Usage
|
||||
|
||||
This package provides a low level interface that allows to compress single independent blocks.
|
||||
|
||||
Each block is separate, and there is no built in integrity checks.
|
||||
This means that the caller should keep track of block sizes and also do checksums if needed.
|
||||
|
||||
Compressing a block is done via the [`Compress1X`](https://godoc.org/github.com/klauspost/compress/huff0#Compress1X) and
|
||||
[`Compress4X`](https://godoc.org/github.com/klauspost/compress/huff0#Compress4X) functions.
|
||||
You must provide input and will receive the output and maybe an error.
|
||||
|
||||
These error values can be returned:
|
||||
|
||||
| Error | Description |
|
||||
|---------------------|-----------------------------------------------------------------------------|
|
||||
| `<nil>` | Everything ok, output is returned |
|
||||
| `ErrIncompressible` | Returned when input is judged to be too hard to compress |
|
||||
| `ErrUseRLE` | Returned from the compressor when the input is a single byte value repeated |
|
||||
| `ErrTooBig` | Returned if the input block exceeds the maximum allowed size (128 Kib) |
|
||||
| `(error)` | An internal error occurred. |
|
||||
|
||||
|
||||
As can be seen above some of there are errors that will be returned even under normal operation so it is important to handle these.
|
||||
|
||||
To reduce allocations you can provide a [`Scratch`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch) object
|
||||
that can be re-used for successive calls. Both compression and decompression accepts a `Scratch` object, and the same
|
||||
object can be used for both.
|
||||
|
||||
Be aware, that when re-using a `Scratch` object that the *output* buffer is also re-used, so if you are still using this
|
||||
you must set the `Out` field in the scratch to nil. The same buffer is used for compression and decompression output.
|
||||
|
||||
The `Scratch` object will retain state that allows to re-use previous tables for encoding and decoding.
|
||||
|
||||
## Tables and re-use
|
||||
|
||||
Huff0 allows for reusing tables from the previous block to save space if that is expected to give better/faster results.
|
||||
|
||||
The Scratch object allows you to set a [`ReusePolicy`](https://godoc.org/github.com/klauspost/compress/huff0#ReusePolicy)
|
||||
that controls this behaviour. See the documentation for details. This can be altered between each block.
|
||||
|
||||
Do however note that this information is *not* stored in the output block and it is up to the users of the package to
|
||||
record whether [`ReadTable`](https://godoc.org/github.com/klauspost/compress/huff0#ReadTable) should be called,
|
||||
based on the boolean reported back from the CompressXX call.
|
||||
|
||||
If you want to store the table separate from the data, you can access them as `OutData` and `OutTable` on the
|
||||
[`Scratch`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch) object.
|
||||
|
||||
## Decompressing
|
||||
|
||||
The first part of decoding is to initialize the decoding table through [`ReadTable`](https://godoc.org/github.com/klauspost/compress/huff0#ReadTable).
|
||||
This will initialize the decoding tables.
|
||||
You can supply the complete block to `ReadTable` and it will return the data part of the block
|
||||
which can be given to the decompressor.
|
||||
|
||||
Decompressing is done by calling the [`Decompress1X`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch.Decompress1X)
|
||||
or [`Decompress4X`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch.Decompress4X) function.
|
||||
|
||||
For concurrently decompressing content with a fixed table a stateless [`Decoder`](https://godoc.org/github.com/klauspost/compress/huff0#Decoder) can be requested which will remain correct as long as the scratch is unchanged. The capacity of the provided slice indicates the expected output size.
|
||||
|
||||
You must provide the output from the compression stage, at exactly the size you got back. If you receive an error back
|
||||
your input was likely corrupted.
|
||||
|
||||
It is important to note that a successful decoding does *not* mean your output matches your original input.
|
||||
There are no integrity checks, so relying on errors from the decompressor does not assure your data is valid.
|
||||
|
||||
# Contributing
|
||||
|
||||
Contributions are always welcome. Be aware that adding public functions will require good justification and breaking
|
||||
changes will likely not be accepted. If in doubt open an issue before writing the PR.
|
||||
# Huff0 entropy compression
|
||||
|
||||
This package provides Huff0 encoding and decoding as used in zstd.
|
||||
|
||||
[Huff0](https://github.com/Cyan4973/FiniteStateEntropy#new-generation-entropy-coders),
|
||||
a Huffman codec designed for modern CPU, featuring OoO (Out of Order) operations on multiple ALU
|
||||
(Arithmetic Logic Unit), achieving extremely fast compression and decompression speeds.
|
||||
|
||||
This can be used for compressing input with a lot of similar input values to the smallest number of bytes.
|
||||
This does not perform any multi-byte [dictionary coding](https://en.wikipedia.org/wiki/Dictionary_coder) as LZ coders,
|
||||
but it can be used as a secondary step to compressors (like Snappy) that does not do entropy encoding.
|
||||
|
||||
* [Godoc documentation](https://godoc.org/github.com/klauspost/compress/huff0)
|
||||
|
||||
## News
|
||||
|
||||
This is used as part of the [zstandard](https://github.com/klauspost/compress/tree/master/zstd#zstd) compression and decompression package.
|
||||
|
||||
This ensures that most functionality is well tested.
|
||||
|
||||
# Usage
|
||||
|
||||
This package provides a low level interface that allows to compress single independent blocks.
|
||||
|
||||
Each block is separate, and there is no built in integrity checks.
|
||||
This means that the caller should keep track of block sizes and also do checksums if needed.
|
||||
|
||||
Compressing a block is done via the [`Compress1X`](https://godoc.org/github.com/klauspost/compress/huff0#Compress1X) and
|
||||
[`Compress4X`](https://godoc.org/github.com/klauspost/compress/huff0#Compress4X) functions.
|
||||
You must provide input and will receive the output and maybe an error.
|
||||
|
||||
These error values can be returned:
|
||||
|
||||
| Error | Description |
|
||||
|---------------------|-----------------------------------------------------------------------------|
|
||||
| `<nil>` | Everything ok, output is returned |
|
||||
| `ErrIncompressible` | Returned when input is judged to be too hard to compress |
|
||||
| `ErrUseRLE` | Returned from the compressor when the input is a single byte value repeated |
|
||||
| `ErrTooBig` | Returned if the input block exceeds the maximum allowed size (128 Kib) |
|
||||
| `(error)` | An internal error occurred. |
|
||||
|
||||
|
||||
As can be seen above some of there are errors that will be returned even under normal operation so it is important to handle these.
|
||||
|
||||
To reduce allocations you can provide a [`Scratch`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch) object
|
||||
that can be re-used for successive calls. Both compression and decompression accepts a `Scratch` object, and the same
|
||||
object can be used for both.
|
||||
|
||||
Be aware, that when re-using a `Scratch` object that the *output* buffer is also re-used, so if you are still using this
|
||||
you must set the `Out` field in the scratch to nil. The same buffer is used for compression and decompression output.
|
||||
|
||||
The `Scratch` object will retain state that allows to re-use previous tables for encoding and decoding.
|
||||
|
||||
## Tables and re-use
|
||||
|
||||
Huff0 allows for reusing tables from the previous block to save space if that is expected to give better/faster results.
|
||||
|
||||
The Scratch object allows you to set a [`ReusePolicy`](https://godoc.org/github.com/klauspost/compress/huff0#ReusePolicy)
|
||||
that controls this behaviour. See the documentation for details. This can be altered between each block.
|
||||
|
||||
Do however note that this information is *not* stored in the output block and it is up to the users of the package to
|
||||
record whether [`ReadTable`](https://godoc.org/github.com/klauspost/compress/huff0#ReadTable) should be called,
|
||||
based on the boolean reported back from the CompressXX call.
|
||||
|
||||
If you want to store the table separate from the data, you can access them as `OutData` and `OutTable` on the
|
||||
[`Scratch`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch) object.
|
||||
|
||||
## Decompressing
|
||||
|
||||
The first part of decoding is to initialize the decoding table through [`ReadTable`](https://godoc.org/github.com/klauspost/compress/huff0#ReadTable).
|
||||
This will initialize the decoding tables.
|
||||
You can supply the complete block to `ReadTable` and it will return the data part of the block
|
||||
which can be given to the decompressor.
|
||||
|
||||
Decompressing is done by calling the [`Decompress1X`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch.Decompress1X)
|
||||
or [`Decompress4X`](https://godoc.org/github.com/klauspost/compress/huff0#Scratch.Decompress4X) function.
|
||||
|
||||
For concurrently decompressing content with a fixed table a stateless [`Decoder`](https://godoc.org/github.com/klauspost/compress/huff0#Decoder) can be requested which will remain correct as long as the scratch is unchanged. The capacity of the provided slice indicates the expected output size.
|
||||
|
||||
You must provide the output from the compression stage, at exactly the size you got back. If you receive an error back
|
||||
your input was likely corrupted.
|
||||
|
||||
It is important to note that a successful decoding does *not* mean your output matches your original input.
|
||||
There are no integrity checks, so relying on errors from the decompressor does not assure your data is valid.
|
||||
|
||||
# Contributing
|
||||
|
||||
Contributions are always welcome. Be aware that adding public functions will require good justification and breaking
|
||||
changes will likely not be accepted. If in doubt open an issue before writing the PR.
|
||||
|
||||
+5
-5
@@ -1,8 +1,8 @@
|
||||
GO_MOD_DIRS := $(shell find . -type f -name 'go.mod' -exec dirname {} \; | sort)
|
||||
REDIS_VERSION ?= 8.6
|
||||
REDIS_VERSION ?= 8.8
|
||||
RE_CLUSTER ?= false
|
||||
RCE_DOCKER ?= true
|
||||
CLIENT_LIBS_TEST_IMAGE ?= redislabs/client-libs-test:custom-21860421418-debian-amd64
|
||||
CLIENT_LIBS_TEST_IMAGE ?= redislabs/client-libs-test:8.8-m02
|
||||
|
||||
docker.start:
|
||||
export RE_CLUSTER=$(RE_CLUSTER) && \
|
||||
@@ -49,7 +49,7 @@ test.ci:
|
||||
export RE_CLUSTER=$(RE_CLUSTER) && \
|
||||
export RCE_DOCKER=$(RCE_DOCKER) && \
|
||||
export REDIS_VERSION=$(REDIS_VERSION) && \
|
||||
go mod tidy -compat=1.18 && \
|
||||
go mod tidy && \
|
||||
go vet && \
|
||||
go test -v -coverprofile=coverage.txt -covermode=atomic ./... -race -skip Example); \
|
||||
done
|
||||
@@ -63,7 +63,7 @@ test.ci.skip-vectorsets:
|
||||
export RE_CLUSTER=$(RE_CLUSTER) && \
|
||||
export RCE_DOCKER=$(RCE_DOCKER) && \
|
||||
export REDIS_VERSION=$(REDIS_VERSION) && \
|
||||
go mod tidy -compat=1.18 && \
|
||||
go mod tidy && \
|
||||
go vet && \
|
||||
go test -v -coverprofile=coverage.txt -covermode=atomic ./... -race \
|
||||
-run '^(?!.*(?:VectorSet|vectorset|ExampleClient_vectorset)).*$$' -skip Example); \
|
||||
@@ -118,5 +118,5 @@ go_mod_tidy:
|
||||
echo "go mod tidy in $${dir}"; \
|
||||
(cd "$${dir}" && \
|
||||
go get -u ./... && \
|
||||
go mod tidy -compat=1.18); \
|
||||
go mod tidy); \
|
||||
done
|
||||
|
||||
+25
-3
@@ -21,9 +21,8 @@ In `go-redis` we are aiming to support the last three releases of Redis. Current
|
||||
- [Redis 8.2](https://raw.githubusercontent.com/redis/redis/8.2/00-RELEASENOTES) - using Redis CE 8.2
|
||||
- [Redis 8.4](https://raw.githubusercontent.com/redis/redis/8.4/00-RELEASENOTES) - using Redis CE 8.4
|
||||
|
||||
Although the `go.mod` states it requires at minimum `go 1.21`, our CI is configured to run the tests against all three
|
||||
versions of Redis and multiple versions of Go ([1.21](https://go.dev/doc/devel/release#go1.21.0),
|
||||
[1.23](https://go.dev/doc/devel/release#go1.23.0), oldstable, and stable). We observe that some modules related test may not pass with
|
||||
Although the `go.mod` states it requires at minimum `go 1.24`, our CI is configured to run the tests against all three
|
||||
versions of Redis and multiple versions of Go ([1.24](https://go.dev/doc/devel/release#go1.24.0), oldstable, and stable). We observe that some modules related test may not pass with
|
||||
Redis Stack 7.2 and some commands are changed with Redis CE 8.0.
|
||||
Although it is not officially supported, `go-redis/v9` should be able to work with any Redis 7.0+.
|
||||
Please do refer to the documentation and the tests if you experience any issues.
|
||||
@@ -136,6 +135,29 @@ func ExampleClient() {
|
||||
}
|
||||
```
|
||||
|
||||
### Dial retries and backoff
|
||||
|
||||
Connection establishment can be retried by the connection pool when dialing fails.
|
||||
|
||||
- **`DialerRetries`**: maximum number of dial attempts (default: 5).
|
||||
- **`DialerRetryTimeout`**: default delay between attempts when no custom backoff is provided (default: 100ms).
|
||||
- **`DialerRetryBackoff`**: optional function hook to control the delay between attempts.
|
||||
|
||||
Example:
|
||||
|
||||
```go
|
||||
rdb := redis.NewClient(&redis.Options{
|
||||
Addr: "localhost:6379",
|
||||
|
||||
DialerRetries: 5,
|
||||
DialerRetryTimeout: 100 * time.Millisecond, // used when DialerRetryBackoff is nil
|
||||
|
||||
// Optional: exponential backoff with jitter and a cap.
|
||||
DialerRetryBackoff: redis.DialRetryBackoffExponential(100*time.Millisecond, 2*time.Second),
|
||||
})
|
||||
defer rdb.Close()
|
||||
```
|
||||
|
||||
### Authentication
|
||||
|
||||
The Redis client supports multiple ways to provide authentication credentials, with a clear priority order. Here are the available options:
|
||||
|
||||
+91
@@ -1,5 +1,96 @@
|
||||
# Release Notes
|
||||
|
||||
# 9.19.0 (2026-04-27)
|
||||
|
||||
## 🚀 Highlights
|
||||
|
||||
### FIPS-Compatible Script Helper
|
||||
|
||||
`Script` now supports a FIPS-safe execution mode that avoids client-side SHA-1 computation, which is blocked in strict FIPS environments. A new `NewScriptServerSHA` constructor uses `SCRIPT LOAD` to obtain and cache the digest from the server, then runs commands via `EVALSHA`/`EVALSHA_RO`. Falls back to `EVAL`/`EVALRO` if loading fails, and transparently retries once on `NOSCRIPT`. The default behavior is unchanged for existing users.
|
||||
|
||||
([#3700](https://github.com/redis/go-redis/pull/3700)) by [@chaitanyabodlapati](https://github.com/chaitanyabodlapati)
|
||||
|
||||
### FT.AGGREGATE Step-Based Pipeline Builder
|
||||
|
||||
Added a new step-based `FT.AGGREGATE` pipeline API via `FTAggregateOptions.Steps`, allowing `LOAD`, `APPLY`, `GROUPBY`, and `SORTBY` (with per-step `MAX`) to be repeated and interleaved in arbitrary order — matching Redis's native multi-stage aggregation semantics. The legacy `Load`/`Apply`/`GroupBy`/`SortBy`/`SortByMax` fields are now deprecated.
|
||||
|
||||
([#3782](https://github.com/redis/go-redis/pull/3782)) by [@ndyakov](https://github.com/ndyakov)
|
||||
|
||||
### Raw RESP Protocol Access
|
||||
|
||||
Added `DoRaw` and `DoRawWriteTo` methods for executing arbitrary commands and reading the raw RESP response. Useful for proxying, custom protocol inspection, and working with commands not yet wrapped by go-redis.
|
||||
|
||||
([#3713](https://github.com/redis/go-redis/pull/3713)) by [@ofekshenawa](https://github.com/ofekshenawa)
|
||||
|
||||
### Configurable Dial Retry Backoff
|
||||
|
||||
Added `DialerRetryBackoff` option (plumbed through `Options`, `ClusterOptions`, `RingOptions`, `FailoverOptions`) to let callers customize the delay between failed dial attempts. Helpers `DialRetryBackoffConstant` and `DialRetryBackoffExponential` (with jitter and cap) are provided out of the box. Dial timeout is now also applied **per attempt** rather than across all retries.
|
||||
|
||||
([#3706](https://github.com/redis/go-redis/pull/3706), [#3705](https://github.com/redis/go-redis/pull/3705)) by [@mwhooker](https://github.com/mwhooker)
|
||||
|
||||
## ✨ New Features
|
||||
|
||||
- **FT.AGGREGATE Steps**: Step-based pipeline builder for `FT.AGGREGATE` with support for repeated/interleaved `LOAD`, `APPLY`, `GROUPBY`, and `SORTBY` stages ([#3782](https://github.com/redis/go-redis/pull/3782)) by [@ndyakov](https://github.com/ndyakov)
|
||||
- **VectorSet commands**: Added `VISMEMBER` and `WITHATTRIBS` support ([#3753](https://github.com/redis/go-redis/pull/3753)) by [@romanpovol](https://github.com/romanpovol)
|
||||
- **FIPS-safe Script**: `NewScriptServerSHA` uses `SCRIPT LOAD` to obtain the digest from the server, avoiding client-side SHA-1 ([#3700](https://github.com/redis/go-redis/pull/3700)) by [@chaitanyabodlapati](https://github.com/chaitanyabodlapati)
|
||||
- **Raw RESP access**: `DoRaw` and `DoRawWriteTo` for raw RESP protocol access ([#3713](https://github.com/redis/go-redis/pull/3713)) by [@ofekshenawa](https://github.com/ofekshenawa)
|
||||
- **Dial retry backoff**: `DialerRetryBackoff` function option with constant and exponential helpers ([#3706](https://github.com/redis/go-redis/pull/3706)) by [@mwhooker](https://github.com/mwhooker)
|
||||
- **Typed NOSCRIPT error**: Redis `NOSCRIPT` replies are now surfaced as a typed error for easier handling ([#3738](https://github.com/redis/go-redis/pull/3738)) by [@LINKIWI](https://github.com/LINKIWI)
|
||||
- **PubSub ClientSetName**: Added `ClientSetName` method to `PubSub` ([#3727](https://github.com/redis/go-redis/pull/3727)) by [@Flack74](https://github.com/Flack74)
|
||||
- **ReplicaOf**: New `ReplicaOf` method replaces the deprecated `SlaveOf` ([#3720](https://github.com/redis/go-redis/pull/3720)) by [@Copilot](https://github.com/apps/copilot-swe-agent)
|
||||
- **HSCAN BinaryUnmarshaler**: `HScan` now supports types implementing `encoding.BinaryUnmarshaler` ([#3768](https://github.com/redis/go-redis/pull/3768)) by [@Aaditya-dubey1](https://github.com/Aaditya-dubey1)
|
||||
|
||||
## 🐛 Bug Fixes
|
||||
|
||||
- **Auto hostname type detection**: Improved endpoint type detection for maintenance notifications using DNS-based classification; handles empty hosts and expanded private-IP ranges ([#3789](https://github.com/redis/go-redis/pull/3789)) by [@ndyakov](https://github.com/ndyakov)
|
||||
- **HELLO fallback**: Don't send `CLIENT MAINT_NOTIFICATIONS` handshake when `HELLO` fails and connection falls back to RESP2; fail fast when explicitly enabled with RESP3 ([#3788](https://github.com/redis/go-redis/pull/3788)) by [@ndyakov](https://github.com/ndyakov)
|
||||
- **Dial TCP retry**: `ShouldRetry` now treats `net.OpError` with `Op == "dial"` timeout errors as safe to retry since no command was sent ([#3787](https://github.com/redis/go-redis/pull/3787)) by [@vladisa88](https://github.com/vladisa88)
|
||||
- **wrappedOnClose leak**: Fixed resource leak caused by repeatedly wrapping `baseClient` close logic; replaced with a bounded, concurrency-safe named-hook registry ([#3785](https://github.com/redis/go-redis/pull/3785)) by [@ndyakov](https://github.com/ndyakov)
|
||||
- **Pool Close() on stale connections**: Suppress close errors (e.g., TLS `closeNotify` timeouts) for connections already dropped by the server due to idle timeout ([#3778](https://github.com/redis/go-redis/pull/3778)) by [@ofekshenawa](https://github.com/ofekshenawa)
|
||||
- **FIFO waiter ordering**: Fixed race in `ConnStateMachine.notifyWaiters` that could wake multiple waiters under a single mutex hold and violate FIFO ordering ([#3777](https://github.com/redis/go-redis/pull/3777)) by [@0x48core](https://github.com/0x48core)
|
||||
- **Lua READONLY detection**: Detect `READONLY` errors embedded in Lua script error messages on read-only replicas so commands are correctly retried ([#3769](https://github.com/redis/go-redis/pull/3769)) by [@zhengjilei](https://github.com/zhengjilei)
|
||||
- **VectorScoreSliceCmd RESP2**: Fixed `VSimWithScores`, `VSimWithArgsWithScores`, and `VLinksWithScores` which were broken on RESP2 connections returning flat arrays instead of maps ([#3767](https://github.com/redis/go-redis/pull/3767)) by [@Copilot](https://github.com/apps/copilot-swe-agent)
|
||||
- **Closed connection handling**: Two fixes for closed connection handling in the pool ([#3764](https://github.com/redis/go-redis/pull/3764)) by [@cxljs](https://github.com/cxljs)
|
||||
- **ZRangeArgs Rev**: Fixed `ZRangeArgs` with `Rev` + `ByScore`/`ByLex` incorrectly swapping `Start`/`Stop`, breaking `ZRANGESTORE` ([#3751](https://github.com/redis/go-redis/pull/3751)) by [@Copilot](https://github.com/apps/copilot-swe-agent)
|
||||
- **OTel metric instrument types**: Fixed metric instrument types in `redisotel-native` ([#3743](https://github.com/redis/go-redis/pull/3743)) by [@ofekshenawa](https://github.com/ofekshenawa)
|
||||
- **Options.clone() data race**: Fixed data race when cloning `Options` ([#3739](https://github.com/redis/go-redis/pull/3739)) by [@rubensayshi](https://github.com/rubensayshi)
|
||||
- **Connection closure metrics**: Fixed connection closure metrics and enabled all metric groups by default in `redisotel-native` ([#3735](https://github.com/redis/go-redis/pull/3735)) by [@ofekshenawa](https://github.com/ofekshenawa)
|
||||
- **OTel semconv v1.38.0**: Use metric definition from `otel/semconv/v1.38.0` in `redisotel-native` ([#3731](https://github.com/redis/go-redis/pull/3731)) by [@wzy9607](https://github.com/wzy9607)
|
||||
- **SETNX semantics**: Use `SET ... NX` instead of the deprecated `SETNX` command ([#3723](https://github.com/redis/go-redis/pull/3723)) by [@ndyakov](https://github.com/ndyakov)
|
||||
- **TIME keyless routing**: Mark `TIME` as a keyless command for correct cluster routing ([#3722](https://github.com/redis/go-redis/pull/3722)) by [@fatal10110](https://github.com/fatal10110)
|
||||
- **Dial timeout per retry**: Dial timeout now applies per attempt instead of across all retry attempts combined ([#3705](https://github.com/redis/go-redis/pull/3705)) by [@mwhooker](https://github.com/mwhooker)
|
||||
- **Cluster metrics attributes**: Fixed `pool.name` being appended per node, which corrupted and dropped user-provided custom attributes ([#3699](https://github.com/redis/go-redis/pull/3699)) by [@Jesse-Bonfire](https://github.com/Jesse-Bonfire)
|
||||
- **initConn nil dereference**: Fixed nil pointer dereference and potential deadlock in `*baseClient.initConn()`; added explicit nil option guards to client constructors ([#3676](https://github.com/redis/go-redis/pull/3676)) by [@olde-ducke](https://github.com/olde-ducke)
|
||||
|
||||
## ⚡ Performance
|
||||
|
||||
- **RESP reader**: Optimized RESP reader by eliminating intermediate string allocations ([#3774](https://github.com/redis/go-redis/pull/3774)) by [@Aaditya-dubey1](https://github.com/Aaditya-dubey1)
|
||||
- **Inline rendezvous hashing**: Replaced `github.com/dgryski/go-rendezvous` dependency with an in-repo implementation in `internal/hashtag`, reducing the dependency graph while preserving algorithm parity ([#3762](https://github.com/redis/go-redis/pull/3762)) by [@bigsk05](https://github.com/bigsk05)
|
||||
|
||||
## 🧪 Testing & Infrastructure
|
||||
|
||||
- **Release automation**: Added `repository`, `ref`, and `client-libs-test-image-tag` inputs to the `run-tests` composite action; `redis-version` is now optional so unstable builds use `REDIS_VERSION` from the Makefile ([#3749](https://github.com/redis/go-redis/pull/3749)) by [@dariaguy](https://github.com/dariaguy)
|
||||
- **Go 1.24**: Updated minimum Go version to 1.24 and use `-compat=1.24` in release scripts ([#3714](https://github.com/redis/go-redis/pull/3714), [#3754](https://github.com/redis/go-redis/pull/3754)) by [@ndyakov](https://github.com/ndyakov), [@cxljs](https://github.com/cxljs)
|
||||
|
||||
## 🧰 Maintenance
|
||||
|
||||
- **Pool state machine**: Removed redundant `Conn.closed` atomic field in favor of the state machine's `StateClosed` ([#3783](https://github.com/redis/go-redis/pull/3783)) by [@cxljs](https://github.com/cxljs)
|
||||
- **OTel SDK**: Updated OpenTelemetry SDK dependencies in `redisotel`/`redisotel-native` ([#3770](https://github.com/redis/go-redis/pull/3770)) by [@ndyakov](https://github.com/ndyakov)
|
||||
- **Go 1.21+ built-ins**: Use `maps.Keys`, `slices.Collect`, `slices.Contains`, `clear()`, and `slices.SortFunc` instead of custom helpers ([#3758](https://github.com/redis/go-redis/pull/3758), [#3746](https://github.com/redis/go-redis/pull/3746)) by [@cxljs](https://github.com/cxljs)
|
||||
- **HGetAll docs**: Added Go doc comment to `HGetAll` describing behavior and complexity ([#3776](https://github.com/redis/go-redis/pull/3776)) by [@0x48core](https://github.com/0x48core)
|
||||
- **Docs links**: Fixed irrelevant docs links ([#3724](https://github.com/redis/go-redis/pull/3724)) by [@olzhas-sabiyev](https://github.com/olzhas-sabiyev)
|
||||
- **Examples cleanup**: Removed throughput binary from examples ([#3733](https://github.com/redis/go-redis/pull/3733)) by [@ndyakov](https://github.com/ndyakov)
|
||||
|
||||
## 👥 Contributors
|
||||
|
||||
We'd like to thank all the contributors who worked on this release!
|
||||
|
||||
[@0x48core](https://github.com/0x48core), [@Aaditya-dubey1](https://github.com/Aaditya-dubey1), [@Copilot](https://github.com/apps/copilot-swe-agent), [@Flack74](https://github.com/Flack74), [@Jesse-Bonfire](https://github.com/Jesse-Bonfire), [@LINKIWI](https://github.com/LINKIWI), [@bigsk05](https://github.com/bigsk05), [@chaitanyabodlapati](https://github.com/chaitanyabodlapati), [@cxljs](https://github.com/cxljs), [@dariaguy](https://github.com/dariaguy), [@fatal10110](https://github.com/fatal10110), [@mwhooker](https://github.com/mwhooker), [@ndyakov](https://github.com/ndyakov), [@ofekshenawa](https://github.com/ofekshenawa), [@olde-ducke](https://github.com/olde-ducke), [@olzhas-sabiyev](https://github.com/olzhas-sabiyev), [@romanpovol](https://github.com/romanpovol), [@rubensayshi](https://github.com/rubensayshi), [@vladisa88](https://github.com/vladisa88), [@wzy9607](https://github.com/wzy9607), [@zhengjilei](https://github.com/zhengjilei)
|
||||
|
||||
---
|
||||
|
||||
**Full Changelog**: https://github.com/redis/go-redis/compare/v9.18.0...v9.19.0
|
||||
|
||||
# 9.18.0 (2026-02-16)
|
||||
|
||||
## 🚀 Highlights
|
||||
|
||||
+137
-6
@@ -1,15 +1,146 @@
|
||||
# Releasing
|
||||
|
||||
1. Run `release.sh` script which updates versions in go.mod files and pushes a new branch to GitHub:
|
||||
This document is the runbook for cutting a go-redis release. It is intended
|
||||
for maintainers with write/tag access to the repository.
|
||||
|
||||
For the format and style of the release notes themselves, see
|
||||
[.github/RELEASE_NOTES_TEMPLATE.md](./.github/RELEASE_NOTES_TEMPLATE.md).
|
||||
|
||||
## Versioning
|
||||
|
||||
go-redis follows [Semantic Versioning](https://semver.org/):
|
||||
|
||||
- **Patch** (`vX.Y.Z+1`) — bug fixes, no API changes.
|
||||
- **Minor** (`vX.Y+1.0`) — backwards-compatible new features, deprecations.
|
||||
- **Major** (`vX+1.0.0`) — breaking changes. Coordinate with the team first.
|
||||
|
||||
Pre-releases use `vX.Y.Z-beta.N` / `vX.Y.Z-rc.N`.
|
||||
|
||||
## Pre-release checklist
|
||||
|
||||
- [ ] Target branch is `master` and CI is green on the latest commit.
|
||||
- [ ] All PRs intended for this release are merged.
|
||||
- [ ] There are no open issues in the release milestone (if used).
|
||||
- [ ] `CHANGELOG` / release notes have been considered; dependabot-only
|
||||
and doc-only changes are excluded per the template.
|
||||
- [ ] Confirm the next version number and decide if it's a patch / minor / major.
|
||||
|
||||
## 1. Draft the release notes
|
||||
|
||||
1. Open the draft release auto-generated by
|
||||
[release-drafter](.github/release-drafter-config.yml) on GitHub.
|
||||
2. Prepend a new section to [`RELEASE-NOTES.md`](./RELEASE-NOTES.md) using
|
||||
[`.github/RELEASE_NOTES_TEMPLATE.md`](./.github/RELEASE_NOTES_TEMPLATE.md)
|
||||
as the format. Keep the file in chronological order (newest first).
|
||||
3. Pick 3–5 **Highlights** — the most user-facing, impactful changes.
|
||||
4. Remove dependabot bumps and doc-only typo fixes from the lists.
|
||||
5. Verify every PR has a contributor attribution and link.
|
||||
6. Open a PR with just the release-notes change if you want review before
|
||||
bumping versions, otherwise include it in the release PR below.
|
||||
|
||||
## 2. Bump versions and open the release PR
|
||||
|
||||
Create a release branch from `master`:
|
||||
|
||||
```shell
|
||||
TAG=v1.0.0 ./scripts/release.sh
|
||||
git checkout master && git pull --ff-only
|
||||
git checkout -b release/vX.Y.Z
|
||||
```
|
||||
|
||||
2. Open a pull request and wait for the build to finish.
|
||||
|
||||
3. Merge the pull request and run `tag.sh` to create tags for packages:
|
||||
Run the release script on that branch:
|
||||
|
||||
```shell
|
||||
TAG=v1.0.0 ./scripts/tag.sh
|
||||
TAG=vX.Y.Z ./scripts/release.sh
|
||||
```
|
||||
|
||||
What the script does (and explicitly does **not** do):
|
||||
|
||||
- ✅ Validates `TAG` matches the semver regex and isn't already a git tag.
|
||||
- ✅ Rewrites every `redis/go-redis*` line in every sub-module `go.mod` to
|
||||
point at the new `TAG`. Trailing `// indirect` markers are preserved.
|
||||
- ✅ Runs `go mod tidy -compat=1.24` in each sub-module.
|
||||
- ✅ Updates the return value in [`version.go`](./version.go).
|
||||
- ❌ Does **not** switch branches (runs in your current branch).
|
||||
- ❌ Does **not** require a clean working tree (so you can mix it with
|
||||
release-notes edits in the same branch).
|
||||
- ❌ Does **not** commit, tag, or push anything.
|
||||
|
||||
Review and commit the changes yourself:
|
||||
|
||||
```shell
|
||||
git diff # sanity-check the bumps
|
||||
git add -u
|
||||
git commit -m "chore: release vX.Y.Z"
|
||||
git push origin release/vX.Y.Z
|
||||
```
|
||||
|
||||
Then on GitHub:
|
||||
|
||||
- [ ] Open a PR from `release/vX.Y.Z` into `master`.
|
||||
- [ ] Wait for all required CI checks (build, golangci-lint, spellcheck,
|
||||
doctests, e2e where applicable) to pass.
|
||||
- [ ] Get at least one maintainer approval.
|
||||
- [ ] Merge the PR (use a merge commit — the tag will point at the merge SHA).
|
||||
|
||||
## 3. Tag the release
|
||||
|
||||
After the release PR is merged, pull the latest `master` and dry-run the
|
||||
tagger:
|
||||
|
||||
```shell
|
||||
git checkout master && git pull --ff-only
|
||||
TAG=vX.Y.Z ./scripts/tag.sh vX.Y.Z
|
||||
```
|
||||
|
||||
The script defaults to **dry-run** and prints the commands it would run.
|
||||
Verify the output, then apply for real with `-t`:
|
||||
|
||||
```shell
|
||||
./scripts/tag.sh vX.Y.Z -t
|
||||
```
|
||||
|
||||
This creates and pushes:
|
||||
- The top-level tag `vX.Y.Z`.
|
||||
- A per-module tag `<module>/vX.Y.Z` for each public sub-module
|
||||
(skipping `example/*` and `internal/*`).
|
||||
|
||||
## 4. Publish the GitHub release
|
||||
|
||||
1. On GitHub, open the draft release created by release-drafter.
|
||||
2. Set the tag to `vX.Y.Z` and the target to `master`.
|
||||
3. Replace the auto-generated body with the curated notes from
|
||||
`RELEASE-NOTES.md` for this version.
|
||||
4. For pre-releases, check **"Set as a pre-release"**.
|
||||
5. Publish.
|
||||
|
||||
## 5. Post-release
|
||||
|
||||
- [ ] Verify the release appears on
|
||||
[pkg.go.dev](https://pkg.go.dev/github.com/redis/go-redis/v9) within
|
||||
a few minutes (trigger a fetch by visiting the version URL if needed).
|
||||
- [ ] Announce on Discord (see the link in `CONTRIBUTING.md`).
|
||||
- [ ] Close the release milestone if one was used.
|
||||
- [ ] Open follow-up issues for anything deferred from this release.
|
||||
|
||||
## Hotfix / patch release
|
||||
|
||||
For an urgent fix on top of the latest release:
|
||||
|
||||
1. Branch from the latest release tag: `git checkout -b hotfix/vX.Y.Z+1 vX.Y.Z`.
|
||||
2. Cherry-pick (or re-apply) only the required fix commits.
|
||||
3. Follow the normal release flow above with `TAG=vX.Y.Z+1`.
|
||||
4. Make sure the fix is also present on `master` (forward-port if necessary).
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **`release.sh` fails with "tag already exists"** — the tag has already
|
||||
been created. Pick the next version, or delete the local tag first if
|
||||
it was created by mistake.
|
||||
- **`tag.sh` reports version mismatch in a `go.mod`** — a sub-module was
|
||||
not updated by `release.sh`. Fix the `go.mod` manually (or re-run
|
||||
`release.sh`), amend the release PR, and re-run the tagger.
|
||||
- **`version.go` does not contain the tag** — `release.sh` did not run or
|
||||
the bump was reverted. Re-run `release.sh` on the release branch.
|
||||
- **pkg.go.dev does not show the new version** — visit
|
||||
`https://pkg.go.dev/github.com/redis/go-redis/v9@vX.Y.Z` once to trigger
|
||||
a fetch from the module proxy.
|
||||
|
||||
+18
@@ -9,12 +9,30 @@ type StreamingCredentialsProvider interface {
|
||||
// Subscribe subscribes to the credentials provider for updates.
|
||||
// It returns the current credentials, a cancel function to unsubscribe from the provider,
|
||||
// and an error if any.
|
||||
//
|
||||
// Implementations MUST be idempotent with respect to listener identity:
|
||||
// subscribing the same listener value more than once must not produce
|
||||
// duplicate notifications and must not create multiple independent
|
||||
// subscriptions that each need to be cancelled separately. Every
|
||||
// UnsubscribeFunc returned for a given listener must cancel that
|
||||
// listener's subscription; calling any one of them must be sufficient to
|
||||
// stop updates to that listener, and calling subsequent ones must be a
|
||||
// safe no-op. Callers (including go-redis internals) may retain only
|
||||
// the most recently returned UnsubscribeFunc and rely on it to fully
|
||||
// unsubscribe the listener.
|
||||
//
|
||||
// TODO(ndyakov): Should we add context to the Subscribe method?
|
||||
Subscribe(listener CredentialsListener) (Credentials, UnsubscribeFunc, error)
|
||||
}
|
||||
|
||||
// UnsubscribeFunc is a function that is used to cancel the subscription to the credentials provider.
|
||||
// It is used to unsubscribe from the provider when the credentials are no longer needed.
|
||||
//
|
||||
// Per the StreamingCredentialsProvider.Subscribe contract, if the same
|
||||
// listener is subscribed multiple times, every UnsubscribeFunc returned for
|
||||
// that listener must fully unsubscribe it on first invocation, and
|
||||
// subsequent invocations (from any of the equivalent UnsubscribeFuncs) must
|
||||
// be a safe no-op.
|
||||
type UnsubscribeFunc func() error
|
||||
|
||||
// CredentialsListener is an interface that defines the methods for a credentials listener.
|
||||
|
||||
+421
-15
@@ -4,6 +4,7 @@ import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"maps"
|
||||
"net"
|
||||
"regexp"
|
||||
@@ -65,6 +66,7 @@ var keylessCommands = map[string]struct{}{
|
||||
"subscribe": {},
|
||||
"swapdb": {},
|
||||
"sync": {},
|
||||
"time": {},
|
||||
"unsubscribe": {},
|
||||
"unwatch": {},
|
||||
"wait": {},
|
||||
@@ -214,6 +216,11 @@ type Cmder interface {
|
||||
SetErr(error)
|
||||
Err() error
|
||||
|
||||
// NoRetry returns true if the command should not be retried on failure.
|
||||
// Commands that write directly to an io.Writer should return true since
|
||||
// partial writes cannot be undone on retry.
|
||||
NoRetry() bool
|
||||
|
||||
// GetCmdType returns the command type for fast value extraction
|
||||
GetCmdType() CmdType
|
||||
}
|
||||
@@ -235,6 +242,18 @@ func cmdsFirstErr(cmds []Cmder) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// cmdsContainNoRetry returns true if any command in the slice has NoRetry() == true.
|
||||
// If a pipeline contains a non-retryable command (e.g., RawWriteToCmd), the entire
|
||||
// pipeline must not be retried to prevent data corruption from partial writes.
|
||||
func cmdsContainNoRetry(cmds []Cmder) bool {
|
||||
for _, cmd := range cmds {
|
||||
if cmd.NoRetry() {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func writeCmds(wr *proto.Writer, cmds []Cmder) error {
|
||||
for _, cmd := range cmds {
|
||||
if err := writeCmd(wr, cmd); err != nil {
|
||||
@@ -397,6 +416,14 @@ func (cmd *baseCmd) readRawReply(rd *proto.Reader) (err error) {
|
||||
return err
|
||||
}
|
||||
|
||||
// NoRetry returns true if the command should not be retried on failure.
|
||||
// By default, commands can be retried. Commands that write directly to an
|
||||
// io.Writer (like RawWriteToCmd) should override this to return true since
|
||||
// partial writes cannot be undone on retry.
|
||||
func (cmd *baseCmd) NoRetry() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (cmd *baseCmd) GetCmdType() CmdType {
|
||||
return cmd.cmdType
|
||||
}
|
||||
@@ -719,6 +746,122 @@ func (cmd *Cmd) Clone() Cmder {
|
||||
|
||||
//------------------------------------------------------------------------------
|
||||
|
||||
// RawCmd returns raw RESP protocol bytes without parsing.
|
||||
type RawCmd struct {
|
||||
baseCmd
|
||||
val []byte
|
||||
}
|
||||
|
||||
var _ Cmder = (*RawCmd)(nil)
|
||||
|
||||
func NewRawCmd(ctx context.Context, args ...interface{}) *RawCmd {
|
||||
return &RawCmd{
|
||||
baseCmd: baseCmd{
|
||||
ctx: ctx,
|
||||
args: args,
|
||||
cmdType: CmdTypeGeneric,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) SetVal(val []byte) {
|
||||
cmd.val = val
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) Val() []byte {
|
||||
return cmd.val
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) Result() ([]byte, error) {
|
||||
return cmd.val, cmd.err
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) Bytes() ([]byte, error) {
|
||||
return cmd.val, cmd.err
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) String() string {
|
||||
return cmdString(cmd, cmd.val)
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) readReply(rd *proto.Reader) (err error) {
|
||||
cmd.val, err = rd.ReadRawReply()
|
||||
return err
|
||||
}
|
||||
|
||||
func (cmd *RawCmd) Clone() Cmder {
|
||||
var val []byte
|
||||
if cmd.val != nil {
|
||||
val = make([]byte, len(cmd.val))
|
||||
copy(val, cmd.val)
|
||||
}
|
||||
return &RawCmd{
|
||||
baseCmd: cmd.cloneBaseCmd(),
|
||||
val: val,
|
||||
}
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------------
|
||||
|
||||
// RawWriteToCmd streams raw RESP protocol bytes directly to an io.Writer without intermediate allocations.
|
||||
type RawWriteToCmd struct {
|
||||
baseCmd
|
||||
w io.Writer
|
||||
written int64
|
||||
}
|
||||
|
||||
var _ Cmder = (*RawWriteToCmd)(nil)
|
||||
|
||||
func NewRawWriteToCmd(ctx context.Context, w io.Writer, args ...interface{}) *RawWriteToCmd {
|
||||
return &RawWriteToCmd{
|
||||
baseCmd: baseCmd{
|
||||
ctx: ctx,
|
||||
args: args,
|
||||
cmdType: CmdTypeGeneric,
|
||||
},
|
||||
w: w,
|
||||
}
|
||||
}
|
||||
|
||||
func (cmd *RawWriteToCmd) SetVal(written int64) {
|
||||
cmd.written = written
|
||||
}
|
||||
|
||||
func (cmd *RawWriteToCmd) Val() int64 {
|
||||
return cmd.written
|
||||
}
|
||||
|
||||
func (cmd *RawWriteToCmd) Result() (int64, error) {
|
||||
return cmd.written, cmd.err
|
||||
}
|
||||
|
||||
func (cmd *RawWriteToCmd) String() string {
|
||||
return cmdString(cmd, cmd.written)
|
||||
}
|
||||
|
||||
func (cmd *RawWriteToCmd) readReply(rd *proto.Reader) (err error) {
|
||||
cmd.written, err = rd.ReadRawReplyWriteTo(cmd.w)
|
||||
return err
|
||||
}
|
||||
|
||||
// NoRetry returns true because RawWriteToCmd writes directly to an io.Writer.
|
||||
// If a retry occurs, partial data from failed attempts would be appended to
|
||||
// the writer, causing data corruption. The caller must handle retries manually
|
||||
// if needed, using a fresh writer for each attempt.
|
||||
func (cmd *RawWriteToCmd) NoRetry() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (cmd *RawWriteToCmd) Clone() Cmder {
|
||||
return &RawWriteToCmd{
|
||||
baseCmd: cmd.cloneBaseCmd(),
|
||||
w: cmd.w,
|
||||
written: cmd.written,
|
||||
}
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------------
|
||||
|
||||
type SliceCmd struct {
|
||||
baseCmd
|
||||
|
||||
@@ -2075,10 +2218,7 @@ func (cmd *XMessageSliceCmd) Clone() Cmder {
|
||||
ID: msg.ID,
|
||||
}
|
||||
if msg.Values != nil {
|
||||
val[i].Values = make(map[string]interface{}, len(msg.Values))
|
||||
for k, v := range msg.Values {
|
||||
val[i].Values[k] = v
|
||||
}
|
||||
val[i].Values = maps.Clone(msg.Values)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -2727,7 +2867,10 @@ func (cmd *XInfoConsumersCmd) readReply(rd *proto.Reader) error {
|
||||
inactive, err = rd.ReadInt()
|
||||
cmd.val[i].Inactive = time.Duration(inactive) * time.Millisecond
|
||||
default:
|
||||
return fmt.Errorf("redis: unexpected content %s in XINFO CONSUMERS reply", key)
|
||||
// skip unknown fields
|
||||
if err = rd.DiscardNext(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2856,7 +2999,10 @@ func (cmd *XInfoGroupsCmd) readReply(rd *proto.Reader) error {
|
||||
group.Lag = -1
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("redis: unexpected key %q in XINFO GROUPS reply", key)
|
||||
// skip unknown fields
|
||||
if err = rd.DiscardNext(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -3025,7 +3171,10 @@ func (cmd *XInfoStreamCmd) readReply(rd *proto.Reader) error {
|
||||
return err
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("redis: unexpected key %q in XINFO STREAM reply", key)
|
||||
// skip unknown fields
|
||||
if err = rd.DiscardNext(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
@@ -3101,6 +3250,7 @@ type XInfoStreamGroup struct {
|
||||
EntriesRead int64
|
||||
Lag int64
|
||||
PelCount int64
|
||||
NackedCount uint64 // redis version 8.8, number of NACK'd messages in the group
|
||||
Pending []XInfoStreamGroupPending
|
||||
Consumers []XInfoStreamConsumer
|
||||
}
|
||||
@@ -3245,7 +3395,10 @@ func (cmd *XInfoStreamFullCmd) readReply(rd *proto.Reader) error {
|
||||
return err
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("redis: unexpected key %q in XINFO STREAM FULL reply", key)
|
||||
// skip unknown fields
|
||||
if err = rd.DiscardNext(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
@@ -3299,6 +3452,11 @@ func readStreamGroups(rd *proto.Reader) ([]XInfoStreamGroup, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case "nacked-count":
|
||||
group.NackedCount, err = rd.ReadUint()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
case "pending":
|
||||
group.Pending, err = readXInfoStreamGroupPending(rd)
|
||||
if err != nil {
|
||||
@@ -3310,7 +3468,10 @@ func readStreamGroups(rd *proto.Reader) ([]XInfoStreamGroup, error) {
|
||||
return nil, err
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("redis: unexpected key %q in XINFO STREAM FULL reply", key)
|
||||
// skip unknown fields
|
||||
if err = rd.DiscardNext(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3435,8 +3596,10 @@ func readXInfoStreamConsumers(rd *proto.Reader) ([]XInfoStreamConsumer, error) {
|
||||
c.Pending = append(c.Pending, p)
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("redis: unexpected content %s "+
|
||||
"in XINFO STREAM FULL reply", cKey)
|
||||
// skip unknown fields
|
||||
if err = rd.DiscardNext(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -6833,6 +6996,9 @@ type ClientInfo struct {
|
||||
Resp int // redis version 7.0, client RESP protocol version
|
||||
LibName string // redis version 7.2, client library name
|
||||
LibVer string // redis version 7.2, client library version
|
||||
ReadEvents uint64 // redis version 8.8, number of read events processed
|
||||
AvgPipelineLenSum uint64 // redis version 8.8, sum of pipeline lengths
|
||||
AvgPipelineLenCnt uint64 // redis version 8.8, count of pipeline operations
|
||||
}
|
||||
|
||||
type ClientInfoCmd struct {
|
||||
@@ -7013,8 +7179,14 @@ func parseClientInfo(txt string) (info *ClientInfo, err error) {
|
||||
info.LibVer = val
|
||||
case "io-thread":
|
||||
info.IoThread, err = strconv.Atoi(val)
|
||||
case "read-events":
|
||||
info.ReadEvents, err = strconv.ParseUint(val, 10, 64)
|
||||
case "avg-pipeline-len-sum":
|
||||
info.AvgPipelineLenSum, err = strconv.ParseUint(val, 10, 64)
|
||||
case "avg-pipeline-len-cnt":
|
||||
info.AvgPipelineLenCnt, err = strconv.ParseUint(val, 10, 64)
|
||||
default:
|
||||
return nil, fmt.Errorf("redis: unexpected client info key(%s)", key)
|
||||
// skip unknown fields
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
@@ -7061,6 +7233,9 @@ func (cmd *ClientInfoCmd) Clone() Cmder {
|
||||
Resp: cmd.val.Resp,
|
||||
LibName: cmd.val.LibName,
|
||||
LibVer: cmd.val.LibVer,
|
||||
ReadEvents: cmd.val.ReadEvents,
|
||||
AvgPipelineLenSum: cmd.val.AvgPipelineLenSum,
|
||||
AvgPipelineLenCnt: cmd.val.AvgPipelineLenCnt,
|
||||
}
|
||||
}
|
||||
return &ClientInfoCmd{
|
||||
@@ -7167,7 +7342,10 @@ func (cmd *ACLLogCmd) readReply(rd *proto.Reader) error {
|
||||
case "timestamp-last-updated":
|
||||
entry.TimestampLastUpdated, err = rd.ReadInt()
|
||||
default:
|
||||
return fmt.Errorf("redis: unexpected key %q in ACL LOG reply", key)
|
||||
// skip unknown fields
|
||||
if err := rd.DiscardNext(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
@@ -7231,6 +7409,9 @@ func (cmd *ACLLogCmd) Clone() Cmder {
|
||||
Resp: entry.ClientInfo.Resp,
|
||||
LibName: entry.ClientInfo.LibName,
|
||||
LibVer: entry.ClientInfo.LibVer,
|
||||
ReadEvents: entry.ClientInfo.ReadEvents,
|
||||
AvgPipelineLenSum: entry.ClientInfo.AvgPipelineLenSum,
|
||||
AvgPipelineLenCnt: entry.ClientInfo.AvgPipelineLenCnt,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -7451,7 +7632,7 @@ type VectorScoreSliceCmd struct {
|
||||
|
||||
var _ Cmder = (*VectorScoreSliceCmd)(nil)
|
||||
|
||||
func NewVectorInfoSliceCmd(ctx context.Context, args ...any) *VectorScoreSliceCmd {
|
||||
func NewVectorScoreSliceCmd(ctx context.Context, args ...any) *VectorScoreSliceCmd {
|
||||
return &VectorScoreSliceCmd{
|
||||
baseCmd: baseCmd{
|
||||
ctx: ctx,
|
||||
@@ -7460,6 +7641,11 @@ func NewVectorInfoSliceCmd(ctx context.Context, args ...any) *VectorScoreSliceCm
|
||||
}
|
||||
}
|
||||
|
||||
// NewVectorInfoSliceCmd is an alias for NewVectorScoreSliceCmd kept for backwards compatibility.
|
||||
func NewVectorInfoSliceCmd(ctx context.Context, args ...any) *VectorScoreSliceCmd {
|
||||
return NewVectorScoreSliceCmd(ctx, args...)
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreSliceCmd) SetVal(val []VectorScore) {
|
||||
cmd.val = val
|
||||
}
|
||||
@@ -7477,11 +7663,29 @@ func (cmd *VectorScoreSliceCmd) String() string {
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreSliceCmd) readReply(rd *proto.Reader) error {
|
||||
n, err := rd.ReadMapLen()
|
||||
typ, err := rd.PeekReplyType()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var n int
|
||||
if typ == proto.RespMap {
|
||||
n, err = rd.ReadMapLen()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
// RESP2 returns a flat array [name, score, name, score, ...]
|
||||
n, err = rd.ReadArrayLen()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n%2 != 0 {
|
||||
return fmt.Errorf("redis: VectorScoreSliceCmd expects even number of elements, got %d", n)
|
||||
}
|
||||
n /= 2
|
||||
}
|
||||
|
||||
cmd.val = make([]VectorScore, n)
|
||||
for i := 0; i < n; i++ {
|
||||
name, err := rd.ReadString()
|
||||
@@ -7507,6 +7711,208 @@ func (cmd *VectorScoreSliceCmd) Clone() Cmder {
|
||||
}
|
||||
}
|
||||
|
||||
func readVectorAttribStringOrNil(rd *proto.Reader) (*string, error) {
|
||||
v, err := rd.ReadReply()
|
||||
if err != nil {
|
||||
if err == proto.Nil {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
s, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("redis: can't parse reply=%T reading string", v)
|
||||
}
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
type VectorAttribSliceCmd struct {
|
||||
baseCmd
|
||||
|
||||
val []VectorAttrib
|
||||
}
|
||||
|
||||
var _ Cmder = (*VectorAttribSliceCmd)(nil)
|
||||
|
||||
func NewVectorAttribSliceCmd(ctx context.Context, args ...any) *VectorAttribSliceCmd {
|
||||
return &VectorAttribSliceCmd{
|
||||
baseCmd: baseCmd{
|
||||
ctx: ctx,
|
||||
args: args,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (cmd *VectorAttribSliceCmd) SetVal(val []VectorAttrib) {
|
||||
cmd.val = val
|
||||
}
|
||||
|
||||
func (cmd *VectorAttribSliceCmd) Val() []VectorAttrib {
|
||||
return cmd.val
|
||||
}
|
||||
|
||||
func (cmd *VectorAttribSliceCmd) Result() ([]VectorAttrib, error) {
|
||||
return cmd.val, cmd.err
|
||||
}
|
||||
|
||||
func (cmd *VectorAttribSliceCmd) String() string {
|
||||
return cmdString(cmd, cmd.val)
|
||||
}
|
||||
|
||||
func (cmd *VectorAttribSliceCmd) readReply(rd *proto.Reader) error {
|
||||
replyType, err := rd.PeekReplyType()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if replyType == proto.RespMap {
|
||||
n, err := rd.ReadMapLen()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.val = make([]VectorAttrib, n)
|
||||
for i := 0; i < n; i++ {
|
||||
name, err := rd.ReadString()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
attrib, err := readVectorAttribStringOrNil(rd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.val[i] = VectorAttrib{Name: name, Attribs: attrib}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
n, err := rd.ReadArrayLen()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n%2 != 0 {
|
||||
return fmt.Errorf("redis: got %d elements in the VSIM array, wanted a multiple of 2", n)
|
||||
}
|
||||
cmd.val = make([]VectorAttrib, n/2)
|
||||
for i := range cmd.val {
|
||||
name, err := rd.ReadString()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
attrib, err := readVectorAttribStringOrNil(rd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.val[i] = VectorAttrib{Name: name, Attribs: attrib}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cmd *VectorAttribSliceCmd) Clone() Cmder {
|
||||
return &VectorAttribSliceCmd{
|
||||
baseCmd: cmd.cloneBaseCmd(),
|
||||
val: cmd.val,
|
||||
}
|
||||
}
|
||||
|
||||
type VectorScoreAttribSliceCmd struct {
|
||||
baseCmd
|
||||
|
||||
val []VectorScoreAttrib
|
||||
}
|
||||
|
||||
var _ Cmder = (*VectorScoreAttribSliceCmd)(nil)
|
||||
|
||||
func NewVectorScoreAttribSliceCmd(ctx context.Context, args ...any) *VectorScoreAttribSliceCmd {
|
||||
return &VectorScoreAttribSliceCmd{
|
||||
baseCmd: baseCmd{
|
||||
ctx: ctx,
|
||||
args: args,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreAttribSliceCmd) SetVal(val []VectorScoreAttrib) {
|
||||
cmd.val = val
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreAttribSliceCmd) Val() []VectorScoreAttrib {
|
||||
return cmd.val
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreAttribSliceCmd) Result() ([]VectorScoreAttrib, error) {
|
||||
return cmd.val, cmd.err
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreAttribSliceCmd) String() string {
|
||||
return cmdString(cmd, cmd.val)
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreAttribSliceCmd) readReply(rd *proto.Reader) error {
|
||||
replyType, err := rd.PeekReplyType()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if replyType == proto.RespMap {
|
||||
n, err := rd.ReadMapLen()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.val = make([]VectorScoreAttrib, n)
|
||||
for i := 0; i < n; i++ {
|
||||
name, err := rd.ReadString()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := rd.ReadFixedArrayLen(2); err != nil {
|
||||
return err
|
||||
}
|
||||
score, err := rd.ReadFloat()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
attrib, err := readVectorAttribStringOrNil(rd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.val[i] = VectorScoreAttrib{Name: name, Score: score, Attribs: attrib}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
n, err := rd.ReadArrayLen()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n%3 != 0 {
|
||||
return fmt.Errorf("redis: got %d elements in the VSIM array, wanted a multiple of 3", n)
|
||||
}
|
||||
cmd.val = make([]VectorScoreAttrib, n/3)
|
||||
for i := range cmd.val {
|
||||
name, err := rd.ReadString()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
score, err := rd.ReadFloat()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
attrib, err := readVectorAttribStringOrNil(rd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.val[i] = VectorScoreAttrib{Name: name, Score: score, Attribs: attrib}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cmd *VectorScoreAttribSliceCmd) Clone() Cmder {
|
||||
return &VectorScoreAttribSliceCmd{
|
||||
baseCmd: cmd.cloneBaseCmd(),
|
||||
val: cmd.val,
|
||||
}
|
||||
}
|
||||
|
||||
func (cmd *MonitorCmd) Clone() Cmder {
|
||||
// MonitorCmd cannot be safely cloned due to channels and goroutines
|
||||
// Return a new MonitorCmd with the same channel
|
||||
|
||||
+22
@@ -215,6 +215,7 @@ type Cmdable interface {
|
||||
ShutdownSave(ctx context.Context) *StatusCmd
|
||||
ShutdownNoSave(ctx context.Context) *StatusCmd
|
||||
SlaveOf(ctx context.Context, host, port string) *StatusCmd
|
||||
ReplicaOf(ctx context.Context, host, port string) *StatusCmd
|
||||
SlowLogGet(ctx context.Context, num int64) *SlowLogCmd
|
||||
SlowLogLen(ctx context.Context) *IntCmd
|
||||
SlowLogReset(ctx context.Context) *StatusCmd
|
||||
@@ -448,6 +449,20 @@ func (c cmdable) Do(ctx context.Context, args ...interface{}) *Cmd {
|
||||
return cmd
|
||||
}
|
||||
|
||||
// DoRaw executes a command and returns the raw RESP protocol bytes without parsing.
|
||||
func (c cmdable) DoRaw(ctx context.Context, args ...interface{}) *RawCmd {
|
||||
cmd := NewRawCmd(ctx, args...)
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// DoRawWriteTo executes a command and streams raw RESP bytes directly to w without intermediate allocations.
|
||||
func (c cmdable) DoRawWriteTo(ctx context.Context, w io.Writer, args ...interface{}) *RawWriteToCmd {
|
||||
cmd := NewRawWriteToCmd(ctx, w, args...)
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// Quit closes the connection.
|
||||
//
|
||||
// Deprecated: Just close the connection instead as of Redis 7.2.0.
|
||||
@@ -682,6 +697,13 @@ func (c cmdable) SlaveOf(ctx context.Context, host, port string) *StatusCmd {
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ReplicaOf sets a Redis server as a replica of another, or promotes it to being a master.
|
||||
func (c cmdable) ReplicaOf(ctx context.Context, host, port string) *StatusCmd {
|
||||
cmd := NewStatusCmd(ctx, "replicaof", host, port)
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func (c cmdable) SlowLogGet(ctx context.Context, num int64) *SlowLogCmd {
|
||||
cmd := NewSlowLogCmd(context.Background(), "slowlog", "get", num)
|
||||
_ = c(ctx, cmd)
|
||||
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
package redis
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9/internal"
|
||||
)
|
||||
|
||||
// DialRetryBackoffConstant returns a dial retry backoff function that always returns d.
|
||||
// attempt is 0-based: attempt=0 is the delay after the 1st failed dial.
|
||||
func DialRetryBackoffConstant(d time.Duration) func(attempt int) time.Duration {
|
||||
if d < 0 {
|
||||
d = 0
|
||||
}
|
||||
return func(int) time.Duration { return d }
|
||||
}
|
||||
|
||||
// DialRetryBackoffExponential returns a dial retry backoff function that uses exponential
|
||||
// backoff with jitter and a cap, using internal.RetryBackoff.
|
||||
//
|
||||
// attempt is 0-based: attempt=0 is the delay after the 1st failed dial.
|
||||
func DialRetryBackoffExponential(minBackoff, maxBackoff time.Duration) func(attempt int) time.Duration {
|
||||
if minBackoff < 0 {
|
||||
minBackoff = 0
|
||||
}
|
||||
if maxBackoff < 0 {
|
||||
maxBackoff = 0
|
||||
}
|
||||
if minBackoff > maxBackoff {
|
||||
minBackoff = maxBackoff
|
||||
}
|
||||
return func(attempt int) time.Duration {
|
||||
// internal.RetryBackoff expects retry >= 0.
|
||||
if attempt < 0 {
|
||||
attempt = 0
|
||||
}
|
||||
return internal.RetryBackoff(attempt, minBackoff, maxBackoff)
|
||||
}
|
||||
}
|
||||
+4
-4
@@ -1,6 +1,6 @@
|
||||
---
|
||||
|
||||
x-default-image: &default-image ${CLIENT_LIBS_TEST_IMAGE:-redislabs/client-libs-test:8.6.0}
|
||||
x-default-image: &default-image ${CLIENT_LIBS_TEST_IMAGE:-redislabs/client-libs-test:8.8-m02}
|
||||
|
||||
services:
|
||||
redis:
|
||||
@@ -164,9 +164,9 @@ services:
|
||||
- PORT=6390
|
||||
command: ${REDIS_EXTRA_ARGS:---enable-debug-command yes --enable-module-command yes --tls-auth-clients optional --save ""}
|
||||
ports:
|
||||
- 6390:6390
|
||||
- 6391:6391
|
||||
- 6392:6392
|
||||
- "6390:6390"
|
||||
- "6391:6391"
|
||||
- "6392:6392"
|
||||
volumes:
|
||||
- "./dockers/ring:/redis/work"
|
||||
profiles:
|
||||
|
||||
+14
@@ -28,6 +28,11 @@ var ErrPoolTimeout = pool.ErrPoolTimeout
|
||||
// is used on a ClusterClient with keys in different slots.
|
||||
var ErrCrossSlot = proto.RedisError("CROSSSLOT Keys in request don't hash to the same slot")
|
||||
|
||||
// ErrNoScript is returned when EVALSHA is requested for a script digest that
|
||||
// is not available in the script cache. Note that this error text is reproduced
|
||||
// literally from that used by Redis.
|
||||
var ErrNoScript = proto.RedisError("NOSCRIPT No matching script. Please use EVAL.")
|
||||
|
||||
// HasErrorPrefix checks if the err is a Redis error and the message contains a prefix.
|
||||
func HasErrorPrefix(err error, prefix string) bool {
|
||||
var rErr Error
|
||||
@@ -100,6 +105,12 @@ func shouldRetry(err error, retryTimeout bool) bool {
|
||||
// Check for timeout errors (works with wrapped errors)
|
||||
if isTimeout, hasTimeoutFlag := isTimeoutError(err); isTimeout {
|
||||
if hasTimeoutFlag {
|
||||
// A dial error means the TCP connection was never established and the
|
||||
// command was never sent to the server, so retry is always safe
|
||||
var opErr *net.OpError
|
||||
if errors.As(err, &opErr) && opErr.Op == "dial" {
|
||||
return true
|
||||
}
|
||||
return retryTimeout
|
||||
}
|
||||
return true
|
||||
@@ -139,6 +150,9 @@ func shouldRetry(err error, retryTimeout bool) bool {
|
||||
if strings.HasPrefix(s, "READONLY ") {
|
||||
return true
|
||||
}
|
||||
if strings.Contains(s, "-READONLY You can't write against a read only replica") {
|
||||
return true
|
||||
}
|
||||
if strings.HasPrefix(s, "CLUSTERDOWN ") {
|
||||
return true
|
||||
}
|
||||
|
||||
+7
@@ -70,6 +70,13 @@ func (c cmdable) HGet(ctx context.Context, key, field string) *StringCmd {
|
||||
return cmd
|
||||
}
|
||||
|
||||
// HGetAll returns a map of all fields and values stored at key.
|
||||
//
|
||||
// Returns an empty map when key does not exist.
|
||||
//
|
||||
// Time complexity: O(N) where N is the size of the hash.
|
||||
//
|
||||
// See https://redis.io/commands/hgetall/
|
||||
func (c cmdable) HGetAll(ctx context.Context, key string) *MapStringStringCmd {
|
||||
cmd := NewMapStringStringCmd(ctx, "hgetall", key)
|
||||
_ = c(ctx, cmd)
|
||||
|
||||
+1
-1
@@ -11,7 +11,7 @@ const slotNumber = 16384
|
||||
// CRC16 implementation according to CCITT standards.
|
||||
// Copyright 2001-2010 Georges Menie (www.menie.org)
|
||||
// Copyright 2013 The Go Authors. All rights reserved.
|
||||
// http://redis.io/topics/cluster-spec#appendix-a-crc16-reference-implementation-in-ansi-c
|
||||
// https://redis.io/docs/latest/operate/oss_and_stack/reference/cluster-spec#appendix-a-crc16-reference-implementation-in-ansi-c.
|
||||
var crc16tab = [256]uint16{
|
||||
0x0000, 0x1021, 0x2042, 0x3063, 0x4084, 0x50a5, 0x60c6, 0x70e7,
|
||||
0x8108, 0x9129, 0xa14a, 0xb16b, 0xc18c, 0xd1ad, 0xe1ce, 0xf1ef,
|
||||
|
||||
+54
@@ -0,0 +1,54 @@
|
||||
package hashtag
|
||||
|
||||
import "github.com/cespare/xxhash/v2"
|
||||
|
||||
// RendezvousHash implements HRW (Highest Random Weight) hashing.
|
||||
type RendezvousHash struct {
|
||||
nodes []node
|
||||
}
|
||||
|
||||
type node struct {
|
||||
name string
|
||||
hash uint64
|
||||
}
|
||||
|
||||
// NewRendezvousHash builds a hash from shard names.
|
||||
func NewRendezvousHash(shards []string) *RendezvousHash {
|
||||
n := make([]node, len(shards))
|
||||
for i, s := range shards {
|
||||
n[i] = node{
|
||||
name: s,
|
||||
hash: xxhash.Sum64String(s),
|
||||
}
|
||||
}
|
||||
return &RendezvousHash{nodes: n}
|
||||
}
|
||||
|
||||
// Get returns the shard name for the given key.
|
||||
func (r *RendezvousHash) Get(key string) string {
|
||||
if len(r.nodes) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
kh := xxhash.Sum64String(key)
|
||||
|
||||
bestIdx := 0
|
||||
bestScore := mix64(kh ^ r.nodes[0].hash)
|
||||
|
||||
for i := 1; i < len(r.nodes); i++ {
|
||||
if score := mix64(kh ^ r.nodes[i].hash); score > bestScore {
|
||||
bestScore = score
|
||||
bestIdx = i
|
||||
}
|
||||
}
|
||||
|
||||
return r.nodes[bestIdx].name
|
||||
}
|
||||
|
||||
// mix64 is a xorshift-based mixing function.
|
||||
func mix64(x uint64) uint64 {
|
||||
x ^= x >> 12
|
||||
x ^= x << 25
|
||||
x ^= x >> 27
|
||||
return x * 2685821657736338717
|
||||
}
|
||||
+2
@@ -109,6 +109,8 @@ func (s StructValue) Scan(key string, value string) error {
|
||||
return scan.ScanRedis(value)
|
||||
case encoding.TextUnmarshaler:
|
||||
return scan.UnmarshalText(util.StringToBytes(value))
|
||||
case encoding.BinaryUnmarshaler:
|
||||
return scan.UnmarshalBinary(util.StringToBytes(value))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+19
@@ -85,6 +85,17 @@ type Recorder interface {
|
||||
// consumerGroup: name of the consumer group
|
||||
// consumerName: name of the consumer
|
||||
RecordStreamLag(ctx context.Context, lag time.Duration, cn *pool.Conn, streamName, consumerGroup, consumerName string)
|
||||
|
||||
// RecordConnectionCount records a change in connection count (UpDownCounter)
|
||||
// delta: +1 when connection added, -1 when connection removed
|
||||
// state: connection state (e.g., "idle", "used")
|
||||
// isPubSub: true if this is a PubSub connection
|
||||
RecordConnectionCount(ctx context.Context, delta int, cn *pool.Conn, state string, isPubSub bool)
|
||||
|
||||
// RecordPendingRequests records a change in pending requests (UpDownCounter)
|
||||
// delta: +1 when request starts waiting, -1 when request stops waiting
|
||||
// poolName is passed explicitly because we may not have a connection yet when request starts
|
||||
RecordPendingRequests(ctx context.Context, delta int, cn *pool.Conn, poolName string)
|
||||
}
|
||||
|
||||
type PubSubPooler interface {
|
||||
@@ -193,6 +204,12 @@ func SetGlobalRecorder(r Recorder) {
|
||||
ConnectionClosed: func(ctx context.Context, cn *pool.Conn, reason string, err error) {
|
||||
getRecorder().RecordConnectionClosed(ctx, cn, reason, err)
|
||||
},
|
||||
ConnectionCount: func(ctx context.Context, delta int, cn *pool.Conn, state string, isPubSub bool) {
|
||||
getRecorder().RecordConnectionCount(ctx, delta, cn, state, isPubSub)
|
||||
},
|
||||
PendingRequests: func(ctx context.Context, delta int, cn *pool.Conn, poolName string) {
|
||||
getRecorder().RecordPendingRequests(ctx, delta, cn, poolName)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
@@ -246,6 +263,8 @@ func (noopRecorder) RecordPubSubMessage(context.Context, *pool.Conn, string, str
|
||||
|
||||
func (noopRecorder) RecordStreamLag(context.Context, time.Duration, *pool.Conn, string, string, string) {
|
||||
}
|
||||
func (noopRecorder) RecordConnectionCount(context.Context, int, *pool.Conn, string, bool) {}
|
||||
func (noopRecorder) RecordPendingRequests(context.Context, int, *pool.Conn, string) {}
|
||||
|
||||
// RegisterPools registers connection pools with the global recorder.
|
||||
func RegisterPools(connPool pool.Pooler, pubSubPool PubSubPooler, addr string) {
|
||||
|
||||
+46
-4
@@ -14,6 +14,7 @@ import (
|
||||
"github.com/redis/go-redis/v9/internal"
|
||||
"github.com/redis/go-redis/v9/internal/maintnotifications/logs"
|
||||
"github.com/redis/go-redis/v9/internal/proto"
|
||||
uberatomic "go.uber.org/atomic"
|
||||
)
|
||||
|
||||
var noDeadline = time.Time{}
|
||||
@@ -102,11 +103,15 @@ type Conn struct {
|
||||
|
||||
pooled bool
|
||||
pubsub bool
|
||||
closed atomic.Bool
|
||||
createdAt time.Time
|
||||
expiresAt time.Time
|
||||
poolName string // Name of the pool this connection belongs to (for metrics)
|
||||
|
||||
// When a goroutine closes a connection, it usually knows the reason, so closeReason is not needed.
|
||||
// closeReason is only used when an in-use connection is closed by another goroutine,
|
||||
// to inform the goroutine using the connection why the connection was closed.
|
||||
closeReason uberatomic.String
|
||||
|
||||
// maintenanceNotifications upgrade support: relaxed timeouts during migrations/failovers
|
||||
|
||||
// Using atomic operations for lock-free access to avoid mutex contention
|
||||
@@ -576,6 +581,41 @@ func (cn *Conn) getEffectiveWriteTimeout(normalTimeout time.Duration) time.Durat
|
||||
}
|
||||
}
|
||||
|
||||
// SetOnClose installs fn as the callback invoked exactly once when this
|
||||
// connection is closed (via Conn.Close).
|
||||
//
|
||||
// IMPORTANT: SetOnClose OVERWRITES any previously installed callback — it
|
||||
// does not compose, chain, or deduplicate. A Conn has room for a single
|
||||
// onClose hook by design, because its lifecycle is bounded (a Conn is
|
||||
// created, optionally re-initialized on its own net.Conn, and then closed
|
||||
// once) and the pool's OnRemove hooks handle any registry-level cleanup
|
||||
// that must survive the net.Conn being swapped.
|
||||
//
|
||||
// This has a subtle implication for per-connection subscriptions such as
|
||||
// the unsubscribe function returned by StreamingCredentialsProvider
|
||||
// (e.g. EntraID token rotation): if SetOnClose is called twice on the
|
||||
// same Conn with DIFFERENT unsubscribe closures — for example because
|
||||
// initConn ran a second time and obtained a fresh Subscribe() —
|
||||
// the previous unsubscribe is dropped and will NEVER run, leaking a
|
||||
// subscription on the provider. Callers must therefore ensure either:
|
||||
//
|
||||
// - the provider's Subscribe is idempotent for the same listener (the
|
||||
// streaming credentials Manager deduplicates listeners by connection
|
||||
// id, so re-Subscribe returns an equivalent unsubscribe), OR
|
||||
// - the previous callback has already been invoked before SetOnClose is
|
||||
// called again.
|
||||
//
|
||||
// Design note: unlike the client-level onCloseHooks registry (see
|
||||
// redis.baseClient), there is intentionally NO named-hook dedup or
|
||||
// multi-callback support on Conn. This is a deliberate trade-off to keep
|
||||
// the Conn object slim — a pool can hold thousands of Conn values and
|
||||
// each one is a hot allocation, so paying for a sync.Mutex plus a
|
||||
// map[string]func() error per connection to support a feature that would
|
||||
// only be used by at most one subsystem today (streaming credentials) is
|
||||
// not worth the per-connection memory and allocation cost. For a single
|
||||
// Conn there is at most one meaningful close callback at any point in
|
||||
// time, and a richer registry here would not even solve the "stale
|
||||
// closure" hazard described above.
|
||||
func (cn *Conn) SetOnClose(fn func() error) {
|
||||
cn.onClose = fn
|
||||
}
|
||||
@@ -882,18 +922,20 @@ func (cn *Conn) WithWriter(
|
||||
}
|
||||
|
||||
func (cn *Conn) IsClosed() bool {
|
||||
return cn.closed.Load() || cn.stateMachine.GetState() == StateClosed
|
||||
return cn.stateMachine.GetState() == StateClosed
|
||||
}
|
||||
|
||||
func (cn *Conn) Close() error {
|
||||
cn.closed.Store(true)
|
||||
|
||||
if cn.IsClosed() {
|
||||
return nil
|
||||
}
|
||||
// Transition to CLOSED state
|
||||
cn.stateMachine.Transition(StateClosed)
|
||||
|
||||
if cn.onClose != nil {
|
||||
// ignore error
|
||||
_ = cn.onClose()
|
||||
cn.onClose = nil
|
||||
}
|
||||
|
||||
// Lock-free netConn access for better performance
|
||||
|
||||
+9
-16
@@ -297,45 +297,38 @@ func (sm *ConnStateMachine) notifyWaiters() {
|
||||
return
|
||||
}
|
||||
|
||||
// Process waiters in FIFO order until no more can be processed
|
||||
// We loop instead of recursing to avoid stack overflow and mutex issues
|
||||
// Track state locally so we only consider transitions made within this
|
||||
// call, not concurrent transitions from woken goroutines. Re-reading the
|
||||
// atomic would let a fast goroutine's Transition(StateIdle) leak into our
|
||||
// view, causing us to wake multiple waiters at once and breaking FIFO
|
||||
// execution ordering.
|
||||
currentState := sm.GetState()
|
||||
|
||||
for {
|
||||
processed := false
|
||||
|
||||
// Find the first waiter that can proceed
|
||||
for elem := sm.waiters.Front(); elem != nil; elem = elem.Next() {
|
||||
w := elem.Value.(*waiter)
|
||||
|
||||
// Read current state inside the loop to get the latest value
|
||||
currentState := sm.GetState()
|
||||
|
||||
// Check if current state is valid for this waiter
|
||||
if _, valid := w.validStates[currentState]; valid {
|
||||
// Remove from queue first
|
||||
sm.waiters.Remove(elem)
|
||||
sm.waiterCount.Add(-1)
|
||||
|
||||
// Use CAS to ensure state hasn't changed since we checked
|
||||
// This prevents race condition where another thread changes state
|
||||
// between our check and our transition
|
||||
if sm.state.CompareAndSwap(uint32(currentState), uint32(w.targetState)) {
|
||||
// Successfully transitioned - notify waiter
|
||||
w.done <- nil
|
||||
currentState = w.targetState
|
||||
processed = true
|
||||
break
|
||||
} else {
|
||||
// State changed - re-add waiter to front of queue to maintain FIFO ordering
|
||||
// This waiter was first in line and should retain priority
|
||||
sm.waiters.PushFront(w)
|
||||
sm.waiterCount.Add(1)
|
||||
// Continue to next iteration to re-read state
|
||||
currentState = sm.GetState()
|
||||
processed = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If we didn't process any waiter, we're done
|
||||
if !processed {
|
||||
break
|
||||
}
|
||||
|
||||
+400
-81
@@ -13,6 +13,41 @@ import (
|
||||
"github.com/redis/go-redis/v9/internal/rand"
|
||||
)
|
||||
|
||||
// Connection close reason constants for metrics.
|
||||
// These are used as the "reason" parameter in CloseConn() calls.
|
||||
const (
|
||||
// CloseReasonStale indicates the connection was closed because it exceeded
|
||||
// the idle timeout or max lifetime.
|
||||
CloseReasonStale = "stale"
|
||||
|
||||
// CloseReasonHookError indicates the connection was closed due to an error
|
||||
// in a pool hook (OnGet or OnPut).
|
||||
CloseReasonHookError = "hook_error"
|
||||
|
||||
// CloseReasonAuthError indicates the connection was closed due to an
|
||||
// authentication error during re-authentication.
|
||||
CloseReasonAuthError = "auth_error"
|
||||
|
||||
// CloseReasonTest is used in tests when closing connections.
|
||||
CloseReasonTest = "test"
|
||||
|
||||
// CloseReasonFailover indicates the connection was closed due to a failover event.
|
||||
CloseReasonFailover = "failover"
|
||||
)
|
||||
|
||||
// Metric state constants for connection state tracking.
|
||||
// These represent the logical state of a connection from a metrics perspective,
|
||||
// not the internal state machine state (ConnState).
|
||||
const (
|
||||
// MetricStateIdle indicates the connection is idle in the pool,
|
||||
// ready to be acquired.
|
||||
MetricStateIdle = "idle"
|
||||
|
||||
// MetricStateUsed indicates the connection is currently being used
|
||||
// by a client operation.
|
||||
MetricStateUsed = "used"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrClosed performs any operation on the closed client will return this error.
|
||||
ErrClosed = errors.New("redis: client is closed")
|
||||
@@ -69,6 +104,15 @@ var (
|
||||
// Parameters: ctx, cn, reason, err
|
||||
metricConnectionClosedCallback func(ctx context.Context, cn *Conn, reason string, err error)
|
||||
|
||||
// Global metric callback for connection count changes (UpDownCounter)
|
||||
// Parameters: ctx, delta (+1/-1), cn, state, isPubSub
|
||||
metricConnectionCountCallback func(ctx context.Context, delta int, cn *Conn, state string, isPubSub bool)
|
||||
|
||||
// Global metric callback for pending requests changes (UpDownCounter)
|
||||
// Parameters: ctx, delta (+1/-1), cn, poolName
|
||||
// poolName is passed explicitly because we may not have a connection yet when request starts
|
||||
metricPendingRequestsCallback func(ctx context.Context, delta int, cn *Conn, poolName string)
|
||||
|
||||
// errPanicInDial is returned when a panic occurs in the dial function.
|
||||
errPanicInQueuedNewConn = errors.New("panic in queuedNewConn")
|
||||
|
||||
@@ -114,6 +158,17 @@ type MetricCallbacks struct {
|
||||
|
||||
// ConnectionClosed is called when a connection is closed
|
||||
ConnectionClosed func(ctx context.Context, cn *Conn, reason string, err error)
|
||||
|
||||
// ConnectionCount is called when connection count changes (UpDownCounter)
|
||||
// delta: +1 when connection added, -1 when connection removed
|
||||
// state: connection state (e.g., "idle", "used")
|
||||
// isPubSub: true if this is a PubSub connection
|
||||
ConnectionCount func(ctx context.Context, delta int, cn *Conn, state string, isPubSub bool)
|
||||
|
||||
// PendingRequests is called when pending requests count changes (UpDownCounter)
|
||||
// delta: +1 when request starts waiting, -1 when request stops waiting
|
||||
// poolName is passed explicitly because we may not have a connection yet when request starts
|
||||
PendingRequests func(ctx context.Context, delta int, cn *Conn, poolName string)
|
||||
}
|
||||
|
||||
// SetAllMetricCallbacks sets all metric callbacks atomically.
|
||||
@@ -138,6 +193,8 @@ func SetAllMetricCallbacks(callbacks *MetricCallbacks) {
|
||||
metricMaintenanceNotificationCallback = nil
|
||||
metricConnectionWaitTimeCallback = nil
|
||||
metricConnectionClosedCallback = nil
|
||||
metricConnectionCountCallback = nil
|
||||
metricPendingRequestsCallback = nil
|
||||
return
|
||||
}
|
||||
|
||||
@@ -148,6 +205,8 @@ func SetAllMetricCallbacks(callbacks *MetricCallbacks) {
|
||||
metricMaintenanceNotificationCallback = callbacks.MaintenanceNotification
|
||||
metricConnectionWaitTimeCallback = callbacks.ConnectionWaitTime
|
||||
metricConnectionClosedCallback = callbacks.ConnectionClosed
|
||||
metricConnectionCountCallback = callbacks.ConnectionCount
|
||||
metricPendingRequestsCallback = callbacks.PendingRequests
|
||||
}
|
||||
|
||||
// getMetricConnectionStateChangeCallback returns the metric callback for connection state changes.
|
||||
@@ -223,6 +282,22 @@ func getMetricConnectionClosedCallback() func(ctx context.Context, cn *Conn, rea
|
||||
return cb
|
||||
}
|
||||
|
||||
// getMetricConnectionCountCallback returns the metric callback for connection count changes (UpDownCounter).
|
||||
func getMetricConnectionCountCallback() func(ctx context.Context, delta int, cn *Conn, state string, isPubSub bool) {
|
||||
metricCallbackMu.RLock()
|
||||
cb := metricConnectionCountCallback
|
||||
metricCallbackMu.RUnlock()
|
||||
return cb
|
||||
}
|
||||
|
||||
// getMetricPendingRequestsCallback returns the metric callback for pending requests changes (UpDownCounter).
|
||||
func getMetricPendingRequestsCallback() func(ctx context.Context, delta int, cn *Conn, poolName string) {
|
||||
metricCallbackMu.RLock()
|
||||
cb := metricPendingRequestsCallback
|
||||
metricCallbackMu.RUnlock()
|
||||
return cb
|
||||
}
|
||||
|
||||
// Stats contains pool state information and accumulated stats.
|
||||
type Stats struct {
|
||||
Hits uint32 // number of times free connection was found in the pool
|
||||
@@ -242,7 +317,7 @@ type Stats struct {
|
||||
|
||||
type Pooler interface {
|
||||
NewConn(context.Context) (*Conn, error)
|
||||
CloseConn(*Conn) error
|
||||
CloseConn(ctx context.Context, cn *Conn, reason string, fromState string) error
|
||||
|
||||
Get(context.Context) (*Conn, error)
|
||||
Put(context.Context, *Conn)
|
||||
@@ -294,6 +369,10 @@ type Options struct {
|
||||
// Default: 100ms
|
||||
DialerRetryTimeout time.Duration
|
||||
|
||||
// DialerRetryBackoff controls the delay between dial retry attempts.
|
||||
// If nil, dial retry backoff is constant and equals DialerRetryTimeout (default: 100ms).
|
||||
DialerRetryBackoff func(attempt int) time.Duration
|
||||
|
||||
// Name is a unique identifier for this pool, used in metrics.
|
||||
// Format: addr_uniqueID (e.g., "localhost:6379_a1b2c3d4")
|
||||
Name string
|
||||
@@ -456,10 +535,8 @@ func (p *ConnPool) checkMinIdleConns() {
|
||||
}
|
||||
|
||||
func (p *ConnPool) addIdleConn() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), p.cfg.DialTimeout)
|
||||
defer cancel()
|
||||
|
||||
cn, err := p.dialConn(ctx, true)
|
||||
// Do not apply DialTimeout via context here; dialConn applies DialTimeout per attempt.
|
||||
cn, err := p.dialConn(context.Background(), true)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -479,6 +556,12 @@ func (p *ConnPool) addIdleConn() error {
|
||||
|
||||
p.conns[cn.GetID()] = cn
|
||||
p.idleConns = append(p.idleConns, cn)
|
||||
|
||||
// Record connection count increment (new idle connection from min-idle prewarm)
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(context.Background(), 1, cn, "idle", false)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -505,9 +588,9 @@ func (p *ConnPool) newConn(ctx context.Context, pooled bool) (*Conn, error) {
|
||||
ctx = context.Background()
|
||||
}
|
||||
|
||||
dialCtx, cancel := context.WithTimeout(ctx, p.cfg.DialTimeout)
|
||||
defer cancel()
|
||||
cn, err := p.dialConn(dialCtx, pooled)
|
||||
// Do not apply DialTimeout via context here; dialConn applies DialTimeout per attempt.
|
||||
// We still propagate ctx so callers can cancel explicitly.
|
||||
cn, err := p.dialConn(ctx, pooled)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -543,9 +626,14 @@ func (p *ConnPool) newConn(ctx context.Context, pooled bool) (*Conn, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// Notify metrics: new connection created and idle
|
||||
// All new connections start as "used" metrically. For the miss path in getConn,
|
||||
// this is the final state. For putIdleConn (undelivered conn), a used→idle
|
||||
// transition is emitted when it's added to idleConns.
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, "", "idle")
|
||||
cb(ctx, cn, "", MetricStateUsed)
|
||||
}
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, 1, cn, "used", false)
|
||||
}
|
||||
|
||||
return cn, nil
|
||||
@@ -569,16 +657,12 @@ func (p *ConnPool) dialConn(ctx context.Context, pooled bool) (*Conn, error) {
|
||||
}
|
||||
|
||||
// Retry dialing with backoff
|
||||
// the context timeout is already handled by the context passed in
|
||||
// so we may never reach the max retries, higher values don't hurt
|
||||
// Dial timeout is applied per attempt (so retries/backoff don't eat into the next
|
||||
// attempt's dial budget), while still honoring caller cancellation via ctx.
|
||||
maxRetries := p.cfg.DialerRetries
|
||||
if maxRetries <= 0 {
|
||||
maxRetries = 5 // Default value
|
||||
}
|
||||
backoffDuration := p.cfg.DialerRetryTimeout
|
||||
if backoffDuration <= 0 {
|
||||
backoffDuration = 100 * time.Millisecond // Default value
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
shouldLoop := true
|
||||
@@ -587,16 +671,32 @@ func (p *ConnPool) dialConn(ctx context.Context, pooled bool) (*Conn, error) {
|
||||
// instead of a generic context deadline exceeded error
|
||||
attempt := 0
|
||||
for attempt = 0; (attempt < maxRetries) && shouldLoop; attempt++ {
|
||||
netConn, err := p.cfg.Dialer(ctx)
|
||||
attemptCtx := ctx
|
||||
var cancel context.CancelFunc
|
||||
if p.cfg.DialTimeout > 0 {
|
||||
// Apply DialTimeout per attempt, but never extend an existing earlier deadline.
|
||||
if deadline, ok := ctx.Deadline(); !ok || time.Until(deadline) > p.cfg.DialTimeout {
|
||||
attemptCtx, cancel = context.WithTimeout(ctx, p.cfg.DialTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
netConn, err := p.cfg.Dialer(attemptCtx)
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
// Add backoff delay for retry attempts
|
||||
// (not for the first attempt, do at least one)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
shouldLoop = false
|
||||
case <-time.After(backoffDuration):
|
||||
// Continue with retry
|
||||
// Do not sleep after the last attempt.
|
||||
if attempt+1 < maxRetries {
|
||||
backoffDuration := p.dialRetryBackoff(attempt)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
shouldLoop = false
|
||||
case <-time.After(backoffDuration):
|
||||
// Continue with retry
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -623,6 +723,22 @@ func (p *ConnPool) dialConn(ctx context.Context, pooled bool) (*Conn, error) {
|
||||
return nil, lastErr
|
||||
}
|
||||
|
||||
func (p *ConnPool) dialRetryBackoff(attempt int) time.Duration {
|
||||
if p.cfg.DialerRetryBackoff != nil {
|
||||
d := p.cfg.DialerRetryBackoff(attempt)
|
||||
if d < 0 {
|
||||
return 0
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
base := p.cfg.DialerRetryTimeout
|
||||
if base <= 0 {
|
||||
base = 100 * time.Millisecond
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
// calcConnExpiresAt calculates the expiration time for a connection.
|
||||
// It applies random jitter to prevent all connections from expiring simultaneously,
|
||||
// avoiding the "thundering herd" problem where all connections expire at once.
|
||||
@@ -648,19 +764,26 @@ func (p *ConnPool) tryDial() {
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), p.cfg.DialTimeout)
|
||||
// Probe dialing even when dialErrorsNum is saturated. Apply DialTimeout per probe
|
||||
// attempt so custom dialers can't hang indefinitely.
|
||||
ctx := context.Background()
|
||||
var cancel context.CancelFunc
|
||||
if p.cfg.DialTimeout > 0 {
|
||||
ctx, cancel = context.WithTimeout(ctx, p.cfg.DialTimeout)
|
||||
}
|
||||
|
||||
conn, err := p.cfg.Dialer(ctx)
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
if err != nil {
|
||||
p.setLastDialError(err)
|
||||
time.Sleep(time.Second)
|
||||
cancel()
|
||||
continue
|
||||
}
|
||||
|
||||
atomic.StoreUint32(&p.dialErrorsNum, 0)
|
||||
_ = conn.Close()
|
||||
cancel()
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -689,12 +812,21 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
}
|
||||
|
||||
// Track pending requests in pool stats
|
||||
// NOTE: We only track in stats, not via callback. The AsyncGauge reads stats directly.
|
||||
atomic.AddUint32(&p.stats.PendingRequests, 1)
|
||||
// Record pending request increment (UpDownCounter)
|
||||
// Pass pool name explicitly since we don't have a connection yet
|
||||
poolName := p.cfg.Name
|
||||
if cb := getMetricPendingRequestsCallback(); cb != nil {
|
||||
cb(ctx, 1, nil, poolName)
|
||||
}
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// Failed to get connection, decrement pending requests
|
||||
atomic.AddUint32(&p.stats.PendingRequests, ^uint32(0)) // -1
|
||||
// Record pending request decrement on failure
|
||||
if cb := getMetricPendingRequestsCallback(); cb != nil {
|
||||
cb(ctx, -1, nil, poolName)
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -732,6 +864,17 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
|
||||
p.connsMu.Lock()
|
||||
cn, err = p.popIdle()
|
||||
if cn != nil {
|
||||
// Emit idle→used transition inside the lock so Close() sees
|
||||
// consistent state (conn removed from idleConns = "used").
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateIdle, MetricStateUsed)
|
||||
}
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "idle", false)
|
||||
cb(ctx, 1, cn, "used", false)
|
||||
}
|
||||
}
|
||||
p.connsMu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
@@ -744,7 +887,8 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
}
|
||||
|
||||
if !p.isHealthyConn(cn, nowNs) {
|
||||
_ = p.CloseConn(cn)
|
||||
// Connection was already transitioned to MetricStateUsed under the lock above.
|
||||
_ = p.CloseConn(ctx, cn, CloseReasonStale, MetricStateUsed)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -755,11 +899,13 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
if hookErr != nil || !acceptConn {
|
||||
if hookErr != nil {
|
||||
internal.Logger.Printf(ctx, "redis: connection pool: failed to process idle connection by hook: %v", hookErr)
|
||||
_ = p.CloseConn(cn)
|
||||
// Connection was already transitioned to MetricStateUsed under the lock above.
|
||||
_ = p.CloseConn(ctx, cn, CloseReasonHookError, MetricStateUsed)
|
||||
} else {
|
||||
internal.Logger.Printf(ctx, "redis: connection pool: conn[%d] rejected by hook, returning to pool", cn.GetID())
|
||||
// Connection is already in MetricStateUsed (transitioned under the lock above).
|
||||
// Return connection to pool without freeing the turn that this Get() call holds.
|
||||
// We use putConnWithoutTurn() to run all the Put hooks and logic without freeing a turn.
|
||||
// putConnWithoutTurn will emit used→idle transition.
|
||||
p.putConnWithoutTurn(ctx, cn)
|
||||
cn = nil
|
||||
}
|
||||
@@ -769,19 +915,17 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
|
||||
atomic.AddUint32(&p.stats.Hits, 1)
|
||||
|
||||
// Notify metrics: connection moved from idle to used
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, "idle", "used")
|
||||
}
|
||||
|
||||
// Record wait time (use cached callback from above)
|
||||
if waitTimeCallback != nil {
|
||||
waitTimeCallback(ctx, waitDuration, cn)
|
||||
}
|
||||
|
||||
// Decrement pending requests (connection acquired successfully)
|
||||
// NOTE: We only track in stats, not via callback. The AsyncGauge reads stats directly.
|
||||
atomic.AddUint32(&p.stats.PendingRequests, ^uint32(0)) // -1
|
||||
// Record pending request decrement (UpDownCounter)
|
||||
if cb := getMetricPendingRequestsCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, poolName)
|
||||
}
|
||||
|
||||
return cn, nil
|
||||
}
|
||||
@@ -802,17 +946,26 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
// both errors and accept=false mean a hook rejected the connection
|
||||
// this should not happen with a new connection, but we handle it gracefully
|
||||
if err != nil || !acceptConn {
|
||||
// Failed to process connection, discard it
|
||||
internal.Logger.Printf(ctx, "redis: connection pool: failed to process new connection conn[%d] by hook: accept=%v, err=%v", newcn.GetID(), acceptConn, err)
|
||||
_ = p.CloseConn(newcn)
|
||||
// newConn emitted +1 used; CloseConn will emit -1 used if we own the removal.
|
||||
_ = p.CloseConn(ctx, newcn, CloseReasonHookError, MetricStateUsed)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Record connection creation time metric when hooks are used.
|
||||
// When hookManager is set, ProcessOnGet initializes the connection (AUTH/HELLO),
|
||||
// causing IsInited()=true. This means _getConn() in redis.go will take the
|
||||
// early return path and never reach its create time recording.
|
||||
// When hookManager is nil, _getConn() handles both initialization and create time recording.
|
||||
if dialStartNs := newcn.GetDialStartNs(); dialStartNs > 0 {
|
||||
if cb := GetMetricConnectionCreateTimeCallback(); cb != nil {
|
||||
duration := time.Duration(time.Now().UnixNano() - dialStartNs)
|
||||
cb(ctx, duration, newcn)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Notify metrics: new connection is created and used
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, newcn, "", "used")
|
||||
}
|
||||
// newConn already emitted +1 used, so no transition needed here.
|
||||
|
||||
// Record wait time (use cached callback from above)
|
||||
if waitTimeCallback != nil {
|
||||
@@ -820,8 +973,11 @@ func (p *ConnPool) getConn(ctx context.Context) (cn *Conn, err error) {
|
||||
}
|
||||
|
||||
// Decrement pending requests (connection acquired successfully)
|
||||
// NOTE: We only track in stats, not via callback. The AsyncGauge reads stats directly.
|
||||
atomic.AddUint32(&p.stats.PendingRequests, ^uint32(0)) // -1
|
||||
// Record pending request decrement (UpDownCounter)
|
||||
if cb := getMetricPendingRequestsCallback(); cb != nil {
|
||||
cb(ctx, -1, newcn, poolName)
|
||||
}
|
||||
|
||||
return newcn, nil
|
||||
}
|
||||
@@ -835,7 +991,8 @@ func (p *ConnPool) queuedNewConn(ctx context.Context) (*Conn, error) {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
dialCtx, cancel := context.WithTimeout(context.Background(), p.cfg.DialTimeout)
|
||||
// Don't apply DialTimeout via context here; dialConn applies DialTimeout per attempt.
|
||||
dialCtx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
w := &wantConn{
|
||||
ctx: dialCtx,
|
||||
@@ -919,14 +1076,24 @@ func (p *ConnPool) putIdleConn(ctx context.Context, cn *Conn) bool {
|
||||
defer p.connsMu.Unlock()
|
||||
|
||||
if p.closed() {
|
||||
_ = cn.Close()
|
||||
// Don't close here — this connection is still in p.conns and Close()
|
||||
// will handle closing it and emitting the correct metric decrements.
|
||||
// We just skip adding it to idleConns.
|
||||
return true
|
||||
}
|
||||
|
||||
// poolSize is increased in newConn
|
||||
p.idleConns = append(p.idleConns, cn)
|
||||
p.idleConnsLen.Add(1)
|
||||
|
||||
// Connection was created as "used" in newConn; transition to idle.
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateUsed, MetricStateIdle)
|
||||
}
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
cb(ctx, 1, cn, "idle", false)
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -1087,6 +1254,7 @@ func (p *ConnPool) putConn(ctx context.Context, cn *Conn, freeTurn bool) {
|
||||
}
|
||||
|
||||
var shouldCloseConn bool
|
||||
var removedFromPool bool
|
||||
|
||||
if p.cfg.MaxIdleConns == 0 || p.idleConnsLen.Load() < p.cfg.MaxIdleConns {
|
||||
// Hot path optimization: try fast IN_USE → IDLE transition
|
||||
@@ -1111,7 +1279,7 @@ func (p *ConnPool) putConn(ctx context.Context, cn *Conn, freeTurn bool) {
|
||||
case StateClosed:
|
||||
internal.Logger.Printf(ctx, "Unexpected conn[%d] state changed by hook to %v, closing it", cn.GetID(), currentState)
|
||||
shouldCloseConn = true
|
||||
p.removeConnWithLock(cn)
|
||||
removedFromPool = p.removeConnWithLock(cn)
|
||||
default:
|
||||
// Pool as-is
|
||||
internal.Logger.Printf(ctx, "Unexpected conn[%d] state changed by hook to %v, pooling as-is", cn.GetID(), currentState)
|
||||
@@ -1122,34 +1290,73 @@ func (p *ConnPool) putConn(ctx context.Context, cn *Conn, freeTurn bool) {
|
||||
// put them at the opposite end of the queue
|
||||
// Optimization: if we just transitioned to IDLE, we know it's usable - skip the check
|
||||
if !transitionedToIdle && !cn.IsUsable() {
|
||||
if p.cfg.PoolFIFO {
|
||||
p.connsMu.Lock()
|
||||
p.idleConns = append(p.idleConns, cn)
|
||||
p.connsMu.Lock()
|
||||
// Check if Close() already removed this connection from p.conns.
|
||||
// If so, skip the append and metrics — Close() already accounted for it.
|
||||
if _, inPool := p.conns[cn.GetID()]; inPool {
|
||||
if p.cfg.PoolFIFO {
|
||||
p.idleConns = append(p.idleConns, cn)
|
||||
} else {
|
||||
p.idleConns = append([]*Conn{cn}, p.idleConns...)
|
||||
}
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateUsed, MetricStateIdle)
|
||||
}
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
cb(ctx, 1, cn, "idle", false)
|
||||
}
|
||||
p.connsMu.Unlock()
|
||||
p.idleConnsLen.Add(1)
|
||||
} else {
|
||||
p.connsMu.Lock()
|
||||
p.idleConns = append([]*Conn{cn}, p.idleConns...)
|
||||
shouldCloseConn = true
|
||||
p.connsMu.Unlock()
|
||||
}
|
||||
p.idleConnsLen.Add(1)
|
||||
} else if !shouldCloseConn {
|
||||
p.connsMu.Lock()
|
||||
p.idleConns = append(p.idleConns, cn)
|
||||
p.connsMu.Unlock()
|
||||
p.idleConnsLen.Add(1)
|
||||
if _, inPool := p.conns[cn.GetID()]; inPool {
|
||||
p.idleConns = append(p.idleConns, cn)
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateUsed, MetricStateIdle)
|
||||
}
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
cb(ctx, 1, cn, "idle", false)
|
||||
}
|
||||
p.connsMu.Unlock()
|
||||
p.idleConnsLen.Add(1)
|
||||
} else {
|
||||
shouldCloseConn = true
|
||||
p.connsMu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
// Notify metrics: connection moved from used to idle
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, "used", "idle")
|
||||
if shouldCloseConn {
|
||||
// Connection was removed (e.g., hook set state to StateClosed).
|
||||
// Only emit if we actually removed it from the map (not already taken by Close()).
|
||||
if removedFromPool {
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateUsed, "")
|
||||
}
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
shouldCloseConn = true
|
||||
p.removeConnWithLock(cn)
|
||||
removedFromPool = p.removeConnWithLock(cn)
|
||||
|
||||
// Notify metrics: connection removed (used -> nothing)
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, "used", "")
|
||||
// Only emit if we actually removed it from the map (not already taken by Close()).
|
||||
if removedFromPool {
|
||||
// Notify metrics: connection removed (used -> nothing)
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateUsed, "")
|
||||
}
|
||||
// Record connection count decrement (connection removed while in used state)
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1158,6 +1365,17 @@ func (p *ConnPool) putConn(ctx context.Context, cn *Conn, freeTurn bool) {
|
||||
}
|
||||
|
||||
if shouldCloseConn {
|
||||
// Only emit connection closed if we actually owned the removal.
|
||||
// If removedFromPool is false, Close() already emitted connectionClosed for this conn.
|
||||
if removedFromPool {
|
||||
if cb := getMetricConnectionClosedCallback(); cb != nil {
|
||||
reason := "conn_pool_close"
|
||||
if r := cn.closeReason.Load(); r != "" {
|
||||
reason = r
|
||||
}
|
||||
cb(ctx, cn, reason, nil)
|
||||
}
|
||||
}
|
||||
_ = p.closeConn(cn)
|
||||
}
|
||||
|
||||
@@ -1185,24 +1403,35 @@ func (p *ConnPool) removeConnInternal(ctx context.Context, cn *Conn, reason erro
|
||||
hookManager.ProcessOnRemove(ctx, cn, reason)
|
||||
}
|
||||
|
||||
p.removeConnWithLock(cn)
|
||||
removed := p.removeConnWithLock(cn)
|
||||
|
||||
if freeTurn {
|
||||
p.freeTurn()
|
||||
}
|
||||
|
||||
// Notify metrics: connection removed (assume from used state)
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, "used", "")
|
||||
// Only emit metric decrements if we actually removed the connection from the map.
|
||||
// If removed is false, Close() already removed it and emitted the -1 delta.
|
||||
if removed {
|
||||
// Notify metrics: connection removed (assume from used state)
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil {
|
||||
cb(ctx, cn, MetricStateUsed, "")
|
||||
}
|
||||
// Record connection count decrement (connection removed, assume from used state)
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
}
|
||||
}
|
||||
|
||||
// Record connection closed
|
||||
if cb := getMetricConnectionClosedCallback(); cb != nil {
|
||||
reasonStr := "unknown"
|
||||
if reason != nil {
|
||||
reasonStr = reason.Error()
|
||||
// Only emit connection closed if we actually owned the removal.
|
||||
// If removed is false, Close() already emitted connectionClosed for this conn.
|
||||
if removed {
|
||||
if cb := getMetricConnectionClosedCallback(); cb != nil {
|
||||
reasonStr := "unknown"
|
||||
if reason != nil {
|
||||
reasonStr = reason.Error()
|
||||
}
|
||||
cb(ctx, cn, reasonStr, reason)
|
||||
}
|
||||
cb(ctx, cn, reasonStr, reason)
|
||||
}
|
||||
|
||||
_ = p.closeConn(cn)
|
||||
@@ -1211,19 +1440,60 @@ func (p *ConnPool) removeConnInternal(ctx context.Context, cn *Conn, reason erro
|
||||
p.checkMinIdleConns()
|
||||
}
|
||||
|
||||
func (p *ConnPool) CloseConn(cn *Conn) error {
|
||||
p.removeConnWithLock(cn)
|
||||
// CloseConn closes a connection and records metrics.
|
||||
// Parameters:
|
||||
// - ctx: context for metric callbacks (enables trace-to-metric correlation)
|
||||
// - cn: the connection to close
|
||||
// - reason: why the connection is being closed (use CloseReason* constants)
|
||||
// - fromState: the metric state the connection was in (use MetricState* constants)
|
||||
func (p *ConnPool) CloseConn(ctx context.Context, cn *Conn, reason string, fromState string) error {
|
||||
removed := p.removeConnWithLock(cn)
|
||||
|
||||
// Only emit UpDownCounter decrements if we actually removed the connection.
|
||||
// If removed is false, Close() already removed it and emitted the -1 delta.
|
||||
// Only emit connection closed if we actually owned the removal.
|
||||
// If removed is false, Close() already emitted connectionClosed for this conn.
|
||||
if removed {
|
||||
p.recordConnectionMetrics(ctx, cn, reason, fromState)
|
||||
}
|
||||
|
||||
return p.closeConn(cn)
|
||||
}
|
||||
|
||||
func (p *ConnPool) removeConnWithLock(cn *Conn) {
|
||||
p.connsMu.Lock()
|
||||
defer p.connsMu.Unlock()
|
||||
p.removeConn(cn)
|
||||
func (p *ConnPool) recordConnectionMetrics(ctx context.Context, cn *Conn, reason string, fromState string) {
|
||||
// Record connection state change: connection is being removed from the specified state
|
||||
if cb := getMetricConnectionStateChangeCallback(); cb != nil && fromState != "" {
|
||||
cb(ctx, cn, fromState, "")
|
||||
}
|
||||
|
||||
// Record connection count decrement (UpDownCounter) for the state the connection was in
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil && fromState != "" {
|
||||
cb(ctx, -1, cn, fromState, false)
|
||||
}
|
||||
|
||||
if cb := getMetricConnectionClosedCallback(); cb != nil {
|
||||
cb(ctx, cn, reason, nil)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *ConnPool) removeConn(cn *Conn) {
|
||||
// removeConnWithLock removes a connection from the pool under the connsMu lock.
|
||||
// Returns true if the connection was actually present in p.conns and was removed,
|
||||
// false if it was already gone (e.g., removed by Close()). Callers must use the
|
||||
// return value to decide whether to emit metric decrements — this eliminates the
|
||||
// shutdown race between Close() and concurrent removal paths.
|
||||
func (p *ConnPool) removeConnWithLock(cn *Conn) bool {
|
||||
p.connsMu.Lock()
|
||||
defer p.connsMu.Unlock()
|
||||
return p.removeConn(cn)
|
||||
}
|
||||
|
||||
// removeConn removes a connection from the pool's internal data structures.
|
||||
// Returns true if the connection was present and removed, false otherwise.
|
||||
func (p *ConnPool) removeConn(cn *Conn) bool {
|
||||
cid := cn.GetID()
|
||||
if _, exists := p.conns[cid]; !exists {
|
||||
return false
|
||||
}
|
||||
delete(p.conns, cid)
|
||||
atomic.AddUint32(&p.stats.StaleConns, 1)
|
||||
|
||||
@@ -1239,6 +1509,7 @@ func (p *ConnPool) removeConn(cn *Conn) {
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (p *ConnPool) closeConn(cn *Conn) error {
|
||||
@@ -1290,13 +1561,33 @@ func (p *ConnPool) closed() bool {
|
||||
}
|
||||
|
||||
func (p *ConnPool) Filter(fn func(*Conn) bool) error {
|
||||
ctx := context.Background()
|
||||
|
||||
p.connsMu.Lock()
|
||||
defer p.connsMu.Unlock()
|
||||
|
||||
idleConnSet := make(map[*Conn]struct{}, len(p.idleConns))
|
||||
for _, ic := range p.idleConns {
|
||||
idleConnSet[ic] = struct{}{}
|
||||
}
|
||||
|
||||
var firstErr error
|
||||
for _, cn := range p.conns {
|
||||
if fn(cn) {
|
||||
if err := p.closeConn(cn); err != nil && firstErr == nil {
|
||||
var err error
|
||||
if _, isIdle := idleConnSet[cn]; isIdle {
|
||||
// Idle connection - remove from pool and close.
|
||||
p.removeConn(cn)
|
||||
p.recordConnectionMetrics(ctx, cn, CloseReasonFailover, MetricStateIdle)
|
||||
err = p.closeConn(cn)
|
||||
} else {
|
||||
// Used connection - set closeReason and close the connection.
|
||||
// The connection remains in p.conns. When putConn() is called later,
|
||||
// it will close the connection instead of pooling it.
|
||||
cn.closeReason.Store(CloseReasonFailover)
|
||||
err = cn.Close()
|
||||
}
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
@@ -1310,10 +1601,38 @@ func (p *ConnPool) Close() error {
|
||||
}
|
||||
|
||||
var firstErr error
|
||||
nowNs := time.Now().UnixNano()
|
||||
p.connsMu.Lock()
|
||||
|
||||
// Emit -1 for each connection. Since all idle↔used transitions happen
|
||||
// under connsMu, the idleConns slice is the source of truth for state.
|
||||
cb := getMetricConnectionCountCallback()
|
||||
idleSet := make(map[uint64]struct{}, len(p.idleConns))
|
||||
for _, cn := range p.idleConns {
|
||||
idleSet[cn.GetID()] = struct{}{}
|
||||
}
|
||||
ctx := context.Background()
|
||||
for _, cn := range p.conns {
|
||||
// Check health before closing, since closeConn invalidates the
|
||||
// underlying fd and would make connCheck (inside isHealthyConn)
|
||||
// always fail with EBADF.
|
||||
healthy := p.isHealthyConn(cn, nowNs)
|
||||
if cb != nil {
|
||||
if _, isIdle := idleSet[cn.GetID()]; isIdle {
|
||||
cb(ctx, -1, cn, "idle", false)
|
||||
} else {
|
||||
cb(ctx, -1, cn, "used", false)
|
||||
}
|
||||
}
|
||||
if closedCb := getMetricConnectionClosedCallback(); closedCb != nil {
|
||||
closedCb(ctx, cn, "pool_shutdown", nil)
|
||||
}
|
||||
if err := p.closeConn(cn); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
// Suppress close errors for stale connections, consistent
|
||||
// with how Get() handles them (see CloseReasonStale path).
|
||||
if healthy {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
}
|
||||
p.conns = nil
|
||||
|
||||
+2
-2
@@ -33,8 +33,8 @@ func (p *SingleConnPool) NewConn(ctx context.Context) (*Conn, error) {
|
||||
return p.pool.NewConn(ctx)
|
||||
}
|
||||
|
||||
func (p *SingleConnPool) CloseConn(cn *Conn) error {
|
||||
return p.pool.CloseConn(cn)
|
||||
func (p *SingleConnPool) CloseConn(ctx context.Context, cn *Conn, reason string, fromState string) error {
|
||||
return p.pool.CloseConn(ctx, cn, reason, fromState)
|
||||
}
|
||||
|
||||
func (p *SingleConnPool) Get(_ context.Context) (*Conn, error) {
|
||||
|
||||
+2
-2
@@ -61,8 +61,8 @@ func (p *StickyConnPool) NewConn(ctx context.Context) (*Conn, error) {
|
||||
return p.pool.NewConn(ctx)
|
||||
}
|
||||
|
||||
func (p *StickyConnPool) CloseConn(cn *Conn) error {
|
||||
return p.pool.CloseConn(cn)
|
||||
func (p *StickyConnPool) CloseConn(ctx context.Context, cn *Conn, reason string, fromState string) error {
|
||||
return p.pool.CloseConn(ctx, cn, reason, fromState)
|
||||
}
|
||||
|
||||
func (p *StickyConnPool) Get(ctx context.Context) (*Conn, error) {
|
||||
|
||||
+25
-1
@@ -53,18 +53,42 @@ func (p *PubSubPool) NewConn(ctx context.Context, network string, addr string, c
|
||||
func (p *PubSubPool) TrackConn(cn *Conn) {
|
||||
atomic.AddUint32(&p.stats.Active, 1)
|
||||
p.activeConns.Store(cn.GetID(), cn)
|
||||
// Emit +1 used for PubSub connection
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(context.Background(), 1, cn, "used", true)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PubSubPool) UntrackConn(cn *Conn) {
|
||||
// LoadAndDelete ensures each connection is only decremented once,
|
||||
// guarding against double-decrement if Close() already untracked it.
|
||||
if _, loaded := p.activeConns.LoadAndDelete(cn.GetID()); !loaded {
|
||||
return
|
||||
}
|
||||
atomic.AddUint32(&p.stats.Active, ^uint32(0))
|
||||
atomic.AddUint32(&p.stats.Untracked, 1)
|
||||
p.activeConns.Delete(cn.GetID())
|
||||
// Emit -1 used for PubSub connection
|
||||
if cb := getMetricConnectionCountCallback(); cb != nil {
|
||||
cb(context.Background(), -1, cn, "used", true)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *PubSubPool) Close() error {
|
||||
p.closed.Store(true)
|
||||
cb := getMetricConnectionCountCallback()
|
||||
p.activeConns.Range(func(key, value interface{}) bool {
|
||||
cn := value.(*Conn)
|
||||
// Use LoadAndDelete to atomically claim ownership of this entry.
|
||||
// If a concurrent UntrackConn already removed it, skip to avoid double-decrement.
|
||||
if _, loaded := p.activeConns.LoadAndDelete(key); !loaded {
|
||||
return true
|
||||
}
|
||||
atomic.AddUint32(&p.stats.Active, ^uint32(0))
|
||||
atomic.AddUint32(&p.stats.Untracked, 1)
|
||||
// Emit -1 used for each PubSub connection being closed
|
||||
if cb != nil {
|
||||
cb(context.Background(), -1, cn, "used", true)
|
||||
}
|
||||
_ = cn.Close()
|
||||
return true
|
||||
})
|
||||
|
||||
+195
-5
@@ -279,8 +279,8 @@ func (r *Reader) ReadReply() (interface{}, error) {
|
||||
}
|
||||
|
||||
func (r *Reader) readFloat(line []byte) (float64, error) {
|
||||
v := string(line[1:])
|
||||
switch string(line[1:]) {
|
||||
v := util.BytesToString(line[1:])
|
||||
switch v {
|
||||
case "inf":
|
||||
return math.Inf(1), nil
|
||||
case "-inf":
|
||||
@@ -292,7 +292,7 @@ func (r *Reader) readFloat(line []byte) (float64, error) {
|
||||
}
|
||||
|
||||
func (r *Reader) readBool(line []byte) (bool, error) {
|
||||
switch string(line[1:]) {
|
||||
switch util.BytesToString(line[1:]) {
|
||||
case "t":
|
||||
return true, nil
|
||||
case "f":
|
||||
@@ -303,7 +303,7 @@ func (r *Reader) readBool(line []byte) (bool, error) {
|
||||
|
||||
func (r *Reader) readBigInt(line []byte) (*big.Int, error) {
|
||||
i := new(big.Int)
|
||||
if i, ok := i.SetString(string(line[1:]), 10); ok {
|
||||
if i, ok := i.SetString(util.BytesToString(line[1:]), 10); ok {
|
||||
return i, nil
|
||||
}
|
||||
return nil, fmt.Errorf("redis: can't parse bigInt reply: %q", line)
|
||||
@@ -453,7 +453,7 @@ func (r *Reader) ReadFloat() (float64, error) {
|
||||
case RespFloat:
|
||||
return r.readFloat(line)
|
||||
case RespStatus:
|
||||
return strconv.ParseFloat(string(line[1:]), 64)
|
||||
return strconv.ParseFloat(util.BytesToString(line[1:]), 64)
|
||||
case RespString:
|
||||
s, err := r.readStringReply(line)
|
||||
if err != nil {
|
||||
@@ -646,3 +646,193 @@ func IsNilReply(line []byte) bool {
|
||||
(line[0] == RespString || line[0] == RespArray) &&
|
||||
line[1] == '-' && line[2] == '1'
|
||||
}
|
||||
|
||||
// ReadRawReply reads the next RESP reply and returns it as raw bytes without parsing.
|
||||
func (r *Reader) ReadRawReply() ([]byte, error) {
|
||||
return r.readRawReplyBuf(nil)
|
||||
}
|
||||
|
||||
func (r *Reader) readRawReplyBuf(buf []byte) ([]byte, error) {
|
||||
line, err := r.readLine()
|
||||
if err != nil {
|
||||
return buf, err
|
||||
}
|
||||
|
||||
buf = append(buf, line...)
|
||||
buf = append(buf, '\r', '\n')
|
||||
|
||||
switch line[0] {
|
||||
case RespStatus, RespError, RespInt, RespNil, RespFloat, RespBool, RespBigInt:
|
||||
return buf, nil
|
||||
|
||||
case RespString, RespVerbatim, RespBlobError:
|
||||
n, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return buf, nil
|
||||
}
|
||||
return buf, err
|
||||
}
|
||||
curLen := len(buf)
|
||||
buf = append(buf, make([]byte, n+2)...)
|
||||
_, err = io.ReadFull(r.rd, buf[curLen:])
|
||||
return buf, err
|
||||
|
||||
case RespArray, RespSet, RespPush:
|
||||
n, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return buf, nil
|
||||
}
|
||||
return buf, err
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
buf, err = r.readRawReplyBuf(buf)
|
||||
if err != nil {
|
||||
return buf, err
|
||||
}
|
||||
}
|
||||
return buf, nil
|
||||
|
||||
case RespMap:
|
||||
n, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return buf, nil
|
||||
}
|
||||
return buf, err
|
||||
}
|
||||
for i := 0; i < n*2; i++ {
|
||||
buf, err = r.readRawReplyBuf(buf)
|
||||
if err != nil {
|
||||
return buf, err
|
||||
}
|
||||
}
|
||||
return buf, nil
|
||||
|
||||
case RespAttr:
|
||||
// Per RESP3 spec, an attribute is always followed by the actual command reply.
|
||||
// We need to read the attribute's key-value pairs AND the following reply.
|
||||
n, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return buf, nil
|
||||
}
|
||||
return buf, err
|
||||
}
|
||||
// Read the attribute key-value pairs
|
||||
for i := 0; i < n*2; i++ {
|
||||
buf, err = r.readRawReplyBuf(buf)
|
||||
if err != nil {
|
||||
return buf, err
|
||||
}
|
||||
}
|
||||
// Read the command reply that follows the attribute
|
||||
return r.readRawReplyBuf(buf)
|
||||
}
|
||||
|
||||
return buf, fmt.Errorf("redis: can't read raw reply: %.100q", line)
|
||||
}
|
||||
|
||||
var crlf = []byte{'\r', '\n'}
|
||||
|
||||
// ReadRawReplyWriteTo streams the next RESP reply directly to w without intermediate allocations.
|
||||
// Returns the number of bytes written and any error encountered.
|
||||
func (r *Reader) ReadRawReplyWriteTo(w io.Writer) (int64, error) {
|
||||
return r.readRawReplyWriteTo(w)
|
||||
}
|
||||
|
||||
func (r *Reader) readRawReplyWriteTo(w io.Writer) (int64, error) {
|
||||
line, err := r.readLine()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
var written int64
|
||||
n, err := w.Write(line)
|
||||
written += int64(n)
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
n, err = w.Write(crlf)
|
||||
written += int64(n)
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
|
||||
switch line[0] {
|
||||
case RespStatus, RespError, RespInt, RespNil, RespFloat, RespBool, RespBigInt:
|
||||
return written, nil
|
||||
|
||||
case RespString, RespVerbatim, RespBlobError:
|
||||
dataLen, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return written, nil
|
||||
}
|
||||
return written, err
|
||||
}
|
||||
copied, err := io.CopyN(w, r.rd, int64(dataLen)+2)
|
||||
written += copied
|
||||
return written, err
|
||||
|
||||
case RespArray, RespSet, RespPush:
|
||||
count, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return written, nil
|
||||
}
|
||||
return written, err
|
||||
}
|
||||
for i := 0; i < count; i++ {
|
||||
n, err := r.readRawReplyWriteTo(w)
|
||||
written += n
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
|
||||
case RespMap:
|
||||
count, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return written, nil
|
||||
}
|
||||
return written, err
|
||||
}
|
||||
for i := 0; i < count*2; i++ {
|
||||
n, err := r.readRawReplyWriteTo(w)
|
||||
written += n
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
}
|
||||
return written, nil
|
||||
|
||||
case RespAttr:
|
||||
// Per RESP3 spec, an attribute is always followed by the actual command reply.
|
||||
// We need to read the attribute's key-value pairs AND the following reply.
|
||||
count, err := replyLen(line)
|
||||
if err != nil {
|
||||
if err == Nil {
|
||||
return written, nil
|
||||
}
|
||||
return written, err
|
||||
}
|
||||
// Read the attribute key-value pairs
|
||||
for i := 0; i < count*2; i++ {
|
||||
n, err := r.readRawReplyWriteTo(w)
|
||||
written += n
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
}
|
||||
// Read the command reply that follows the attribute
|
||||
n, err := r.readRawReplyWriteTo(w)
|
||||
written += n
|
||||
return written, err
|
||||
}
|
||||
|
||||
return written, fmt.Errorf("redis: can't read raw reply: %.100q", line)
|
||||
}
|
||||
|
||||
+16
-4
@@ -309,13 +309,25 @@ func IsReadOnlyError(err error) bool {
|
||||
if errors.As(err, &readOnlyErr) {
|
||||
return true
|
||||
}
|
||||
// Check if wrapped error is a RedisError with READONLY prefix
|
||||
// Check if wrapped error is a RedisError with READONLY prefix or Lua script READONLY
|
||||
var redisErr RedisError
|
||||
if errors.As(err, &redisErr) && strings.HasPrefix(redisErr.Error(), "READONLY ") {
|
||||
return true
|
||||
if errors.As(err, &redisErr) {
|
||||
s := redisErr.Error()
|
||||
if strings.HasPrefix(s, "READONLY ") {
|
||||
return true
|
||||
}
|
||||
// Lua script wrapped READONLY errors:
|
||||
// "ERR Error running script (call to f_<sha>): @user_script:N: -READONLY You can't write against a read only replica."
|
||||
if strings.Contains(s, "-READONLY You can't write against a read only replica") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
// Fallback to string checking for backward compatibility
|
||||
return strings.HasPrefix(err.Error(), "READONLY ")
|
||||
s := err.Error()
|
||||
if strings.HasPrefix(s, "READONLY ") {
|
||||
return true
|
||||
}
|
||||
return strings.Contains(s, "-READONLY You can't write against a read only replica")
|
||||
}
|
||||
|
||||
// IsMovedError checks if an error is a MovedError, even if wrapped.
|
||||
|
||||
+88
-43
@@ -4,7 +4,6 @@ import (
|
||||
"context"
|
||||
"net"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9/internal"
|
||||
@@ -364,20 +363,46 @@ func (c *Config) applyWorkerDefaults(poolSize int) {
|
||||
}
|
||||
}
|
||||
|
||||
// endpointDetectResolveTimeout bounds the DNS lookup performed by
|
||||
// DetectEndpointType so a slow or broken resolver cannot block client
|
||||
// construction for the full system resolver timeout (often 5-30s).
|
||||
const endpointDetectResolveTimeout = 2 * time.Second
|
||||
|
||||
// cgnatNet is RFC6598 shared address space (100.64.0.0/10), used by many
|
||||
// cloud/carrier NATs and not covered by net.IP.IsPrivate.
|
||||
var cgnatNet = &net.IPNet{IP: net.IPv4(100, 64, 0, 0), Mask: net.CIDRMask(10, 32)}
|
||||
|
||||
// isPrivateIP reports whether ip belongs to a range that should be treated
|
||||
// as "internal" for the purpose of endpoint type detection. It extends
|
||||
// net.IP.IsPrivate (RFC1918 + RFC4193) with loopback, link-local and
|
||||
// RFC6598 shared address space (CGNAT).
|
||||
func isPrivateIP(ip net.IP) bool {
|
||||
if ip == nil {
|
||||
return false
|
||||
}
|
||||
if ip.IsPrivate() || ip.IsLoopback() || ip.IsLinkLocalUnicast() {
|
||||
return true
|
||||
}
|
||||
if v4 := ip.To4(); v4 != nil && cgnatNet.Contains(v4) {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DetectEndpointType automatically detects the appropriate endpoint type
|
||||
// based on the connection address and TLS configuration.
|
||||
//
|
||||
// For IP addresses:
|
||||
// TLS behaviour:
|
||||
// - If TLS is enabled: requests FQDN for proper certificate validation
|
||||
// - If TLS is disabled: requests IP for better performance
|
||||
//
|
||||
// For hostnames:
|
||||
// - If TLS is enabled: always requests FQDN for proper certificate validation
|
||||
// - If TLS is disabled: requests IP for better performance
|
||||
// (SNI / hostname verification).
|
||||
// - If TLS is disabled: always requests IP for better performance, even
|
||||
// when the configured address is a hostname. In that case the hostname
|
||||
// is resolved to determine whether it belongs to an internal or
|
||||
// external network range.
|
||||
//
|
||||
// Internal vs External detection:
|
||||
// - For IPs: uses private IP range detection
|
||||
// - For hostnames: uses heuristics based on common internal naming patterns
|
||||
// - For hostnames: resolves the hostname to an IP address and uses the IP range detection
|
||||
func DetectEndpointType(addr string, tlsEnabled bool) EndpointType {
|
||||
// Extract host from "host:port" format
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
@@ -385,6 +410,16 @@ func DetectEndpointType(addr string, tlsEnabled bool) EndpointType {
|
||||
host = addr // Assume no port
|
||||
}
|
||||
|
||||
// An empty host (e.g., ":6379") conventionally means the loopback
|
||||
// interface and is treated as internal. With TLS off we return an IP
|
||||
// endpoint; with TLS on the caller still needs an FQDN for SNI.
|
||||
if host == "" {
|
||||
if tlsEnabled {
|
||||
return EndpointTypeInternalFQDN
|
||||
}
|
||||
return EndpointTypeInternalIP
|
||||
}
|
||||
|
||||
// Check if the host is an IP address or hostname
|
||||
ip := net.ParseIP(host)
|
||||
isIPAddress := ip != nil
|
||||
@@ -392,7 +427,7 @@ func DetectEndpointType(addr string, tlsEnabled bool) EndpointType {
|
||||
|
||||
if isIPAddress {
|
||||
// Address is an IP - determine if it's private or public
|
||||
isPrivate := ip.IsPrivate() || ip.IsLoopback() || ip.IsLinkLocalUnicast()
|
||||
isPrivate := isPrivateIP(ip)
|
||||
|
||||
if tlsEnabled {
|
||||
// TLS with IP addresses - still prefer FQDN for certificate validation
|
||||
@@ -410,48 +445,58 @@ func DetectEndpointType(addr string, tlsEnabled bool) EndpointType {
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Address is a hostname
|
||||
isInternalHostname := isInternalHostname(host)
|
||||
if isInternalHostname {
|
||||
endpointType = EndpointTypeInternalFQDN
|
||||
// Address is a hostname - resolve it under a bounded timeout so a
|
||||
// slow/broken DNS server cannot stall client construction.
|
||||
ctx, cancel := context.WithTimeout(context.Background(), endpointDetectResolveTimeout)
|
||||
defer cancel()
|
||||
|
||||
isInternal, err := isInternalHostname(ctx, host)
|
||||
// Will fallback to external classification if we can't determine
|
||||
// whether the hostname is internal.
|
||||
if err != nil && internal.LogLevel.WarnOrAbove() {
|
||||
internal.Logger.Printf(ctx, "Failed to determine if hostname %q is internal: %v", host, err)
|
||||
}
|
||||
|
||||
if tlsEnabled {
|
||||
// With TLS the server name must be preserved for certificate
|
||||
// validation, so request an FQDN endpoint.
|
||||
if isInternal {
|
||||
endpointType = EndpointTypeInternalFQDN
|
||||
} else {
|
||||
endpointType = EndpointTypeExternalFQDN
|
||||
}
|
||||
} else {
|
||||
endpointType = EndpointTypeExternalFQDN
|
||||
// Without TLS we always prefer IP endpoints for performance,
|
||||
// even if the configured address is a hostname.
|
||||
if isInternal {
|
||||
endpointType = EndpointTypeInternalIP
|
||||
} else {
|
||||
endpointType = EndpointTypeExternalIP
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return endpointType
|
||||
}
|
||||
|
||||
// isInternalHostname determines if a hostname appears to be internal/private.
|
||||
// This is a heuristic based on common naming patterns.
|
||||
func isInternalHostname(hostname string) bool {
|
||||
// Convert to lowercase for comparison
|
||||
hostname = strings.ToLower(hostname)
|
||||
|
||||
// Common internal hostname patterns
|
||||
internalPatterns := []string{
|
||||
"localhost",
|
||||
".local",
|
||||
".internal",
|
||||
".corp",
|
||||
".lan",
|
||||
".intranet",
|
||||
".private",
|
||||
// isInternalHostname resolves the hostname (both IPv4 and IPv6) under the
|
||||
// given context and reports whether every resolved address is in a
|
||||
// private/internal range. If any address is public the hostname is treated
|
||||
// as external. A resolution error returns (false, err). An empty result set
|
||||
// returns (false, nil); callers are expected to fall back to an external
|
||||
// classification when the hostname cannot be determined to be internal.
|
||||
func isInternalHostname(ctx context.Context, hostname string) (bool, error) {
|
||||
ips, err := net.DefaultResolver.LookupIPAddr(ctx, hostname)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// Check for exact match or suffix match
|
||||
for _, pattern := range internalPatterns {
|
||||
if hostname == pattern || strings.HasSuffix(hostname, pattern) {
|
||||
return true
|
||||
if len(ips) == 0 {
|
||||
return false, nil
|
||||
}
|
||||
for _, ia := range ips {
|
||||
if !isPrivateIP(ia.IP) {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Check for RFC 1918 style hostnames (e.g., redis-1, db-server, etc.)
|
||||
// If hostname doesn't contain dots, it's likely internal
|
||||
if !strings.Contains(hostname, ".") {
|
||||
return true
|
||||
}
|
||||
|
||||
// Default to external for fully qualified domain names
|
||||
return false
|
||||
return true, nil
|
||||
}
|
||||
|
||||
+13
-6
@@ -5,10 +5,11 @@ import (
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"maps"
|
||||
"net"
|
||||
"net/url"
|
||||
"runtime"
|
||||
"sort"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
@@ -143,6 +144,13 @@ type Options struct {
|
||||
// default: 100 milliseconds
|
||||
DialerRetryTimeout time.Duration
|
||||
|
||||
// DialerRetryBackoff controls the delay between dial retry attempts.
|
||||
//
|
||||
// attempt is 0-based: attempt=0 is the delay after the 1st failed dial (before the 2nd attempt).
|
||||
//
|
||||
// If nil, dial retry backoff is constant and equals DialerRetryTimeout (default: 100ms).
|
||||
DialerRetryBackoff func(attempt int) time.Duration
|
||||
|
||||
// ReadTimeout for socket reads. If reached, commands will fail
|
||||
// with a timeout instead of blocking. Supported values:
|
||||
//
|
||||
@@ -644,11 +652,8 @@ func (o *queryOptions) remaining() []string {
|
||||
if len(o.q) == 0 {
|
||||
return nil
|
||||
}
|
||||
keys := make([]string, 0, len(o.q))
|
||||
for k := range o.q {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
keys := slices.Collect(maps.Keys(o.q))
|
||||
slices.Sort(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
@@ -755,6 +760,7 @@ func newConnPool(
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
MinIdleConns: minIdleConns,
|
||||
MaxIdleConns: maxIdleConns,
|
||||
MaxActiveConns: maxActiveConns,
|
||||
@@ -801,6 +807,7 @@ func newPubSubPool(
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
MinIdleConns: minIdleConns,
|
||||
MaxIdleConns: maxIdleConns,
|
||||
MaxActiveConns: maxActiveConns,
|
||||
|
||||
+22
-22
@@ -1,6 +1,7 @@
|
||||
package redis
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
"net"
|
||||
"net/url"
|
||||
"runtime"
|
||||
"slices"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -104,6 +106,10 @@ type ClusterOptions struct {
|
||||
// default: 100 milliseconds
|
||||
DialerRetryTimeout time.Duration
|
||||
|
||||
// DialerRetryBackoff controls the delay between dial retry attempts.
|
||||
// See Options.DialerRetryBackoff for details.
|
||||
DialerRetryBackoff func(attempt int) time.Duration
|
||||
|
||||
ReadTimeout time.Duration
|
||||
WriteTimeout time.Duration
|
||||
ContextTimeoutEnabled bool
|
||||
@@ -429,6 +435,7 @@ func (opt *ClusterOptions) clientOptions() *Options {
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
ReadTimeout: opt.ReadTimeout,
|
||||
WriteTimeout: opt.WriteTimeout,
|
||||
|
||||
@@ -786,20 +793,6 @@ type clusterSlot struct {
|
||||
nodes []*clusterNode
|
||||
}
|
||||
|
||||
type clusterSlotSlice []*clusterSlot
|
||||
|
||||
func (p clusterSlotSlice) Len() int {
|
||||
return len(p)
|
||||
}
|
||||
|
||||
func (p clusterSlotSlice) Less(i, j int) bool {
|
||||
return p[i].start < p[j].start
|
||||
}
|
||||
|
||||
func (p clusterSlotSlice) Swap(i, j int) {
|
||||
p[i], p[j] = p[j], p[i]
|
||||
}
|
||||
|
||||
type clusterState struct {
|
||||
nodes *clusterNodes
|
||||
Masters []*clusterNode
|
||||
@@ -858,7 +851,9 @@ func newClusterState(
|
||||
})
|
||||
}
|
||||
|
||||
sort.Sort(clusterSlotSlice(c.slots))
|
||||
slices.SortFunc(c.slots, func(a, b *clusterSlot) int {
|
||||
return cmp.Compare(a.start, b.start)
|
||||
})
|
||||
|
||||
time.AfterFunc(time.Minute, func() {
|
||||
nodes.GC(c.generation)
|
||||
@@ -1139,8 +1134,12 @@ type ClusterClient struct {
|
||||
}
|
||||
|
||||
// NewClusterClient returns a Redis Cluster client as described in
|
||||
// http://redis.io/topics/cluster-spec.
|
||||
// https://redis.io/docs/latest/operate/oss_and_stack/reference/cluster-spec.
|
||||
// Passing nil ClusterOptions will cause a panic.
|
||||
func NewClusterClient(opt *ClusterOptions) *ClusterClient {
|
||||
if opt == nil {
|
||||
panic("redis: NewClusterClient nil options")
|
||||
}
|
||||
opt.init()
|
||||
|
||||
c := &ClusterClient{
|
||||
@@ -1185,7 +1184,8 @@ func NewClusterClient(opt *ClusterOptions) *ClusterClient {
|
||||
return c
|
||||
}
|
||||
|
||||
// Options returns read-only Options that were used to create the client.
|
||||
// Options returns read-only *ClusterOptions that were used to create the client.
|
||||
// Any alteration of the returned *ClusterOptions may result in undefined behaviour.
|
||||
func (c *ClusterClient) Options() *ClusterOptions {
|
||||
return c.opt
|
||||
}
|
||||
@@ -1295,7 +1295,7 @@ func (c *ClusterClient) process(ctx context.Context, cmd Cmder) error {
|
||||
continue
|
||||
}
|
||||
|
||||
if shouldRetry(lastErr, cmd.readTimeout() == nil) {
|
||||
if shouldRetry(lastErr, cmd.readTimeout() == nil) && !cmd.NoRetry() {
|
||||
// First retry the same node.
|
||||
if attempt == 0 {
|
||||
continue
|
||||
@@ -1711,7 +1711,7 @@ func (c *ClusterClient) processPipelineNodeConn(
|
||||
if isBadConn(err, false, node.Client.getAddr()) {
|
||||
node.MarkAsFailing()
|
||||
}
|
||||
if shouldRetry(err, true) {
|
||||
if shouldRetry(err, true) && !cmdsContainNoRetry(cmds) {
|
||||
_ = c.mapCmdsByNode(ctx, failedCmds, cmds)
|
||||
}
|
||||
setCmdsErr(cmds, err)
|
||||
@@ -1747,7 +1747,7 @@ func (c *ClusterClient) pipelineReadCmds(
|
||||
}
|
||||
|
||||
if !isRedisError(err) {
|
||||
if shouldRetry(err, true) {
|
||||
if shouldRetry(err, true) && !cmdsContainNoRetry(cmds) {
|
||||
_ = c.mapCmdsByNode(ctx, failedCmds, cmds)
|
||||
}
|
||||
setCmdsErr(cmds[i+1:], err)
|
||||
@@ -1755,7 +1755,7 @@ func (c *ClusterClient) pipelineReadCmds(
|
||||
}
|
||||
}
|
||||
|
||||
if err := cmds[0].Err(); err != nil && shouldRetry(err, true) {
|
||||
if err := cmds[0].Err(); err != nil && shouldRetry(err, true) && !cmdsContainNoRetry(cmds) {
|
||||
_ = c.mapCmdsByNode(ctx, failedCmds, cmds)
|
||||
return err
|
||||
}
|
||||
@@ -1958,7 +1958,7 @@ func (c *ClusterClient) processTxPipelineNodeConn(
|
||||
if err := cn.WithWriter(c.context(ctx), c.opt.WriteTimeout, func(wr *proto.Writer) error {
|
||||
return writeCmds(wr, cmds)
|
||||
}); err != nil {
|
||||
if shouldRetry(err, true) {
|
||||
if shouldRetry(err, true) && !cmdsContainNoRetry(cmds) {
|
||||
_ = c.mapCmdsByNode(ctx, failedCmds, cmds)
|
||||
}
|
||||
setCmdsErr(cmds, err)
|
||||
|
||||
+31
@@ -79,6 +79,25 @@ type OTelRecorder interface {
|
||||
RecordStreamLag(ctx context.Context, lag time.Duration, cn ConnInfo, streamName, consumerGroup, consumerName string)
|
||||
}
|
||||
|
||||
// OTelConnectionCounter is an optional capability interface for recording
|
||||
// connection count and pending request changes via UpDownCounters.
|
||||
// Implementations of OTelRecorder can optionally implement this interface
|
||||
// to receive connection count and pending request delta notifications.
|
||||
// This is kept separate from OTelRecorder to avoid breaking existing
|
||||
// third-party implementations when new methods are added.
|
||||
type OTelConnectionCounter interface {
|
||||
// RecordConnectionCount records a change in connection count (UpDownCounter)
|
||||
// delta: +1 when connection added, -1 when connection removed
|
||||
// state: connection state (e.g., "idle", "used")
|
||||
// isPubSub: true if this is a PubSub connection
|
||||
RecordConnectionCount(ctx context.Context, delta int, cn ConnInfo, state string, isPubSub bool)
|
||||
|
||||
// RecordPendingRequests records a change in pending requests (UpDownCounter)
|
||||
// delta: +1 when request starts waiting, -1 when request stops waiting
|
||||
// poolName is passed explicitly because we may not have a connection yet when request starts
|
||||
RecordPendingRequests(ctx context.Context, delta int, cn ConnInfo, poolName string)
|
||||
}
|
||||
|
||||
// This is used for async gauge metrics that need to pull stats from pools periodically.
|
||||
type OTelPoolRegistrar interface {
|
||||
// RegisterPool is called when a new client is created with its main connection pool.
|
||||
@@ -163,6 +182,18 @@ func (a *otelRecorderAdapter) RecordStreamLag(ctx context.Context, lag time.Dura
|
||||
a.recorder.RecordStreamLag(ctx, lag, toConnInfo(cn), streamName, consumerGroup, consumerName)
|
||||
}
|
||||
|
||||
func (a *otelRecorderAdapter) RecordConnectionCount(ctx context.Context, delta int, cn *pool.Conn, state string, isPubSub bool) {
|
||||
if counter, ok := a.recorder.(OTelConnectionCounter); ok {
|
||||
counter.RecordConnectionCount(ctx, delta, toConnInfo(cn), state, isPubSub)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *otelRecorderAdapter) RecordPendingRequests(ctx context.Context, delta int, cn *pool.Conn, poolName string) {
|
||||
if counter, ok := a.recorder.(OTelConnectionCounter); ok {
|
||||
counter.RecordPendingRequests(ctx, delta, toConnInfo(cn), poolName)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *otelRecorderAdapter) RegisterPool(poolName string, p pool.Pooler) {
|
||||
if registrar, ok := a.recorder.(OTelPoolRegistrar); ok {
|
||||
registrar.RegisterPool(poolName, &poolerAdapter{p})
|
||||
|
||||
+1
-1
@@ -49,7 +49,7 @@ type Pipeliner interface {
|
||||
var _ Pipeliner = (*Pipeline)(nil)
|
||||
|
||||
// Pipeline implements pipelining as described in
|
||||
// http://redis.io/topics/pipelining.
|
||||
// https://redis.io/docs/latest/develop/using-commands/pipelining.
|
||||
// Please note: it is not safe for concurrent use by multiple goroutines.
|
||||
type Pipeline struct {
|
||||
cmdable
|
||||
|
||||
+32
-27
@@ -3,6 +3,8 @@ package redis
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"maps"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -15,7 +17,7 @@ import (
|
||||
)
|
||||
|
||||
// PubSub implements Pub/Sub commands as described in
|
||||
// http://redis.io/topics/pubsub. Message receiving is NOT safe
|
||||
// https://redis.io/docs/latest/develop/pubsub. Message receiving is NOT safe
|
||||
// for concurrent use by multiple goroutines.
|
||||
//
|
||||
// PubSub automatically reconnects to Redis Server and resubscribes
|
||||
@@ -56,9 +58,9 @@ func (c *PubSub) String() string {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
channels := mapKeys(c.channels)
|
||||
channels = append(channels, mapKeys(c.patterns)...)
|
||||
channels = append(channels, mapKeys(c.schannels)...)
|
||||
channels := slices.Collect(maps.Keys(c.channels))
|
||||
channels = append(channels, slices.Collect(maps.Keys(c.patterns))...)
|
||||
channels = append(channels, slices.Collect(maps.Keys(c.schannels))...)
|
||||
return fmt.Sprintf("PubSub(%s)", strings.Join(channels, ", "))
|
||||
}
|
||||
|
||||
@@ -85,7 +87,7 @@ func (c *PubSub) conn(ctx context.Context, newChannels []string) (*pool.Conn, er
|
||||
c.opt.Addr = internal.RedisNull
|
||||
}
|
||||
|
||||
channels := mapKeys(c.channels)
|
||||
channels := slices.Collect(maps.Keys(c.channels))
|
||||
channels = append(channels, newChannels...)
|
||||
|
||||
cn, err := c.newConn(ctx, c.opt.Addr, channels)
|
||||
@@ -112,18 +114,18 @@ func (c *PubSub) resubscribe(ctx context.Context, cn *pool.Conn) error {
|
||||
var firstErr error
|
||||
|
||||
if len(c.channels) > 0 {
|
||||
firstErr = c._subscribe(ctx, cn, "subscribe", mapKeys(c.channels))
|
||||
firstErr = c._subscribe(ctx, cn, "subscribe", slices.Collect(maps.Keys(c.channels)))
|
||||
}
|
||||
|
||||
if len(c.patterns) > 0 {
|
||||
err := c._subscribe(ctx, cn, "psubscribe", mapKeys(c.patterns))
|
||||
err := c._subscribe(ctx, cn, "psubscribe", slices.Collect(maps.Keys(c.patterns)))
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
|
||||
if len(c.schannels) > 0 {
|
||||
err := c._subscribe(ctx, cn, "ssubscribe", mapKeys(c.schannels))
|
||||
err := c._subscribe(ctx, cn, "ssubscribe", slices.Collect(maps.Keys(c.schannels)))
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
@@ -132,16 +134,6 @@ func (c *PubSub) resubscribe(ctx context.Context, cn *pool.Conn) error {
|
||||
return firstErr
|
||||
}
|
||||
|
||||
func mapKeys(m map[string]struct{}) []string {
|
||||
s := make([]string, len(m))
|
||||
i := 0
|
||||
for k := range m {
|
||||
s[i] = k
|
||||
i++
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (c *PubSub) _subscribe(
|
||||
ctx context.Context, cn *pool.Conn, redisCmd string, channels []string,
|
||||
) error {
|
||||
@@ -284,9 +276,7 @@ func (c *PubSub) Unsubscribe(ctx context.Context, channels ...string) error {
|
||||
}
|
||||
} else {
|
||||
// Unsubscribe from all channels.
|
||||
for channel := range c.channels {
|
||||
delete(c.channels, channel)
|
||||
}
|
||||
clear(c.channels)
|
||||
}
|
||||
|
||||
err := c.subscribe(ctx, "unsubscribe", channels...)
|
||||
@@ -305,9 +295,7 @@ func (c *PubSub) PUnsubscribe(ctx context.Context, patterns ...string) error {
|
||||
}
|
||||
} else {
|
||||
// Unsubscribe from all patterns.
|
||||
for pattern := range c.patterns {
|
||||
delete(c.patterns, pattern)
|
||||
}
|
||||
clear(c.patterns)
|
||||
}
|
||||
|
||||
err := c.subscribe(ctx, "punsubscribe", patterns...)
|
||||
@@ -326,9 +314,7 @@ func (c *PubSub) SUnsubscribe(ctx context.Context, channels ...string) error {
|
||||
}
|
||||
} else {
|
||||
// Unsubscribe from all channels.
|
||||
for channel := range c.schannels {
|
||||
delete(c.schannels, channel)
|
||||
}
|
||||
clear(c.schannels)
|
||||
}
|
||||
|
||||
err := c.subscribe(ctx, "sunsubscribe", channels...)
|
||||
@@ -366,6 +352,25 @@ func (c *PubSub) Ping(ctx context.Context, payload ...string) error {
|
||||
return err
|
||||
}
|
||||
|
||||
// ClientSetName assigns a namee to the PubSub connection using CLIENT SETNAME,
|
||||
// The name is visible in CLIENT LIST output and is useful for debugging
|
||||
// and identifying connections in a redis instance.
|
||||
func (c *PubSub) ClientSetName(ctx context.Context, name string) error {
|
||||
cmd := NewStatusCmd(ctx, "client", "setname", name)
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
cn, err := c.conn(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = c.writeCmd(ctx, cn, cmd)
|
||||
c.releaseConn(ctx, cn, err, false)
|
||||
return err
|
||||
}
|
||||
|
||||
// Subscription received after a successful subscription to channel.
|
||||
type Subscription struct {
|
||||
// Can be "subscribe", "unsubscribe", "psubscribe" or "punsubscribe".
|
||||
|
||||
+176
-33
@@ -215,6 +215,96 @@ func (hs *hooksMixin) processTxPipelineHook(ctx context.Context, cmds []Cmder) e
|
||||
|
||||
//------------------------------------------------------------------------------
|
||||
|
||||
// Stable identifiers for baseClient.onClose hooks. Each component that
|
||||
// registers a close callback owns a dedicated id here so the set of known
|
||||
// hooks is discoverable in one place and id collisions are caught at
|
||||
// compile time. New ids should be added as additional constants.
|
||||
const (
|
||||
// onCloseHookIDSentinelFailover identifies the close callback installed
|
||||
// by NewFailoverClient to tear down sentinel failover background work.
|
||||
onCloseHookIDSentinelFailover = "sentinel-failover"
|
||||
)
|
||||
|
||||
// onCloseHooks is a small registry of named close callbacks attached to a
|
||||
// baseClient. Each callback is identified by a stable string id; registering
|
||||
// the same id twice replaces the previous callback rather than chaining onto
|
||||
// it. This guarantees the registry stays bounded regardless of how often a
|
||||
// hook is (re)registered and avoids the unbounded closure chain that
|
||||
// motivated issue #3772.
|
||||
//
|
||||
// Hooks are invoked in registration order. All hooks run regardless of
|
||||
// individual errors; the first non-nil error is returned.
|
||||
//
|
||||
// A zero-value onCloseHooks is ready to use. It is safe for concurrent use.
|
||||
// Clones of a baseClient share the same *onCloseHooks so registrations and
|
||||
// close semantics are preserved across WithTimeout / WithContext / etc.
|
||||
type onCloseHooks struct {
|
||||
mu sync.Mutex
|
||||
order []string
|
||||
hooks map[string]func() error
|
||||
}
|
||||
|
||||
// register adds or replaces the callback associated with id. Re-registering
|
||||
// an existing id overwrites the previous callback in place; new ids are
|
||||
// appended to the invocation order.
|
||||
func (h *onCloseHooks) register(id string, fn func() error) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if h.hooks == nil {
|
||||
h.hooks = make(map[string]func() error)
|
||||
}
|
||||
if _, exists := h.hooks[id]; !exists {
|
||||
h.order = append(h.order, id)
|
||||
}
|
||||
h.hooks[id] = fn
|
||||
}
|
||||
|
||||
// unregister removes the callback associated with id, if any. It is kept
|
||||
// for API symmetry with register so future callers (e.g. dynamic hook
|
||||
// owners that need to detach before client Close) do not have to
|
||||
// reinvent it.
|
||||
//
|
||||
//nolint:unused // kept for API symmetry with register; see comment above.
|
||||
func (h *onCloseHooks) unregister(id string) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if _, exists := h.hooks[id]; !exists {
|
||||
return
|
||||
}
|
||||
delete(h.hooks, id)
|
||||
for i, x := range h.order {
|
||||
if x == id {
|
||||
h.order = append(h.order[:i], h.order[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// run invokes all registered callbacks in registration order and returns
|
||||
// the first non-nil error encountered. All callbacks are executed even if
|
||||
// an earlier one returns an error.
|
||||
func (h *onCloseHooks) run() error {
|
||||
if h == nil {
|
||||
return nil
|
||||
}
|
||||
h.mu.Lock()
|
||||
fns := make([]func() error, 0, len(h.order))
|
||||
for _, id := range h.order {
|
||||
if fn := h.hooks[id]; fn != nil {
|
||||
fns = append(fns, fn)
|
||||
}
|
||||
}
|
||||
h.mu.Unlock()
|
||||
|
||||
var firstErr error
|
||||
for _, fn := range fns {
|
||||
if err := fn(); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
|
||||
type baseClient struct {
|
||||
opt *Options
|
||||
optLock sync.RWMutex
|
||||
@@ -222,7 +312,13 @@ type baseClient struct {
|
||||
pubSubPool *pool.PubSubPool
|
||||
hooksMixin
|
||||
|
||||
onClose func() error // hook called when client is closed
|
||||
// onClose holds named callbacks invoked when the client is closed.
|
||||
// Registering a new callback never removes previously registered ones;
|
||||
// only re-registering the same id replaces the existing callback. This
|
||||
// lets composing components (e.g. sentinel failover) add close logic
|
||||
// safely without fear of overwriting each other and without building
|
||||
// unbounded closure chains on repeated registration.
|
||||
onClose *onCloseHooks
|
||||
|
||||
// Push notification processing
|
||||
pushProcessor push.NotificationProcessor
|
||||
@@ -252,8 +348,17 @@ func (c *baseClient) clone() *baseClient {
|
||||
return clone
|
||||
}
|
||||
|
||||
// cloneOpt clones c.opt while holding optLock to prevent races with initConn
|
||||
// which writes to MaintNotificationsConfig.Mode under the same lock.
|
||||
func (c *baseClient) cloneOpt() *Options {
|
||||
c.optLock.RLock()
|
||||
clone := c.opt.clone()
|
||||
c.optLock.RUnlock()
|
||||
return clone
|
||||
}
|
||||
|
||||
func (c *baseClient) withTimeout(timeout time.Duration) *baseClient {
|
||||
opt := c.opt.clone()
|
||||
opt := c.cloneOpt()
|
||||
opt.ReadTimeout = timeout
|
||||
opt.WriteTimeout = timeout
|
||||
|
||||
@@ -347,7 +452,11 @@ func (c *baseClient) onAuthenticationErr() func(poolCn *pool.Conn, err error) {
|
||||
if err != nil {
|
||||
if isBadConn(err, false, c.opt.Addr) {
|
||||
// Close the connection to force a reconnection.
|
||||
err := c.connPool.CloseConn(poolCn)
|
||||
// Re-auth happens on connections that were idle in the pool (the pool hook
|
||||
// waits for IDLE state before transitioning to UNUSABLE for re-auth).
|
||||
// From metrics perspective, the connection was never "used" by a client.
|
||||
// Note: Using context.Background() as this callback doesn't have access to caller's context.
|
||||
err := c.connPool.CloseConn(context.Background(), poolCn, pool.CloseReasonAuthError, pool.MetricStateIdle)
|
||||
if err != nil {
|
||||
internal.Logger.Printf(context.Background(), "redis: failed to close connection: %v", err)
|
||||
// try to close the network connection directly
|
||||
@@ -363,27 +472,6 @@ func (c *baseClient) onAuthenticationErr() func(poolCn *pool.Conn, err error) {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *baseClient) wrappedOnClose(newOnClose func() error) func() error {
|
||||
onClose := c.onClose
|
||||
return func() error {
|
||||
var firstErr error
|
||||
err := newOnClose()
|
||||
// Even if we have an error we would like to execute the onClose hook
|
||||
// if it exists. We will return the first error that occurred.
|
||||
// This is to keep error handling consistent with the rest of the code.
|
||||
if err != nil {
|
||||
firstErr = err
|
||||
}
|
||||
if onClose != nil {
|
||||
err = onClose()
|
||||
if err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
}
|
||||
|
||||
func (c *baseClient) initConn(ctx context.Context, cn *pool.Conn) error {
|
||||
// This function is called in two scenarios:
|
||||
// 1. First-time init: Connection is in CREATED state (from pool.Get())
|
||||
@@ -483,7 +571,22 @@ func (c *baseClient) initConn(ctx context.Context, cn *pool.Conn) error {
|
||||
return fmt.Errorf("failed to subscribe to streaming credentials: %w", initErr)
|
||||
}
|
||||
|
||||
c.onClose = c.wrappedOnClose(unsubscribeFromCredentialsProvider)
|
||||
// Per-connection unsubscribe is attached to the connection itself so it
|
||||
// runs when this specific connection is closed. Do not register it on
|
||||
// c.onClose: initConn runs for every (re)initialized connection, and
|
||||
// attaching per-connection state to the shared baseClient registry would
|
||||
// either leak entries (one per connection id, never trimmed) or — with
|
||||
// the pre-fix wrappedOnClose approach — build an unbounded closure chain
|
||||
// retaining every prior connection's unsubscribe (see issue #3772).
|
||||
//
|
||||
// Note: pool.Conn.SetOnClose OVERWRITES any prior callback (see the
|
||||
// doc on that method). That is safe here because the streaming
|
||||
// credentials Manager deduplicates listeners by connection id, so a
|
||||
// second initConn on the same cn re-Subscribes the SAME listener and
|
||||
// the returned unsubscribe is equivalent to the one already installed.
|
||||
// Any future code path that could hand out a distinct unsubscribe on
|
||||
// re-initialization must first invoke the existing one to avoid
|
||||
// orphaning the old subscription on the credentials provider.
|
||||
cn.SetOnClose(unsubscribeFromCredentialsProvider)
|
||||
|
||||
username, password = credentials.BasicAuth()
|
||||
@@ -501,8 +604,13 @@ func (c *baseClient) initConn(ctx context.Context, cn *pool.Conn) error {
|
||||
|
||||
// for redis-server versions that do not support the HELLO command,
|
||||
// RESP2 will continue to be used.
|
||||
// helloOK tracks whether HELLO succeeded. If it did not, the connection
|
||||
// falls back to RESP2 regardless of c.opt.Protocol, and features that
|
||||
// require RESP3 (e.g. maintenance notifications) must be skipped.
|
||||
helloOK := false
|
||||
if initErr = conn.Hello(ctx, c.opt.Protocol, username, password, c.opt.ClientName).Err(); initErr == nil {
|
||||
// Authentication successful with HELLO command
|
||||
helloOK = true
|
||||
} else if !isRedisError(initErr) {
|
||||
// When the server responds with the RESP protocol and the result is not a normal
|
||||
// execution result of the HELLO command, we consider it to be an indication that
|
||||
@@ -551,10 +659,38 @@ func (c *baseClient) initConn(ctx context.Context, cn *pool.Conn) error {
|
||||
maintNotifEnabled := c.opt.MaintNotificationsConfig != nil && c.opt.MaintNotificationsConfig.Mode != maintnotifications.ModeDisabled
|
||||
protocol := c.opt.Protocol
|
||||
var endpointType maintnotifications.EndpointType
|
||||
var maintNotifMode maintnotifications.Mode
|
||||
if maintNotifEnabled {
|
||||
endpointType = c.opt.MaintNotificationsConfig.EndpointType
|
||||
maintNotifMode = c.opt.MaintNotificationsConfig.Mode
|
||||
}
|
||||
c.optLock.RUnlock()
|
||||
|
||||
// Maintenance notifications require RESP3 push frames. If HELLO failed
|
||||
// and the connection fell back to RESP2, there is no point in sending
|
||||
// CLIENT MAINT_NOTIFICATIONS: the server either rejects it (making the
|
||||
// error misleading) or accepts it silently, leaving the client unable
|
||||
// to receive any notifications. Decide based on the actual negotiated
|
||||
// protocol rather than the requested one.
|
||||
if maintNotifEnabled && protocol == 3 && !helloOK {
|
||||
if maintNotifMode == maintnotifications.ModeEnabled {
|
||||
// Explicitly requested - fail fast with a clear reason.
|
||||
cn.GetStateMachine().Transition(pool.StateClosed)
|
||||
if errorCallback := pool.GetMetricErrorCallback(); errorCallback != nil {
|
||||
errorCallback(ctx, "HANDSHAKE_FAILED", cn, "HANDSHAKE_FAILED", true, 0)
|
||||
}
|
||||
return fmt.Errorf("failed to enable maintnotifications: server does not support RESP3 (HELLO command failed)")
|
||||
}
|
||||
// auto/other modes: silently disable maintnotifications for this client.
|
||||
c.optLock.Lock()
|
||||
c.opt.MaintNotificationsConfig.Mode = maintnotifications.ModeDisabled
|
||||
c.optLock.Unlock()
|
||||
if err := c.disableMaintNotificationsUpgrades(); err != nil {
|
||||
internal.Logger.Printf(ctx, "failed to disable maintnotifications in auto mode: %v", err)
|
||||
}
|
||||
maintNotifEnabled = false
|
||||
}
|
||||
|
||||
var maintNotifHandshakeErr error
|
||||
if maintNotifEnabled && protocol == 3 {
|
||||
maintNotifHandshakeErr = conn.ClientMaintNotifications(
|
||||
@@ -703,7 +839,9 @@ func (c *baseClient) process(ctx context.Context, cmd Cmder) error {
|
||||
if cn != nil {
|
||||
lastConn = cn
|
||||
}
|
||||
if err == nil || !retry {
|
||||
// Don't retry if command explicitly disables retries (e.g., RawWriteToCmd
|
||||
// which writes directly to an io.Writer and cannot undo partial writes)
|
||||
if err == nil || !retry || cmd.NoRetry() {
|
||||
// Record total operation duration
|
||||
if opDurationCallback != nil {
|
||||
operationDuration := time.Since(operationStart)
|
||||
@@ -948,10 +1086,8 @@ func (c *baseClient) Close() error {
|
||||
firstErr = err
|
||||
}
|
||||
|
||||
if c.onClose != nil {
|
||||
if err := c.onClose(); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
if err := c.onClose.run(); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
|
||||
// Unregister pools from OTel before closing them
|
||||
@@ -1028,7 +1164,10 @@ func (c *baseClient) generalProcessPipeline(
|
||||
canRetry, err = p(ctx, cn, cmds)
|
||||
return err
|
||||
})
|
||||
if lastErr == nil || !canRetry || !shouldRetry(lastErr, true) {
|
||||
// Don't retry if any command in the pipeline explicitly disables retries
|
||||
// (e.g., RawWriteToCmd which writes directly to an io.Writer and cannot
|
||||
// undo partial writes on retry)
|
||||
if lastErr == nil || !canRetry || !shouldRetry(lastErr, true) || cmdsContainNoRetry(cmds) {
|
||||
// The error should be set here only when failing to obtain the conn.
|
||||
if !isRedisError(lastErr) {
|
||||
setCmdsErr(cmds, lastErr)
|
||||
@@ -1196,6 +1335,7 @@ type Client struct {
|
||||
}
|
||||
|
||||
// NewClient returns a client to the Redis Server specified by Options.
|
||||
// Passing nil Options will cause a panic.
|
||||
func NewClient(opt *Options) *Client {
|
||||
if opt == nil {
|
||||
panic("redis: NewClient nil options")
|
||||
@@ -1208,7 +1348,8 @@ func NewClient(opt *Options) *Client {
|
||||
|
||||
c := Client{
|
||||
baseClient: &baseClient{
|
||||
opt: opt,
|
||||
opt: opt,
|
||||
onClose: &onCloseHooks{},
|
||||
},
|
||||
}
|
||||
c.init()
|
||||
@@ -1295,7 +1436,8 @@ func (c *Client) Process(ctx context.Context, cmd Cmder) error {
|
||||
return err
|
||||
}
|
||||
|
||||
// Options returns read-only Options that were used to create the client.
|
||||
// Options returns read-only *Options that were used to create the client.
|
||||
// Any alteration of the returned *Options may result in undefined behaviour.
|
||||
func (c *Client) Options() *Options {
|
||||
return c.opt
|
||||
}
|
||||
@@ -1490,6 +1632,7 @@ func newConn(opt *Options, connPool pool.Pooler, parentHooks *hooksMixin) *Conn
|
||||
baseClient: baseClient{
|
||||
opt: opt,
|
||||
connPool: connPool,
|
||||
onClose: &onCloseHooks{},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
+11
-13
@@ -11,8 +11,6 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/cespare/xxhash/v2"
|
||||
"github.com/dgryski/go-rendezvous" //nolint
|
||||
"github.com/redis/go-redis/v9/auth"
|
||||
|
||||
"github.com/redis/go-redis/v9/internal"
|
||||
@@ -36,16 +34,8 @@ type ConsistentHash interface {
|
||||
Get(string) string
|
||||
}
|
||||
|
||||
type rendezvousWrapper struct {
|
||||
*rendezvous.Rendezvous
|
||||
}
|
||||
|
||||
func (w rendezvousWrapper) Get(key string) string {
|
||||
return w.Lookup(key)
|
||||
}
|
||||
|
||||
func newRendezvous(shards []string) ConsistentHash {
|
||||
return rendezvousWrapper{rendezvous.New(shards, xxhash.Sum64String)}
|
||||
return hashtag.NewRendezvousHash(shards)
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------------
|
||||
@@ -120,6 +110,10 @@ type RingOptions struct {
|
||||
// default: 100 milliseconds
|
||||
DialerRetryTimeout time.Duration
|
||||
|
||||
// DialerRetryBackoff controls the delay between dial retry attempts.
|
||||
// See Options.DialerRetryBackoff for details.
|
||||
DialerRetryBackoff func(attempt int) time.Duration
|
||||
|
||||
ReadTimeout time.Duration
|
||||
WriteTimeout time.Duration
|
||||
ContextTimeoutEnabled bool
|
||||
@@ -233,6 +227,7 @@ func (opt *RingOptions) clientOptions() *Options {
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
ReadTimeout: opt.ReadTimeout,
|
||||
WriteTimeout: opt.WriteTimeout,
|
||||
ContextTimeoutEnabled: opt.ContextTimeoutEnabled,
|
||||
@@ -600,6 +595,8 @@ type Ring struct {
|
||||
heartbeatCancelFn context.CancelFunc
|
||||
}
|
||||
|
||||
// NewRing returns a Redis Ring client to the Redis Server specified by RingOptions.
|
||||
// Passing nil RingOptions will cause a panic.
|
||||
func NewRing(opt *RingOptions) *Ring {
|
||||
if opt == nil {
|
||||
panic("redis: NewRing nil options")
|
||||
@@ -642,7 +639,8 @@ func (c *Ring) Process(ctx context.Context, cmd Cmder) error {
|
||||
return err
|
||||
}
|
||||
|
||||
// Options returns read-only Options that were used to create the client.
|
||||
// Options returns read-only *RingOptions that were used to create the client.
|
||||
// Any alteration of the returned *RingOptions may result in undefined behaviour.
|
||||
func (c *Ring) Options() *RingOptions {
|
||||
return c.opt
|
||||
}
|
||||
@@ -797,7 +795,7 @@ func (c *Ring) process(ctx context.Context, cmd Cmder) error {
|
||||
}
|
||||
|
||||
lastErr = shard.Client.Process(ctx, cmd)
|
||||
if lastErr == nil || !shouldRetry(lastErr, cmd.readTimeout() == nil) {
|
||||
if lastErr == nil || !shouldRetry(lastErr, cmd.readTimeout() == nil) || cmd.NoRetry() {
|
||||
return lastErr
|
||||
}
|
||||
}
|
||||
|
||||
+134
-9
@@ -4,7 +4,9 @@ import (
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"io"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type Scripter interface {
|
||||
@@ -23,28 +25,69 @@ var (
|
||||
)
|
||||
|
||||
type Script struct {
|
||||
src, hash string
|
||||
src string
|
||||
mu sync.RWMutex
|
||||
hash string
|
||||
serverSHA bool // if true: do not compute SHA-1 in Go; load digest from Redis (SCRIPT LOAD)
|
||||
}
|
||||
|
||||
func NewScript(src string) *Script {
|
||||
h := sha1.New()
|
||||
_, _ = io.WriteString(h, src)
|
||||
|
||||
return &Script{
|
||||
src: src,
|
||||
hash: hex.EncodeToString(h.Sum(nil)),
|
||||
src: src,
|
||||
hash: hex.EncodeToString(h.Sum(nil)),
|
||||
serverSHA: false,
|
||||
}
|
||||
}
|
||||
|
||||
// NewScriptServerSHA creates a Script that avoids computing SHA-1 in Go.
|
||||
// The digest is obtained from Redis via SCRIPT LOAD (server-side hashing),
|
||||
// then EVALSHA/EVALSHA_RO is used.
|
||||
func NewScriptServerSHA(src string) *Script {
|
||||
return &Script{
|
||||
src: src,
|
||||
serverSHA: true,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Script) Hash() string {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.hash
|
||||
}
|
||||
|
||||
func (s *Script) Load(ctx context.Context, c Scripter) *StringCmd {
|
||||
return c.ScriptLoad(ctx, s.src)
|
||||
cmd := c.ScriptLoad(ctx, s.src)
|
||||
if err := cmd.Err(); err == nil {
|
||||
s.mu.Lock()
|
||||
s.hash = cmd.Val()
|
||||
s.mu.Unlock()
|
||||
}
|
||||
return cmd
|
||||
}
|
||||
|
||||
func (s *Script) Exists(ctx context.Context, c Scripter) *BoolSliceCmd {
|
||||
return c.ScriptExists(ctx, s.hash)
|
||||
s.mu.RLock()
|
||||
hash := s.hash
|
||||
serverSHA := s.serverSHA
|
||||
s.mu.RUnlock()
|
||||
if hash == "" && serverSHA {
|
||||
// For server-side scripts, obtain digest from Redis first.
|
||||
// If hash is empty, it means SCRIPT LOAD was not called yet, so we check existence of empty hash which will return false.
|
||||
// This avoids unnecessary SCRIPT LOAD just to check existence.
|
||||
if err := s.ensureHash(ctx, c); err != nil {
|
||||
return c.ScriptExists(ctx, "")
|
||||
}
|
||||
s.mu.RLock()
|
||||
hash = s.hash
|
||||
s.mu.RUnlock()
|
||||
}
|
||||
if hash == "" {
|
||||
return c.ScriptExists(ctx, "")
|
||||
}
|
||||
return c.ScriptExists(ctx, hash)
|
||||
}
|
||||
|
||||
func (s *Script) Eval(ctx context.Context, c Scripter, keys []string, args ...interface{}) *Cmd {
|
||||
@@ -55,19 +98,101 @@ func (s *Script) EvalRO(ctx context.Context, c Scripter, keys []string, args ...
|
||||
return c.EvalRO(ctx, s.src, keys, args...)
|
||||
}
|
||||
|
||||
// ensureHash ensures that s.hash is populated by using SCRIPT LOAD.
|
||||
// It never calls SHA-1 in Go; Redis computes and returns the digest.
|
||||
func (s *Script) ensureHash(ctx context.Context, c Scripter) error {
|
||||
// Fast path: read lock, return if hash is already set.
|
||||
s.mu.RLock()
|
||||
if s.hash != "" {
|
||||
s.mu.RUnlock()
|
||||
return nil
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
|
||||
// Slow path: acquire write lock and load.
|
||||
s.mu.Lock()
|
||||
if s.hash != "" {
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
cmd := c.ScriptLoad(ctx, s.src)
|
||||
if err := cmd.Err(); err != nil {
|
||||
s.mu.Unlock()
|
||||
return err
|
||||
}
|
||||
s.hash = cmd.Val()
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Script) EvalSha(ctx context.Context, c Scripter, keys []string, args ...interface{}) *Cmd {
|
||||
return c.EvalSha(ctx, s.hash, keys, args...)
|
||||
// Default behavior: use client-side SHA-1 computed in NewScript.
|
||||
if !s.serverSHA {
|
||||
s.mu.RLock()
|
||||
hash := s.hash
|
||||
s.mu.RUnlock()
|
||||
return c.EvalSha(ctx, hash, keys, args...)
|
||||
}
|
||||
|
||||
// Server-side SHA via SCRIPT LOAD + EVALSHA.
|
||||
if err := s.ensureHash(ctx, c); err != nil {
|
||||
return s.Eval(ctx, c, keys, args...)
|
||||
}
|
||||
|
||||
s.mu.RLock()
|
||||
hash := s.hash
|
||||
s.mu.RUnlock()
|
||||
|
||||
r := c.EvalSha(ctx, hash, keys, args...)
|
||||
if HasErrorPrefix(r.Err(), "NOSCRIPT") {
|
||||
// Script cache was flushed; reload and retry once.
|
||||
if err := s.ensureHash(ctx, c); err != nil {
|
||||
return s.Eval(ctx, c, keys, args...)
|
||||
}
|
||||
s.mu.RLock()
|
||||
hash = s.hash
|
||||
s.mu.RUnlock()
|
||||
return c.EvalSha(ctx, hash, keys, args...)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
func (s *Script) EvalShaRO(ctx context.Context, c Scripter, keys []string, args ...interface{}) *Cmd {
|
||||
return c.EvalShaRO(ctx, s.hash, keys, args...)
|
||||
if !s.serverSHA {
|
||||
s.mu.RLock()
|
||||
hash := s.hash
|
||||
s.mu.RUnlock()
|
||||
return c.EvalShaRO(ctx, hash, keys, args...)
|
||||
}
|
||||
|
||||
if err := s.ensureHash(ctx, c); err != nil {
|
||||
return s.EvalRO(ctx, c, keys, args...)
|
||||
}
|
||||
|
||||
s.mu.RLock()
|
||||
hash := s.hash
|
||||
s.mu.RUnlock()
|
||||
|
||||
r := c.EvalShaRO(ctx, hash, keys, args...)
|
||||
if HasErrorPrefix(r.Err(), "NOSCRIPT") {
|
||||
if err := s.ensureHash(ctx, c); err != nil {
|
||||
return s.EvalRO(ctx, c, keys, args...)
|
||||
}
|
||||
s.mu.RLock()
|
||||
hash = s.hash
|
||||
s.mu.RUnlock()
|
||||
return c.EvalShaRO(ctx, hash, keys, args...)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// Run optimistically uses EVALSHA to run the script. If script does not exist
|
||||
// it is retried using EVAL.
|
||||
func (s *Script) Run(ctx context.Context, c Scripter, keys []string, args ...interface{}) *Cmd {
|
||||
r := s.EvalSha(ctx, c, keys, args...)
|
||||
if HasErrorPrefix(r.Err(), "NOSCRIPT") {
|
||||
if errors.Is(r.Err(), ErrNoScript) {
|
||||
return s.Eval(ctx, c, keys, args...)
|
||||
}
|
||||
return r
|
||||
@@ -77,7 +202,7 @@ func (s *Script) Run(ctx context.Context, c Scripter, keys []string, args ...int
|
||||
// it is retried using EVAL_RO.
|
||||
func (s *Script) RunRO(ctx context.Context, c Scripter, keys []string, args ...interface{}) *Cmd {
|
||||
r := s.EvalShaRO(ctx, c, keys, args...)
|
||||
if HasErrorPrefix(r.Err(), "NOSCRIPT") {
|
||||
if errors.Is(r.Err(), ErrNoScript) {
|
||||
return s.EvalRO(ctx, c, keys, args...)
|
||||
}
|
||||
return r
|
||||
|
||||
+5
@@ -60,6 +60,11 @@ func (c cmdable) eval(ctx context.Context, name, payload string, keys []string,
|
||||
cmd.SetFirstKeyPos(3)
|
||||
}
|
||||
_ = c(ctx, cmd)
|
||||
if err := cmd.Err(); err != nil {
|
||||
if HasErrorPrefix(err, "NOSCRIPT") {
|
||||
cmd.SetErr(ErrNoScript)
|
||||
}
|
||||
}
|
||||
return cmd
|
||||
}
|
||||
|
||||
|
||||
+65
-32
@@ -2,6 +2,7 @@ package redis
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// ----------------------
|
||||
@@ -215,6 +216,7 @@ type AggregateBuilder struct {
|
||||
index string
|
||||
query string
|
||||
options *FTAggregateOptions
|
||||
err error
|
||||
}
|
||||
|
||||
// NewAggregateBuilder creates a new AggregateBuilder for FT.AGGREGATE commands.
|
||||
@@ -223,6 +225,14 @@ func (c *Client) NewAggregateBuilder(ctx context.Context, index, query string) *
|
||||
return &AggregateBuilder{c: c, ctx: ctx, index: index, query: query, options: &FTAggregateOptions{LimitOffset: -1}}
|
||||
}
|
||||
|
||||
// setErr records the first error produced while building the pipeline.
|
||||
// Subsequent errors are ignored; the first error is returned from Run.
|
||||
func (b *AggregateBuilder) setErr(err error) {
|
||||
if b.err == nil {
|
||||
b.err = err
|
||||
}
|
||||
}
|
||||
|
||||
// Verbatim includes VERBATIM.
|
||||
func (b *AggregateBuilder) Verbatim() *AggregateBuilder { b.options.Verbatim = true; return b }
|
||||
|
||||
@@ -241,15 +251,15 @@ func (b *AggregateBuilder) LoadAll() *AggregateBuilder {
|
||||
return b
|
||||
}
|
||||
|
||||
// Load adds LOAD <n> <field> [AS alias]...
|
||||
// You can call it multiple times for multiple fields.
|
||||
// Load adds a LOAD <field> [AS alias] step.
|
||||
// You can call it multiple times; each call becomes a separate LOAD clause
|
||||
// at its position in the pipeline.
|
||||
func (b *AggregateBuilder) Load(field string, alias ...string) *AggregateBuilder {
|
||||
// each Load entry becomes one element in options.Load
|
||||
l := FTAggregateLoad{Field: field}
|
||||
l := &FTAggregateLoad{Field: field}
|
||||
if len(alias) > 0 {
|
||||
l.As = alias[0]
|
||||
}
|
||||
b.options.Load = append(b.options.Load, l)
|
||||
b.options.Steps = append(b.options.Steps, FTAggregateStep{Load: l})
|
||||
return b
|
||||
}
|
||||
|
||||
@@ -259,62 +269,79 @@ func (b *AggregateBuilder) Timeout(ms int) *AggregateBuilder {
|
||||
return b
|
||||
}
|
||||
|
||||
// Apply adds APPLY <field> [AS alias].
|
||||
// Apply adds an APPLY <field> [AS alias] step.
|
||||
func (b *AggregateBuilder) Apply(field string, alias ...string) *AggregateBuilder {
|
||||
a := FTAggregateApply{Field: field}
|
||||
a := &FTAggregateApply{Field: field}
|
||||
if len(alias) > 0 {
|
||||
a.As = alias[0]
|
||||
}
|
||||
b.options.Apply = append(b.options.Apply, a)
|
||||
b.options.Steps = append(b.options.Steps, FTAggregateStep{Apply: a})
|
||||
return b
|
||||
}
|
||||
|
||||
// GroupBy starts a new GROUPBY <fields...> clause.
|
||||
// GroupBy adds a new GROUPBY <fields...> step.
|
||||
func (b *AggregateBuilder) GroupBy(fields ...interface{}) *AggregateBuilder {
|
||||
b.options.GroupBy = append(b.options.GroupBy, FTAggregateGroupBy{
|
||||
Fields: fields,
|
||||
b.options.Steps = append(b.options.Steps, FTAggregateStep{
|
||||
GroupBy: &FTAggregateGroupBy{Fields: fields},
|
||||
})
|
||||
return b
|
||||
}
|
||||
|
||||
// Reduce adds a REDUCE <fn> [<#args> <args...>] clause to the *last* GROUPBY.
|
||||
// Reduce adds a REDUCE <fn> [<#args> <args...>] clause to the last step,
|
||||
// which must be a GROUPBY. If it is not, Run will return an error.
|
||||
func (b *AggregateBuilder) Reduce(fn SearchAggregator, args ...interface{}) *AggregateBuilder {
|
||||
if len(b.options.GroupBy) == 0 {
|
||||
// no GROUPBY yet — nothing to attach to
|
||||
n := len(b.options.Steps)
|
||||
if n == 0 || b.options.Steps[n-1].GroupBy == nil {
|
||||
b.setErr(fmt.Errorf("FT.AGGREGATE: Reduce must follow a GroupBy step"))
|
||||
return b
|
||||
}
|
||||
idx := len(b.options.GroupBy) - 1
|
||||
b.options.GroupBy[idx].Reduce = append(b.options.GroupBy[idx].Reduce, FTAggregateReducer{
|
||||
Reducer: fn,
|
||||
Args: args,
|
||||
})
|
||||
g := b.options.Steps[n-1].GroupBy
|
||||
g.Reduce = append(g.Reduce, FTAggregateReducer{Reducer: fn, Args: args})
|
||||
return b
|
||||
}
|
||||
|
||||
// ReduceAs does the same but also sets an alias: REDUCE <fn> … AS <alias>
|
||||
// ReduceAs does the same but also sets an alias: REDUCE <fn> … AS <alias>.
|
||||
// The last step must be a GROUPBY; otherwise Run will return an error.
|
||||
func (b *AggregateBuilder) ReduceAs(fn SearchAggregator, alias string, args ...interface{}) *AggregateBuilder {
|
||||
if len(b.options.GroupBy) == 0 {
|
||||
n := len(b.options.Steps)
|
||||
if n == 0 || b.options.Steps[n-1].GroupBy == nil {
|
||||
b.setErr(fmt.Errorf("FT.AGGREGATE: ReduceAs must follow a GroupBy step"))
|
||||
return b
|
||||
}
|
||||
idx := len(b.options.GroupBy) - 1
|
||||
b.options.GroupBy[idx].Reduce = append(b.options.GroupBy[idx].Reduce, FTAggregateReducer{
|
||||
Reducer: fn,
|
||||
Args: args,
|
||||
As: alias,
|
||||
})
|
||||
g := b.options.Steps[n-1].GroupBy
|
||||
g.Reduce = append(g.Reduce, FTAggregateReducer{Reducer: fn, Args: args, As: alias})
|
||||
return b
|
||||
}
|
||||
|
||||
// SortBy adds SORTBY <field> ASC|DESC.
|
||||
// SortBy adds SORTBY <field> ASC|DESC. Consecutive SortBy calls (with no
|
||||
// other step in between) are merged into a single SORTBY clause so fields
|
||||
// act as tiebreakers. A SortBy call after a non-SortBy step starts a new
|
||||
// SORTBY step.
|
||||
//
|
||||
// Note: this is a semantics change from earlier experimental versions of
|
||||
// the builder, where SortBy always accumulated into a single SORTBY clause
|
||||
// regardless of position in the pipeline.
|
||||
func (b *AggregateBuilder) SortBy(field string, asc bool) *AggregateBuilder {
|
||||
sb := FTAggregateSortBy{FieldName: field, Asc: asc, Desc: !asc}
|
||||
b.options.SortBy = append(b.options.SortBy, sb)
|
||||
if n := len(b.options.Steps); n > 0 && b.options.Steps[n-1].SortBy != nil {
|
||||
b.options.Steps[n-1].SortBy.Fields = append(b.options.Steps[n-1].SortBy.Fields, sb)
|
||||
return b
|
||||
}
|
||||
b.options.Steps = append(b.options.Steps, FTAggregateStep{
|
||||
SortBy: &FTAggregateSortByStep{Fields: []FTAggregateSortBy{sb}},
|
||||
})
|
||||
return b
|
||||
}
|
||||
|
||||
// SortByMax sets MAX <n> (only if SortBy was called).
|
||||
// SortByMax sets MAX <n> on the last SORTBY step. The last step must be a
|
||||
// SORTBY; otherwise Run will return an error.
|
||||
func (b *AggregateBuilder) SortByMax(max int) *AggregateBuilder {
|
||||
b.options.SortByMax = max
|
||||
n := len(b.options.Steps)
|
||||
if n == 0 || b.options.Steps[n-1].SortBy == nil {
|
||||
b.setErr(fmt.Errorf("FT.AGGREGATE: SortByMax must follow a SortBy step"))
|
||||
return b
|
||||
}
|
||||
b.options.Steps[n-1].SortBy.Max = max
|
||||
return b
|
||||
}
|
||||
|
||||
@@ -352,8 +379,14 @@ func (b *AggregateBuilder) Dialect(version int) *AggregateBuilder {
|
||||
return b
|
||||
}
|
||||
|
||||
// Run executes FT.AGGREGATE and returns a typed result.
|
||||
// Run executes FT.AGGREGATE and returns a typed result. If the builder
|
||||
// recorded a validation error while constructing the pipeline (for example,
|
||||
// calling SortByMax when the last step is not a SortBy), that error is
|
||||
// returned without issuing the command.
|
||||
func (b *AggregateBuilder) Run() (*FTAggregateResult, error) {
|
||||
if b.err != nil {
|
||||
return nil, b.err
|
||||
}
|
||||
cmd := b.c.FTAggregateWithArgs(b.ctx, b.index, b.query, b.options)
|
||||
return cmd.Result()
|
||||
}
|
||||
|
||||
+270
-133
@@ -3,6 +3,8 @@ package redis
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"maps"
|
||||
"slices"
|
||||
"strconv"
|
||||
|
||||
"github.com/redis/go-redis/v9/internal"
|
||||
@@ -256,22 +258,42 @@ type FTAggregateWithCursor struct {
|
||||
MaxIdle int
|
||||
}
|
||||
|
||||
// FTAggregateSortByStep represents a SORTBY operation with optional MAX.
|
||||
// Used inside FTAggregateStep to place SORTBY at an arbitrary position in
|
||||
// the aggregation pipeline.
|
||||
type FTAggregateSortByStep struct {
|
||||
Fields []FTAggregateSortBy
|
||||
Max int // 0 means no MAX
|
||||
}
|
||||
|
||||
// FTAggregateStep represents a single operation in the aggregation pipeline.
|
||||
// LOAD, APPLY, SORTBY and GROUPBY can all appear multiple times in any order.
|
||||
// Exactly one of the fields should be set per step.
|
||||
type FTAggregateStep struct {
|
||||
Load *FTAggregateLoad
|
||||
Apply *FTAggregateApply
|
||||
GroupBy *FTAggregateGroupBy
|
||||
SortBy *FTAggregateSortByStep
|
||||
}
|
||||
|
||||
type FTAggregateOptions struct {
|
||||
Verbatim bool
|
||||
LoadAll bool
|
||||
Load []FTAggregateLoad
|
||||
Timeout int
|
||||
GroupBy []FTAggregateGroupBy
|
||||
SortBy []FTAggregateSortBy
|
||||
SortByMax int
|
||||
Verbatim bool
|
||||
LoadAll bool
|
||||
Timeout int
|
||||
// Scorer is used to set scoring function, if not set passed, a default will be used.
|
||||
// The default scorer depends on the Redis version:
|
||||
// - `BM25` for Redis >= 8
|
||||
// - `TFIDF` for Redis < 8
|
||||
Scorer string
|
||||
// AddScores is available in Redis CE 8
|
||||
AddScores bool
|
||||
Apply []FTAggregateApply
|
||||
AddScores bool
|
||||
|
||||
// Steps is the ordered sequence of aggregation pipeline operations.
|
||||
// It can contain LOAD, APPLY, GROUPBY and SORTBY in any order, multiple times.
|
||||
// Steps cannot be combined with the deprecated Load, Apply, GroupBy, SortBy
|
||||
// and SortByMax fields: doing so returns an error.
|
||||
Steps []FTAggregateStep
|
||||
|
||||
LimitOffset int
|
||||
Limit int
|
||||
Filter string
|
||||
@@ -280,6 +302,17 @@ type FTAggregateOptions struct {
|
||||
Params map[string]interface{}
|
||||
// Dialect 1,3 and 4 are deprecated since redis 8.0
|
||||
DialectVersion int
|
||||
|
||||
// Deprecated: Use Steps instead.
|
||||
Load []FTAggregateLoad
|
||||
// Deprecated: Use Steps instead.
|
||||
GroupBy []FTAggregateGroupBy
|
||||
// Deprecated: Use Steps instead.
|
||||
SortBy []FTAggregateSortBy
|
||||
// Deprecated: Use Steps instead.
|
||||
SortByMax int
|
||||
// Deprecated: Use Steps instead.
|
||||
Apply []FTAggregateApply
|
||||
}
|
||||
|
||||
type FTSearchFilter struct {
|
||||
@@ -615,9 +648,112 @@ func (c cmdable) FTAggregate(ctx context.Context, index string, query string) *M
|
||||
return cmd
|
||||
}
|
||||
|
||||
// validateFTAggregateOptions validates mutually exclusive combinations of
|
||||
// FTAggregateOptions fields before any command arguments are constructed.
|
||||
func validateFTAggregateOptions(options *FTAggregateOptions) error {
|
||||
if len(options.Steps) > 0 {
|
||||
if options.Load != nil || options.Apply != nil || options.GroupBy != nil ||
|
||||
options.SortBy != nil || options.SortByMax != 0 {
|
||||
return fmt.Errorf("FT.AGGREGATE: Steps cannot be combined with the deprecated Load, Apply, GroupBy, SortBy and SortByMax fields")
|
||||
}
|
||||
if options.LoadAll {
|
||||
for _, step := range options.Steps {
|
||||
if step.Load != nil {
|
||||
return fmt.Errorf("FT.AGGREGATE: LOADALL and LOAD are mutually exclusive")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if options.LoadAll && options.Load != nil {
|
||||
return fmt.Errorf("FT.AGGREGATE: LOADALL and LOAD are mutually exclusive")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// appendFTAggregateStep appends the Redis command arguments for a single
|
||||
// aggregation pipeline step. Each step must set exactly one of Load, Apply,
|
||||
// GroupBy or SortBy.
|
||||
func appendFTAggregateStep(args []interface{}, step FTAggregateStep) ([]interface{}, error) {
|
||||
set := 0
|
||||
if step.Load != nil {
|
||||
set++
|
||||
}
|
||||
if step.Apply != nil {
|
||||
set++
|
||||
}
|
||||
if step.GroupBy != nil {
|
||||
set++
|
||||
}
|
||||
if step.SortBy != nil {
|
||||
set++
|
||||
}
|
||||
if set != 1 {
|
||||
return args, fmt.Errorf("FT.AGGREGATE: each step must set exactly one of Load, Apply, GroupBy, SortBy (got %d)", set)
|
||||
}
|
||||
|
||||
switch {
|
||||
case step.Load != nil:
|
||||
args = append(args, "LOAD")
|
||||
countIdx := len(args)
|
||||
args = append(args, 0)
|
||||
count := 0
|
||||
args = append(args, step.Load.Field)
|
||||
count++
|
||||
if step.Load.As != "" {
|
||||
args = append(args, "AS", step.Load.As)
|
||||
count += 2
|
||||
}
|
||||
args[countIdx] = count
|
||||
case step.Apply != nil:
|
||||
args = append(args, "APPLY", step.Apply.Field)
|
||||
if step.Apply.As != "" {
|
||||
args = append(args, "AS", step.Apply.As)
|
||||
}
|
||||
case step.GroupBy != nil:
|
||||
args = append(args, "GROUPBY", len(step.GroupBy.Fields))
|
||||
args = append(args, step.GroupBy.Fields...)
|
||||
for _, reducer := range step.GroupBy.Reduce {
|
||||
args = append(args, "REDUCE", reducer.Reducer.String())
|
||||
if reducer.Args != nil {
|
||||
args = append(args, len(reducer.Args))
|
||||
args = append(args, reducer.Args...)
|
||||
} else {
|
||||
args = append(args, 0)
|
||||
}
|
||||
if reducer.As != "" {
|
||||
args = append(args, "AS", reducer.As)
|
||||
}
|
||||
}
|
||||
case step.SortBy != nil:
|
||||
args = append(args, "SORTBY")
|
||||
sortByOptions := []interface{}{}
|
||||
for _, sortBy := range step.SortBy.Fields {
|
||||
if sortBy.Asc && sortBy.Desc {
|
||||
return args, fmt.Errorf("FT.AGGREGATE: ASC and DESC are mutually exclusive")
|
||||
}
|
||||
sortByOptions = append(sortByOptions, sortBy.FieldName)
|
||||
if sortBy.Asc {
|
||||
sortByOptions = append(sortByOptions, "ASC")
|
||||
}
|
||||
if sortBy.Desc {
|
||||
sortByOptions = append(sortByOptions, "DESC")
|
||||
}
|
||||
}
|
||||
args = append(args, len(sortByOptions))
|
||||
args = append(args, sortByOptions...)
|
||||
if step.SortBy.Max > 0 {
|
||||
args = append(args, "MAX", step.SortBy.Max)
|
||||
}
|
||||
}
|
||||
return args, nil
|
||||
}
|
||||
|
||||
func FTAggregateQuery(query string, options *FTAggregateOptions) (AggregateQuery, error) {
|
||||
queryArgs := []interface{}{query}
|
||||
if options != nil {
|
||||
if err := validateFTAggregateOptions(options); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if options.Verbatim {
|
||||
queryArgs = append(queryArgs, "VERBATIM")
|
||||
}
|
||||
@@ -630,13 +766,10 @@ func FTAggregateQuery(query string, options *FTAggregateOptions) (AggregateQuery
|
||||
queryArgs = append(queryArgs, "ADDSCORES")
|
||||
}
|
||||
|
||||
if options.LoadAll && options.Load != nil {
|
||||
return nil, fmt.Errorf("FT.AGGREGATE: LOADALL and LOAD are mutually exclusive")
|
||||
}
|
||||
if options.LoadAll {
|
||||
queryArgs = append(queryArgs, "LOAD", "*")
|
||||
}
|
||||
if options.Load != nil {
|
||||
if len(options.Steps) == 0 && options.Load != nil {
|
||||
queryArgs = append(queryArgs, "LOAD", len(options.Load))
|
||||
index, count := len(queryArgs)-1, 0
|
||||
for _, load := range options.Load {
|
||||
@@ -654,53 +787,63 @@ func FTAggregateQuery(query string, options *FTAggregateOptions) (AggregateQuery
|
||||
queryArgs = append(queryArgs, "TIMEOUT", options.Timeout)
|
||||
}
|
||||
|
||||
for _, apply := range options.Apply {
|
||||
queryArgs = append(queryArgs, "APPLY", apply.Field)
|
||||
if apply.As != "" {
|
||||
queryArgs = append(queryArgs, "AS", apply.As)
|
||||
if len(options.Steps) > 0 {
|
||||
for _, step := range options.Steps {
|
||||
var err error
|
||||
queryArgs, err = appendFTAggregateStep(queryArgs, step)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for _, apply := range options.Apply {
|
||||
queryArgs = append(queryArgs, "APPLY", apply.Field)
|
||||
if apply.As != "" {
|
||||
queryArgs = append(queryArgs, "AS", apply.As)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if options.GroupBy != nil {
|
||||
for _, groupBy := range options.GroupBy {
|
||||
queryArgs = append(queryArgs, "GROUPBY", len(groupBy.Fields))
|
||||
queryArgs = append(queryArgs, groupBy.Fields...)
|
||||
if options.GroupBy != nil {
|
||||
for _, groupBy := range options.GroupBy {
|
||||
queryArgs = append(queryArgs, "GROUPBY", len(groupBy.Fields))
|
||||
queryArgs = append(queryArgs, groupBy.Fields...)
|
||||
|
||||
for _, reducer := range groupBy.Reduce {
|
||||
queryArgs = append(queryArgs, "REDUCE")
|
||||
queryArgs = append(queryArgs, reducer.Reducer.String())
|
||||
if reducer.Args != nil {
|
||||
queryArgs = append(queryArgs, len(reducer.Args))
|
||||
queryArgs = append(queryArgs, reducer.Args...)
|
||||
} else {
|
||||
queryArgs = append(queryArgs, 0)
|
||||
}
|
||||
if reducer.As != "" {
|
||||
queryArgs = append(queryArgs, "AS", reducer.As)
|
||||
for _, reducer := range groupBy.Reduce {
|
||||
queryArgs = append(queryArgs, "REDUCE")
|
||||
queryArgs = append(queryArgs, reducer.Reducer.String())
|
||||
if reducer.Args != nil {
|
||||
queryArgs = append(queryArgs, len(reducer.Args))
|
||||
queryArgs = append(queryArgs, reducer.Args...)
|
||||
} else {
|
||||
queryArgs = append(queryArgs, 0)
|
||||
}
|
||||
if reducer.As != "" {
|
||||
queryArgs = append(queryArgs, "AS", reducer.As)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if options.SortBy != nil {
|
||||
queryArgs = append(queryArgs, "SORTBY")
|
||||
sortByOptions := []interface{}{}
|
||||
for _, sortBy := range options.SortBy {
|
||||
sortByOptions = append(sortByOptions, sortBy.FieldName)
|
||||
if sortBy.Asc && sortBy.Desc {
|
||||
return nil, fmt.Errorf("FT.AGGREGATE: ASC and DESC are mutually exclusive")
|
||||
}
|
||||
if sortBy.Asc {
|
||||
sortByOptions = append(sortByOptions, "ASC")
|
||||
}
|
||||
if sortBy.Desc {
|
||||
sortByOptions = append(sortByOptions, "DESC")
|
||||
if options.SortBy != nil {
|
||||
queryArgs = append(queryArgs, "SORTBY")
|
||||
sortByOptions := []interface{}{}
|
||||
for _, sortBy := range options.SortBy {
|
||||
sortByOptions = append(sortByOptions, sortBy.FieldName)
|
||||
if sortBy.Asc && sortBy.Desc {
|
||||
return nil, fmt.Errorf("FT.AGGREGATE: ASC and DESC are mutually exclusive")
|
||||
}
|
||||
if sortBy.Asc {
|
||||
sortByOptions = append(sortByOptions, "ASC")
|
||||
}
|
||||
if sortBy.Desc {
|
||||
sortByOptions = append(sortByOptions, "DESC")
|
||||
}
|
||||
}
|
||||
queryArgs = append(queryArgs, len(sortByOptions))
|
||||
queryArgs = append(queryArgs, sortByOptions...)
|
||||
}
|
||||
if options.SortByMax > 0 {
|
||||
queryArgs = append(queryArgs, "MAX", options.SortByMax)
|
||||
}
|
||||
queryArgs = append(queryArgs, len(sortByOptions))
|
||||
queryArgs = append(queryArgs, sortByOptions...)
|
||||
}
|
||||
if options.SortByMax > 0 {
|
||||
queryArgs = append(queryArgs, "MAX", options.SortByMax)
|
||||
}
|
||||
if options.LimitOffset >= 0 && options.Limit > 0 {
|
||||
queryArgs = append(queryArgs, "LIMIT", options.LimitOffset, options.Limit)
|
||||
@@ -850,6 +993,11 @@ func (cmd *AggregateCmd) Clone() Cmder {
|
||||
func (c cmdable) FTAggregateWithArgs(ctx context.Context, index string, query string, options *FTAggregateOptions) *AggregateCmd {
|
||||
args := []interface{}{"FT.AGGREGATE", index, query}
|
||||
if options != nil {
|
||||
if err := validateFTAggregateOptions(options); err != nil {
|
||||
cmd := NewAggregateCmd(ctx, args...)
|
||||
cmd.SetErr(err)
|
||||
return cmd
|
||||
}
|
||||
if options.Verbatim {
|
||||
args = append(args, "VERBATIM")
|
||||
}
|
||||
@@ -859,15 +1007,10 @@ func (c cmdable) FTAggregateWithArgs(ctx context.Context, index string, query st
|
||||
if options.AddScores {
|
||||
args = append(args, "ADDSCORES")
|
||||
}
|
||||
if options.LoadAll && options.Load != nil {
|
||||
cmd := NewAggregateCmd(ctx, args...)
|
||||
cmd.SetErr(fmt.Errorf("FT.AGGREGATE: LOADALL and LOAD are mutually exclusive"))
|
||||
return cmd
|
||||
}
|
||||
if options.LoadAll {
|
||||
args = append(args, "LOAD", "*")
|
||||
}
|
||||
if options.Load != nil {
|
||||
if len(options.Steps) == 0 && options.Load != nil {
|
||||
args = append(args, "LOAD", len(options.Load))
|
||||
index, count := len(args)-1, 0
|
||||
for _, load := range options.Load {
|
||||
@@ -883,54 +1026,66 @@ func (c cmdable) FTAggregateWithArgs(ctx context.Context, index string, query st
|
||||
if options.Timeout > 0 {
|
||||
args = append(args, "TIMEOUT", options.Timeout)
|
||||
}
|
||||
for _, apply := range options.Apply {
|
||||
args = append(args, "APPLY", apply.Field)
|
||||
if apply.As != "" {
|
||||
args = append(args, "AS", apply.As)
|
||||
}
|
||||
}
|
||||
if options.GroupBy != nil {
|
||||
for _, groupBy := range options.GroupBy {
|
||||
args = append(args, "GROUPBY", len(groupBy.Fields))
|
||||
args = append(args, groupBy.Fields...)
|
||||
|
||||
for _, reducer := range groupBy.Reduce {
|
||||
args = append(args, "REDUCE")
|
||||
args = append(args, reducer.Reducer.String())
|
||||
if reducer.Args != nil {
|
||||
args = append(args, len(reducer.Args))
|
||||
args = append(args, reducer.Args...)
|
||||
} else {
|
||||
args = append(args, 0)
|
||||
}
|
||||
if reducer.As != "" {
|
||||
args = append(args, "AS", reducer.As)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if options.SortBy != nil {
|
||||
args = append(args, "SORTBY")
|
||||
sortByOptions := []interface{}{}
|
||||
for _, sortBy := range options.SortBy {
|
||||
sortByOptions = append(sortByOptions, sortBy.FieldName)
|
||||
if sortBy.Asc && sortBy.Desc {
|
||||
if len(options.Steps) > 0 {
|
||||
for _, step := range options.Steps {
|
||||
var err error
|
||||
args, err = appendFTAggregateStep(args, step)
|
||||
if err != nil {
|
||||
cmd := NewAggregateCmd(ctx, args...)
|
||||
cmd.SetErr(fmt.Errorf("FT.AGGREGATE: ASC and DESC are mutually exclusive"))
|
||||
cmd.SetErr(err)
|
||||
return cmd
|
||||
}
|
||||
if sortBy.Asc {
|
||||
sortByOptions = append(sortByOptions, "ASC")
|
||||
}
|
||||
if sortBy.Desc {
|
||||
sortByOptions = append(sortByOptions, "DESC")
|
||||
}
|
||||
} else {
|
||||
for _, apply := range options.Apply {
|
||||
args = append(args, "APPLY", apply.Field)
|
||||
if apply.As != "" {
|
||||
args = append(args, "AS", apply.As)
|
||||
}
|
||||
}
|
||||
args = append(args, len(sortByOptions))
|
||||
args = append(args, sortByOptions...)
|
||||
}
|
||||
if options.SortByMax > 0 {
|
||||
args = append(args, "MAX", options.SortByMax)
|
||||
if options.GroupBy != nil {
|
||||
for _, groupBy := range options.GroupBy {
|
||||
args = append(args, "GROUPBY", len(groupBy.Fields))
|
||||
args = append(args, groupBy.Fields...)
|
||||
|
||||
for _, reducer := range groupBy.Reduce {
|
||||
args = append(args, "REDUCE")
|
||||
args = append(args, reducer.Reducer.String())
|
||||
if reducer.Args != nil {
|
||||
args = append(args, len(reducer.Args))
|
||||
args = append(args, reducer.Args...)
|
||||
} else {
|
||||
args = append(args, 0)
|
||||
}
|
||||
if reducer.As != "" {
|
||||
args = append(args, "AS", reducer.As)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if options.SortBy != nil {
|
||||
args = append(args, "SORTBY")
|
||||
sortByOptions := []interface{}{}
|
||||
for _, sortBy := range options.SortBy {
|
||||
sortByOptions = append(sortByOptions, sortBy.FieldName)
|
||||
if sortBy.Asc && sortBy.Desc {
|
||||
cmd := NewAggregateCmd(ctx, args...)
|
||||
cmd.SetErr(fmt.Errorf("FT.AGGREGATE: ASC and DESC are mutually exclusive"))
|
||||
return cmd
|
||||
}
|
||||
if sortBy.Asc {
|
||||
sortByOptions = append(sortByOptions, "ASC")
|
||||
}
|
||||
if sortBy.Desc {
|
||||
sortByOptions = append(sortByOptions, "DESC")
|
||||
}
|
||||
}
|
||||
args = append(args, len(sortByOptions))
|
||||
args = append(args, sortByOptions...)
|
||||
}
|
||||
if options.SortByMax > 0 {
|
||||
args = append(args, "MAX", options.SortByMax)
|
||||
}
|
||||
}
|
||||
if options.LimitOffset >= 0 && options.Limit > 0 {
|
||||
args = append(args, "LIMIT", options.LimitOffset, options.Limit)
|
||||
@@ -1728,26 +1883,19 @@ func (cmd *FTInfoCmd) Clone() Cmder {
|
||||
}
|
||||
// Clone slices and maps
|
||||
if cmd.val.Attributes != nil {
|
||||
val.Attributes = make([]FTAttribute, len(cmd.val.Attributes))
|
||||
copy(val.Attributes, cmd.val.Attributes)
|
||||
val.Attributes = slices.Clone(cmd.val.Attributes)
|
||||
}
|
||||
if cmd.val.DialectStats != nil {
|
||||
val.DialectStats = make(map[string]int, len(cmd.val.DialectStats))
|
||||
for k, v := range cmd.val.DialectStats {
|
||||
val.DialectStats[k] = v
|
||||
}
|
||||
val.DialectStats = maps.Clone(cmd.val.DialectStats)
|
||||
}
|
||||
if cmd.val.FieldStatistics != nil {
|
||||
val.FieldStatistics = make([]FieldStatistic, len(cmd.val.FieldStatistics))
|
||||
copy(val.FieldStatistics, cmd.val.FieldStatistics)
|
||||
val.FieldStatistics = slices.Clone(cmd.val.FieldStatistics)
|
||||
}
|
||||
if cmd.val.IndexOptions != nil {
|
||||
val.IndexOptions = make([]string, len(cmd.val.IndexOptions))
|
||||
copy(val.IndexOptions, cmd.val.IndexOptions)
|
||||
val.IndexOptions = slices.Clone(cmd.val.IndexOptions)
|
||||
}
|
||||
if cmd.val.IndexDefinition.Prefixes != nil {
|
||||
val.IndexDefinition.Prefixes = make([]string, len(cmd.val.IndexDefinition.Prefixes))
|
||||
copy(val.IndexDefinition.Prefixes, cmd.val.IndexDefinition.Prefixes)
|
||||
val.IndexDefinition.Prefixes = slices.Clone(cmd.val.IndexDefinition.Prefixes)
|
||||
}
|
||||
return &FTInfoCmd{
|
||||
baseCmd: cmd.cloneBaseCmd(),
|
||||
@@ -1918,8 +2066,7 @@ func (cmd *FTSpellCheckCmd) Clone() Cmder {
|
||||
Term: result.Term,
|
||||
}
|
||||
if result.Suggestions != nil {
|
||||
val[i].Suggestions = make([]SpellCheckSuggestion, len(result.Suggestions))
|
||||
copy(val[i].Suggestions, result.Suggestions)
|
||||
val[i].Suggestions = slices.Clone(result.Suggestions)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -2115,34 +2262,25 @@ func (cmd *FTSearchCmd) Clone() Cmder {
|
||||
}
|
||||
// Clone slices and maps
|
||||
if cmd.options.Filters != nil {
|
||||
options.Filters = make([]FTSearchFilter, len(cmd.options.Filters))
|
||||
copy(options.Filters, cmd.options.Filters)
|
||||
options.Filters = slices.Clone(cmd.options.Filters)
|
||||
}
|
||||
if cmd.options.GeoFilter != nil {
|
||||
options.GeoFilter = make([]FTSearchGeoFilter, len(cmd.options.GeoFilter))
|
||||
copy(options.GeoFilter, cmd.options.GeoFilter)
|
||||
options.GeoFilter = slices.Clone(cmd.options.GeoFilter)
|
||||
}
|
||||
if cmd.options.InKeys != nil {
|
||||
options.InKeys = make([]interface{}, len(cmd.options.InKeys))
|
||||
copy(options.InKeys, cmd.options.InKeys)
|
||||
options.InKeys = slices.Clone(cmd.options.InKeys)
|
||||
}
|
||||
if cmd.options.InFields != nil {
|
||||
options.InFields = make([]interface{}, len(cmd.options.InFields))
|
||||
copy(options.InFields, cmd.options.InFields)
|
||||
options.InFields = slices.Clone(cmd.options.InFields)
|
||||
}
|
||||
if cmd.options.Return != nil {
|
||||
options.Return = make([]FTSearchReturn, len(cmd.options.Return))
|
||||
copy(options.Return, cmd.options.Return)
|
||||
options.Return = slices.Clone(cmd.options.Return)
|
||||
}
|
||||
if cmd.options.SortBy != nil {
|
||||
options.SortBy = make([]FTSearchSortBy, len(cmd.options.SortBy))
|
||||
copy(options.SortBy, cmd.options.SortBy)
|
||||
options.SortBy = slices.Clone(cmd.options.SortBy)
|
||||
}
|
||||
if cmd.options.Params != nil {
|
||||
options.Params = make(map[string]interface{}, len(cmd.options.Params))
|
||||
for k, v := range cmd.options.Params {
|
||||
options.Params[k] = v
|
||||
}
|
||||
options.Params = maps.Clone(cmd.options.Params)
|
||||
}
|
||||
}
|
||||
return &FTSearchCmd{
|
||||
@@ -2368,8 +2506,7 @@ func (cmd *FTHybridCmd) Clone() Cmder {
|
||||
}
|
||||
}
|
||||
if cmd.val.Warnings != nil {
|
||||
val.Warnings = make([]string, len(cmd.val.Warnings))
|
||||
copy(val.Warnings, cmd.val.Warnings)
|
||||
val.Warnings = slices.Clone(cmd.val.Warnings)
|
||||
}
|
||||
|
||||
var cursorVal *FTHybridCursorResult
|
||||
|
||||
+18
-13
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -100,6 +101,10 @@ type FailoverOptions struct {
|
||||
// default: 100 milliseconds
|
||||
DialerRetryTimeout time.Duration
|
||||
|
||||
// DialerRetryBackoff controls the delay between dial retry attempts.
|
||||
// See Options.DialerRetryBackoff for details.
|
||||
DialerRetryBackoff func(attempt int) time.Duration
|
||||
|
||||
ReadTimeout time.Duration
|
||||
WriteTimeout time.Duration
|
||||
ContextTimeoutEnabled bool
|
||||
@@ -197,6 +202,7 @@ func (opt *FailoverOptions) clientOptions() *Options {
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
ReadTimeout: opt.ReadTimeout,
|
||||
WriteTimeout: opt.WriteTimeout,
|
||||
|
||||
@@ -251,6 +257,7 @@ func (opt *FailoverOptions) sentinelOptions(addr string) *Options {
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
ReadTimeout: opt.ReadTimeout,
|
||||
WriteTimeout: opt.WriteTimeout,
|
||||
|
||||
@@ -311,6 +318,7 @@ func (opt *FailoverOptions) clusterOptions() *ClusterOptions {
|
||||
DialTimeout: opt.DialTimeout,
|
||||
DialerRetries: opt.DialerRetries,
|
||||
DialerRetryTimeout: opt.DialerRetryTimeout,
|
||||
DialerRetryBackoff: opt.DialerRetryBackoff,
|
||||
ReadTimeout: opt.ReadTimeout,
|
||||
WriteTimeout: opt.WriteTimeout,
|
||||
|
||||
@@ -494,6 +502,7 @@ func setupFailoverConnParams(u *url.URL, o *FailoverOptions) (*FailoverOptions,
|
||||
// NewFailoverClient returns a Redis client that uses Redis Sentinel
|
||||
// for automatic failover. It's safe for concurrent use by multiple
|
||||
// goroutines.
|
||||
// Passing nil FailoverOptions will cause a panic.
|
||||
func NewFailoverClient(failoverOpt *FailoverOptions) *Client {
|
||||
if failoverOpt == nil {
|
||||
panic("redis: NewFailoverClient nil options")
|
||||
@@ -524,7 +533,8 @@ func NewFailoverClient(failoverOpt *FailoverOptions) *Client {
|
||||
|
||||
rdb := &Client{
|
||||
baseClient: &baseClient{
|
||||
opt: opt,
|
||||
opt: opt,
|
||||
onClose: &onCloseHooks{},
|
||||
},
|
||||
}
|
||||
rdb.init()
|
||||
@@ -548,7 +558,7 @@ func NewFailoverClient(failoverOpt *FailoverOptions) *Client {
|
||||
panic(fmt.Errorf("redis: failed to create pubsub pool: %w", err))
|
||||
}
|
||||
|
||||
rdb.onClose = rdb.wrappedOnClose(failover.Close)
|
||||
rdb.onClose.register(onCloseHookIDSentinelFailover, failover.Close)
|
||||
|
||||
failover.mu.Lock()
|
||||
failover.onFailover = func(ctx context.Context, addr string) {
|
||||
@@ -603,6 +613,8 @@ type SentinelClient struct {
|
||||
*baseClient
|
||||
}
|
||||
|
||||
// NewSentinelClient returns a Redis Sentinel client.
|
||||
// Passing nil Options will cause a panic.
|
||||
func NewSentinelClient(opt *Options) *SentinelClient {
|
||||
if opt == nil {
|
||||
panic("redis: NewSentinelClient nil options")
|
||||
@@ -610,7 +622,8 @@ func NewSentinelClient(opt *Options) *SentinelClient {
|
||||
opt.init()
|
||||
c := &SentinelClient{
|
||||
baseClient: &baseClient{
|
||||
opt: opt,
|
||||
opt: opt,
|
||||
onClose: &onCloseHooks{},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1128,7 +1141,7 @@ func (c *sentinelFailover) discoverSentinels(ctx context.Context) {
|
||||
}
|
||||
if ip != "" && port != "" {
|
||||
sentinelAddr := net.JoinHostPort(ip, port)
|
||||
if !contains(c.sentinelAddrs, sentinelAddr) {
|
||||
if !slices.Contains(c.sentinelAddrs, sentinelAddr) {
|
||||
internal.Logger.Printf(ctx, "sentinel: discovered new sentinel=%q for master=%q",
|
||||
sentinelAddr, c.opt.MasterName)
|
||||
c.sentinelAddrs = append(c.sentinelAddrs, sentinelAddr)
|
||||
@@ -1162,19 +1175,11 @@ func (c *sentinelFailover) listen(pubsub *PubSub) {
|
||||
}
|
||||
}
|
||||
|
||||
func contains(slice []string, str string) bool {
|
||||
for _, s := range slice {
|
||||
if s == str {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------------
|
||||
|
||||
// NewFailoverClusterClient returns a client that supports routing read-only commands
|
||||
// to a replica node.
|
||||
// Passing nil FailoverOptions will cause a panic.
|
||||
func NewFailoverClusterClient(failoverOpt *FailoverOptions) *ClusterClient {
|
||||
if failoverOpt == nil {
|
||||
panic("redis: NewFailoverClusterClient nil options")
|
||||
|
||||
+23
-6
@@ -373,6 +373,17 @@ type ZRangeArgs struct {
|
||||
// }
|
||||
// cmd: "ZRange example-key (3 8 ByScore" (3 < score <= 8).
|
||||
//
|
||||
// When the Rev option is also provided, <Start> should be the higher score value and
|
||||
// <Stop> should be the lower score value (i.e. reversed order):
|
||||
// ZRangeArgs{
|
||||
// Key: "example-key",
|
||||
// Start: 8,
|
||||
// Stop: "(3",
|
||||
// ByScore: true,
|
||||
// Rev: true,
|
||||
// }
|
||||
// cmd: "ZRange example-key 8 (3 ByScore Rev" (8 >= score > 3, in reverse order).
|
||||
//
|
||||
// For the ByLex option, it is similar to the deprecated(6.2.0+) ZRangeByLex command.
|
||||
// You can set the <Start> and <Stop> options as follows:
|
||||
// ZRangeArgs{
|
||||
@@ -383,6 +394,17 @@ type ZRangeArgs struct {
|
||||
// }
|
||||
// cmd: "ZRange example-key [abc (def ByLex"
|
||||
//
|
||||
// When the Rev option is also provided, <Start> should be the lexicographically higher
|
||||
// value and <Stop> should be the lower value:
|
||||
// ZRangeArgs{
|
||||
// Key: "example-key",
|
||||
// Start: "(def",
|
||||
// Stop: "[abc",
|
||||
// ByLex: true,
|
||||
// Rev: true,
|
||||
// }
|
||||
// cmd: "ZRange example-key (def [abc ByLex Rev"
|
||||
//
|
||||
// For normal cases (ByScore==false && ByLex==false), <Start> and <Stop> should be set to the index range (int).
|
||||
// You can read the documentation for more information: https://redis.io/commands/zrange
|
||||
Start interface{}
|
||||
@@ -400,12 +422,7 @@ type ZRangeArgs struct {
|
||||
}
|
||||
|
||||
func (z ZRangeArgs) appendArgs(args []interface{}) []interface{} {
|
||||
// For Rev+ByScore/ByLex, we need to adjust the position of <Start> and <Stop>.
|
||||
if z.Rev && (z.ByScore || z.ByLex) {
|
||||
args = append(args, z.Key, z.Stop, z.Start)
|
||||
} else {
|
||||
args = append(args, z.Key, z.Start, z.Stop)
|
||||
}
|
||||
args = append(args, z.Key, z.Start, z.Stop)
|
||||
|
||||
if z.ByScore {
|
||||
args = append(args, "byscore")
|
||||
|
||||
+1
-4
@@ -429,8 +429,6 @@ func (c cmdable) SetEx(ctx context.Context, key string, value interface{}, expir
|
||||
|
||||
// SetNX sets the value of a key only if the key does not exist.
|
||||
//
|
||||
// Deprecated: Use Set with NX option instead as of Redis 2.6.12.
|
||||
//
|
||||
// Zero expiration means the key has no expiration time.
|
||||
// KeepTTL is a Redis KEEPTTL option to keep existing TTL, it requires your redis-server version >= 6.0,
|
||||
// otherwise you will receive an error: (error) ERR syntax error.
|
||||
@@ -438,8 +436,7 @@ func (c cmdable) SetNX(ctx context.Context, key string, value interface{}, expir
|
||||
var cmd *BoolCmd
|
||||
switch expiration {
|
||||
case 0:
|
||||
// Use old `SETNX` to support old Redis versions.
|
||||
cmd = NewBoolCmd(ctx, "setnx", key, value)
|
||||
cmd = NewBoolCmd(ctx, "set", key, value, "nx")
|
||||
case KeepTTL:
|
||||
cmd = NewBoolCmd(ctx, "set", key, value, "keepttl", "nx")
|
||||
default:
|
||||
|
||||
+3
-2
@@ -11,7 +11,7 @@ import (
|
||||
const TxFailedErr = proto.RedisError("redis: transaction failed")
|
||||
|
||||
// Tx implements Redis transactions as described in
|
||||
// http://redis.io/topics/transactions. It's NOT safe for concurrent use
|
||||
// https://redis.io/docs/latest/develop/using-commands/transactions. It's NOT safe for concurrent use
|
||||
// by multiple goroutines, because Exec resets list of watched keys.
|
||||
//
|
||||
// If you don't need WATCH, use Pipeline instead.
|
||||
@@ -24,10 +24,11 @@ type Tx struct {
|
||||
func (c *Client) newTx() *Tx {
|
||||
tx := Tx{
|
||||
baseClient: baseClient{
|
||||
opt: c.opt.clone(), // Clone options to avoid sharing mutable state between transaction and parent client
|
||||
opt: c.cloneOpt(), // Clone options under optLock to avoid race with initConn
|
||||
connPool: pool.NewStickyConnPool(c.connPool),
|
||||
hooksMixin: c.hooksMixin.clone(),
|
||||
pushProcessor: c.pushProcessor, // Copy push processor from parent client
|
||||
onClose: &onCloseHooks{},
|
||||
},
|
||||
}
|
||||
tx.init()
|
||||
|
||||
+2
@@ -372,6 +372,8 @@ var (
|
||||
// 3. If the number of Addrs is two or more, or IsClusterMode option is specified,
|
||||
// a ClusterClient is returned.
|
||||
// 4. Otherwise, a single-node Client is returned.
|
||||
//
|
||||
// Passing nil UniversalOptions will cause a panic.
|
||||
func NewUniversalClient(opts *UniversalOptions) UniversalClient {
|
||||
if opts == nil {
|
||||
panic("redis: NewUniversalClient nil options")
|
||||
|
||||
+58
-1
@@ -26,7 +26,10 @@ type VectorSetCmdable interface {
|
||||
VSimWithScores(ctx context.Context, key string, val Vector) *VectorScoreSliceCmd
|
||||
VSimWithArgs(ctx context.Context, key string, val Vector, args *VSimArgs) *StringSliceCmd
|
||||
VSimWithArgsWithScores(ctx context.Context, key string, val Vector, args *VSimArgs) *VectorScoreSliceCmd
|
||||
VSimWithArgsWithAttribs(ctx context.Context, key string, val Vector, args *VSimArgs) *VectorAttribSliceCmd
|
||||
VSimWithArgsWithScoresWithAttribs(ctx context.Context, key string, val Vector, args *VSimArgs) *VectorScoreAttribSliceCmd
|
||||
VRange(ctx context.Context, key, start, end string, count int64) *StringSliceCmd
|
||||
VIsMember(ctx context.Context, key, element string) *BoolCmd
|
||||
}
|
||||
|
||||
type Vector interface {
|
||||
@@ -79,6 +82,17 @@ type VectorScore struct {
|
||||
Score float64
|
||||
}
|
||||
|
||||
type VectorAttrib struct {
|
||||
Name string
|
||||
Attribs *string
|
||||
}
|
||||
|
||||
type VectorScoreAttrib struct {
|
||||
Name string
|
||||
Score float64
|
||||
Attribs *string
|
||||
}
|
||||
|
||||
// `VADD key (FP32 | VALUES num) vector element`
|
||||
// note: the API is experimental and may be subject to change.
|
||||
func (c cmdable) VAdd(ctx context.Context, key, element string, val Vector) *BoolCmd {
|
||||
@@ -311,7 +325,7 @@ func (v VSimArgs) appendArgs(args []any) []any {
|
||||
args = append(args, "nothread")
|
||||
}
|
||||
if v.Epsilon > 0 {
|
||||
args = append(args, "Epsilon", v.Epsilon)
|
||||
args = append(args, "epsilon", v.Epsilon)
|
||||
}
|
||||
return args
|
||||
}
|
||||
@@ -347,6 +361,40 @@ func (c cmdable) VSimWithArgsWithScores(ctx context.Context, key string, val Vec
|
||||
return cmd
|
||||
}
|
||||
|
||||
// `VSIM key (ELE | FP32 | VALUES num) (vector | element) [WITHATTRIBS] [COUNT num] [EPSILON delta]
|
||||
// [EF search-exploration-factor] [FILTER expression] [FILTER-EF max-filtering-effort] [TRUTH] [NOTHREAD]`
|
||||
// WITHATTRIBS is only available in Redis v8.2.0+
|
||||
// note: the API is experimental and may be subject to change.
|
||||
func (c cmdable) VSimWithArgsWithAttribs(ctx context.Context, key string, val Vector, simArgs *VSimArgs) *VectorAttribSliceCmd {
|
||||
if simArgs == nil {
|
||||
simArgs = &VSimArgs{}
|
||||
}
|
||||
args := []any{"vsim", key}
|
||||
args = append(args, val.Value()...)
|
||||
args = append(args, "withattribs")
|
||||
args = simArgs.appendArgs(args)
|
||||
cmd := NewVectorAttribSliceCmd(ctx, args...)
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// `VSIM key (ELE | FP32 | VALUES num) (vector | element) [WITHSCORES] [WITHATTRIBS] [COUNT num] [EPSILON delta]
|
||||
// [EF search-exploration-factor] [FILTER expression] [FILTER-EF max-filtering-effort] [TRUTH] [NOTHREAD]`
|
||||
// WITHATTRIBS is only available in Redis v8.2.0+
|
||||
// note: the API is experimental and may be subject to change.
|
||||
func (c cmdable) VSimWithArgsWithScoresWithAttribs(ctx context.Context, key string, val Vector, simArgs *VSimArgs) *VectorScoreAttribSliceCmd {
|
||||
if simArgs == nil {
|
||||
simArgs = &VSimArgs{}
|
||||
}
|
||||
args := []any{"vsim", key}
|
||||
args = append(args, val.Value()...)
|
||||
args = append(args, "withscores", "withattribs")
|
||||
args = simArgs.appendArgs(args)
|
||||
cmd := NewVectorScoreAttribSliceCmd(ctx, args...)
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// `VRANGE key start end count`
|
||||
// a negative count means to return all the elements in the vector set.
|
||||
// note: the API is experimental and may be subject to change.
|
||||
@@ -356,3 +404,12 @@ func (c cmdable) VRange(ctx context.Context, key, start, end string, count int64
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// `VISMEMBER key element`
|
||||
// Check if an element exists in a vector set.
|
||||
// note: the API is experimental and may be subject to change.
|
||||
func (c cmdable) VIsMember(ctx context.Context, key, element string) *BoolCmd {
|
||||
cmd := NewBoolCmd(ctx, "vismember", key, element)
|
||||
_ = c(ctx, cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
+1
-1
@@ -2,5 +2,5 @@ package redis
|
||||
|
||||
// Version is the current release version.
|
||||
func Version() string {
|
||||
return "9.18.0"
|
||||
return "9.19.0"
|
||||
}
|
||||
|
||||
+81
-1
@@ -7,6 +7,83 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [0.23.1] - 2026-05-10
|
||||
|
||||
### Fixed
|
||||
|
||||
- `NewSecretsVerifier` now rejects empty signing secrets to avoid accepting forged request
|
||||
signatures when applications are misconfigured.
|
||||
|
||||
## [0.23.0] - 2026-04-22
|
||||
|
||||
### Added
|
||||
|
||||
- **Block Kit: `CardBlock` and `CarouselBlock`** — Support for two of the new
|
||||
agent-UI blocks announced in the
|
||||
[April 16 Slack changelog](https://docs.slack.dev/changelog/2026/04/16/block-kit-new-blocks).
|
||||
`CardBlock` is constructed via `NewCardBlock` with a functional-options
|
||||
pattern and fluent `With*` builders (`WithTitle`, `WithSubtitle`, `WithBody`,
|
||||
`WithIcon`, `WithHeroImage`, `WithActions`). `CarouselBlock` is constructed
|
||||
via `NewCarouselBlock` with a variadic `*CardBlock` list plus `WithBlockID`
|
||||
and `AddCard` helpers. Both blocks wire into `Blocks.UnmarshalJSON` for
|
||||
round-trip fidelity, and reuse existing `ImageBlockElement` /
|
||||
`ButtonBlockElement` / `BlockElements` types rather than introducing new
|
||||
composition objects.
|
||||
- **Block Kit: `AlertBlock`** — Support for the third of the new agent-UI
|
||||
blocks from the
|
||||
[April 16 Slack changelog](https://docs.slack.dev/changelog/2026/04/16/block-kit-new-blocks).
|
||||
`AlertBlock` is constructed via `NewAlertBlock` with a `*TextBlockObject`
|
||||
body and a functional-options pattern. Severity is set via
|
||||
`AlertBlockOptionLevel` (`AlertLevelDefault`, `AlertLevelInfo`,
|
||||
`AlertLevelWarning`, `AlertLevelError`, `AlertLevelSuccess`) and the block
|
||||
ID via `AlertBlockOptionBlockID`. Wires into `Blocks.UnmarshalJSON` for
|
||||
round-trip fidelity. Must be delivered via the streaming chunks API —
|
||||
`chat.postMessage` rejects it as an unsupported block type.
|
||||
- **Streaming-message chunks API** — `chat.startStream` / `chat.appendStream` /
|
||||
`chat.stopStream` now accept a `chunks` parameter. Added `MsgOptionChunks`
|
||||
along with a `StreamChunk` interface and four chunk types:
|
||||
`MarkdownTextChunk`, `TaskUpdateChunk`, `PlanUpdateChunk`, and `BlocksChunk`
|
||||
(each with a `New*Chunk` constructor). This is the supported transport for
|
||||
streaming Block Kit content and the new agent-UI blocks in particular
|
||||
(which `chat.postMessage` rejects as `Unsupported block type`).
|
||||
- **`MsgOptionTaskDisplayMode`** — New option for `chat.startStream` controlling
|
||||
whether task chunks render as a sequential timeline or a grouped plan.
|
||||
Accepts `TaskDisplayModeTimeline` or `TaskDisplayModePlan`.
|
||||
- Added `Username`, `IconURL`, and `IconEmoji` fields to
|
||||
`AssistantThreadsSetStatusParameters`, forwarded by
|
||||
`SetAssistantThreadsStatusContext`, matching the new optional parameters on
|
||||
[`assistant.threads.setStatus`](https://docs.slack.dev/reference/methods/assistant.threads.setStatus)
|
||||
for customising the status-update presentation.
|
||||
- Exposed `SocketmodeHandler.DispatchEvent` (previously the unexported
|
||||
`dispatcher`), enabling integration tests to exercise registered handlers
|
||||
without a live WebSocket connection. The unexported `dispatcher` is kept as
|
||||
a thin wrapper for backwards compatibility. Closes #1549.
|
||||
|
||||
## [0.22.0] - 2026-04-12
|
||||
|
||||
### Added
|
||||
|
||||
- Added missing parameters to `assistant.search.context` (`Sort`, `SortDir`, `Before`,
|
||||
`After`, `Highlight`, `IncludeContextMessages`, `IncludeDeletedUsers`,
|
||||
`IncludeMessageBlocks`, `IncludeArchivedChannels`, `DisableSemanticSearch`, `Modifiers`,
|
||||
`TermClauses`) and new response types (`AssistantSearchContextFile`,
|
||||
`AssistantSearchContextChannel`, `AssistantSearchContextMessageContext`) to match the
|
||||
full Real-Time Search API surface.
|
||||
- Added `Underline`, `Highlight`, `ClientHighlight`, and `Unlink` fields to
|
||||
`RichTextSectionTextStyle`. Added `Style` field to `RichTextSectionUserGroupElement`.
|
||||
- Added `BotOptional` and `UserOptional` fields to `OAuthScopes` for app manifests.
|
||||
- Added PKCE support for OAuth: `OAuthOptionCodeVerifier` option for
|
||||
`GetOAuthV2Response`, `GenerateCodeVerifier()` and `GenerateCodeChallenge()`
|
||||
helper functions (RFC 7636). `client_secret` is now conditionally omitted when
|
||||
empty in both `GetOAuthV2ResponseContext` and `RefreshOAuthV2TokenContext`.
|
||||
|
||||
### Fixed
|
||||
|
||||
- `ChannelTypes` and `ContentTypes` now send comma-separated values instead of repeated
|
||||
form keys, matching the convention used by every other method in the library.
|
||||
- In `socketmode` malformed JSON messages no longer force an unnecessary reconnect.
|
||||
Instead the error is emitted and the connection continues as normal.
|
||||
|
||||
## [0.21.1] - 2026-04-08
|
||||
|
||||
### Added
|
||||
@@ -489,7 +566,10 @@ for details.
|
||||
[#1196]: https://github.com/slack-go/slack/issues/1196
|
||||
[#1547]: https://github.com/slack-go/slack/pull/1547
|
||||
|
||||
[Unreleased]: https://github.com/slack-go/slack/compare/v0.21.1...HEAD
|
||||
[Unreleased]: https://github.com/slack-go/slack/compare/v0.23.1...HEAD
|
||||
[0.23.1]: https://github.com/slack-go/slack/compare/v0.23.0...v0.23.1
|
||||
[0.23.0]: https://github.com/slack-go/slack/compare/v0.22.0...v0.23.0
|
||||
[0.22.0]: https://github.com/slack-go/slack/compare/v0.21.1...0.22.0
|
||||
[0.21.1]: https://github.com/slack-go/slack/compare/v0.21.0...v0.21.1
|
||||
[0.21.0]: https://github.com/slack-go/slack/compare/v0.20.0...v0.21.0
|
||||
[0.20.0]: https://github.com/slack-go/slack/compare/v0.19.0...v0.20.0
|
||||
|
||||
+139
-24
@@ -14,6 +14,9 @@ type AssistantThreadsSetStatusParameters struct {
|
||||
Status string `json:"status"`
|
||||
ThreadTS string `json:"thread_ts"`
|
||||
LoadingMessages []string `json:"loading_messages,omitempty"`
|
||||
Username string `json:"username,omitempty"`
|
||||
IconURL string `json:"icon_url,omitempty"`
|
||||
IconEmoji string `json:"icon_emoji,omitempty"`
|
||||
}
|
||||
|
||||
// AssistantThreadSetTitleParameters are the parameters for AssistantThreadSetTitle
|
||||
@@ -39,30 +42,82 @@ type AssistantThreadsPrompt struct {
|
||||
|
||||
// AssistantSearchContextParameters are the parameters for AssistantSearchContext
|
||||
type AssistantSearchContextParameters struct {
|
||||
Query string `json:"query"`
|
||||
ActionToken string `json:"action_token,omitempty"`
|
||||
ChannelTypes []string `json:"channel_types,omitempty"`
|
||||
ContentTypes []string `json:"content_types,omitempty"`
|
||||
ContextChannelID string `json:"context_channel_id,omitempty"`
|
||||
Cursor string `json:"cursor,omitempty"`
|
||||
IncludeBots bool `json:"include_bots,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
Query string `json:"query"`
|
||||
ActionToken string `json:"action_token,omitempty"`
|
||||
ChannelTypes []string `json:"channel_types,omitempty"`
|
||||
ContentTypes []string `json:"content_types,omitempty"`
|
||||
ContextChannelID string `json:"context_channel_id,omitempty"`
|
||||
Cursor string `json:"cursor,omitempty"`
|
||||
IncludeBots bool `json:"include_bots,omitempty"`
|
||||
Limit int `json:"limit,omitempty"`
|
||||
IncludeDeletedUsers bool `json:"include_deleted_users,omitempty"`
|
||||
Before int64 `json:"before,omitempty"`
|
||||
After int64 `json:"after,omitempty"`
|
||||
IncludeContextMessages bool `json:"include_context_messages,omitempty"`
|
||||
Sort string `json:"sort,omitempty"`
|
||||
SortDir string `json:"sort_dir,omitempty"`
|
||||
IncludeMessageBlocks bool `json:"include_message_blocks,omitempty"`
|
||||
Highlight bool `json:"highlight,omitempty"`
|
||||
TermClauses []string `json:"term_clauses,omitempty"`
|
||||
Modifiers string `json:"modifiers,omitempty"`
|
||||
IncludeArchivedChannels bool `json:"include_archived_channels,omitempty"`
|
||||
DisableSemanticSearch bool `json:"disable_semantic_search,omitempty"`
|
||||
}
|
||||
|
||||
// AssistantSearchContextMessage represents a search result message
|
||||
type AssistantSearchContextMessage struct {
|
||||
AuthorUserID string `json:"author_user_id"`
|
||||
TeamID string `json:"team_id"`
|
||||
ChannelID string `json:"channel_id"`
|
||||
MessageTS string `json:"message_ts"`
|
||||
Content string `json:"content"`
|
||||
IsAuthorBot bool `json:"is_author_bot"`
|
||||
Permalink string `json:"permalink"`
|
||||
AuthorUserID string `json:"author_user_id"`
|
||||
AuthorName string `json:"author_name,omitempty"`
|
||||
TeamID string `json:"team_id"`
|
||||
ChannelID string `json:"channel_id"`
|
||||
ChannelName string `json:"channel_name,omitempty"`
|
||||
MessageTS string `json:"message_ts"`
|
||||
Content string `json:"content"`
|
||||
IsAuthorBot bool `json:"is_author_bot"`
|
||||
Permalink string `json:"permalink"`
|
||||
Blocks Blocks `json:"blocks,omitempty"`
|
||||
ContextMessages *AssistantSearchContextMessageContext `json:"context_messages,omitempty"`
|
||||
}
|
||||
|
||||
// AssistantSearchContextMessageContext contains context messages surrounding a search result
|
||||
type AssistantSearchContextMessageContext struct {
|
||||
Before []AssistantSearchContextMessage `json:"before"`
|
||||
After []AssistantSearchContextMessage `json:"after"`
|
||||
}
|
||||
|
||||
// AssistantSearchContextFile represents a search result file
|
||||
type AssistantSearchContextFile struct {
|
||||
UploaderUserID string `json:"uploader_user_id"`
|
||||
AuthorUserID string `json:"author_user_id"`
|
||||
AuthorName string `json:"author_name"`
|
||||
TeamID string `json:"team_id"`
|
||||
FileID string `json:"file_id"`
|
||||
DateCreated int64 `json:"date_created"`
|
||||
DateUpdated int64 `json:"date_updated"`
|
||||
Title string `json:"title"`
|
||||
FileType string `json:"file_type"`
|
||||
Permalink string `json:"permalink"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// AssistantSearchContextChannel represents a search result channel
|
||||
type AssistantSearchContextChannel struct {
|
||||
TeamID string `json:"team_id"`
|
||||
CreatorUserID string `json:"creator_user_id"`
|
||||
CreatorName string `json:"creator_name"`
|
||||
DateCreated int64 `json:"date_created"`
|
||||
DateUpdated int64 `json:"date_updated"`
|
||||
Name string `json:"name"`
|
||||
Topic string `json:"topic"`
|
||||
Purpose string `json:"purpose"`
|
||||
Permalink string `json:"permalink"`
|
||||
}
|
||||
|
||||
// AssistantSearchContextResults contains the search results
|
||||
type AssistantSearchContextResults struct {
|
||||
Messages []AssistantSearchContextMessage `json:"messages"`
|
||||
Messages []AssistantSearchContextMessage `json:"messages,omitempty"`
|
||||
Files []AssistantSearchContextFile `json:"files,omitempty"`
|
||||
Channels []AssistantSearchContextChannel `json:"channels,omitempty"`
|
||||
}
|
||||
|
||||
// AssistantSearchContextResponse is the response from assistant.search.context
|
||||
@@ -126,13 +181,17 @@ func (api *Client) SetAssistantThreadsSuggestedPromptsContext(ctx context.Contex
|
||||
return response.Err()
|
||||
}
|
||||
|
||||
// SetAssistantThreadStatus sets the status of a thread
|
||||
// SetAssistantThreadsStatus sets the status of a thread.
|
||||
// This method accepts either the chat:write or assistant:write scope.
|
||||
// Note: the assistant:write scope is being deprecated in favor of chat:write.
|
||||
// @see https://api.slack.com/methods/assistant.threads.setStatus
|
||||
func (api *Client) SetAssistantThreadsStatus(params AssistantThreadsSetStatusParameters) (err error) {
|
||||
return api.SetAssistantThreadsStatusContext(context.Background(), params)
|
||||
}
|
||||
|
||||
// SetAssistantThreadStatusContext sets the status of a thread with a custom context
|
||||
// SetAssistantThreadsStatusContext sets the status of a thread with a custom context.
|
||||
// This method accepts either the chat:write or assistant:write scope.
|
||||
// Note: the assistant:write scope is being deprecated in favor of chat:write.
|
||||
// @see https://api.slack.com/methods/assistant.threads.setStatus
|
||||
func (api *Client) SetAssistantThreadsStatusContext(ctx context.Context, params AssistantThreadsSetStatusParameters) (err error) {
|
||||
|
||||
@@ -153,6 +212,18 @@ func (api *Client) SetAssistantThreadsStatusContext(ctx context.Context, params
|
||||
values.Add("loading_messages", strings.Join(params.LoadingMessages, ","))
|
||||
}
|
||||
|
||||
if params.Username != "" {
|
||||
values.Add("username", params.Username)
|
||||
}
|
||||
|
||||
if params.IconURL != "" {
|
||||
values.Add("icon_url", params.IconURL)
|
||||
}
|
||||
|
||||
if params.IconEmoji != "" {
|
||||
values.Add("icon_emoji", params.IconEmoji)
|
||||
}
|
||||
|
||||
response := struct {
|
||||
SlackResponse
|
||||
}{}
|
||||
@@ -224,15 +295,11 @@ func (api *Client) SearchAssistantContextContext(ctx context.Context, params Ass
|
||||
}
|
||||
|
||||
if len(params.ChannelTypes) > 0 {
|
||||
for _, channelType := range params.ChannelTypes {
|
||||
values.Add("channel_types", channelType)
|
||||
}
|
||||
values.Add("channel_types", strings.Join(params.ChannelTypes, ","))
|
||||
}
|
||||
|
||||
if len(params.ContentTypes) > 0 {
|
||||
for _, contentType := range params.ContentTypes {
|
||||
values.Add("content_types", contentType)
|
||||
}
|
||||
values.Add("content_types", strings.Join(params.ContentTypes, ","))
|
||||
}
|
||||
|
||||
if params.ContextChannelID != "" {
|
||||
@@ -251,6 +318,54 @@ func (api *Client) SearchAssistantContextContext(ctx context.Context, params Ass
|
||||
values.Add("limit", strconv.Itoa(params.Limit))
|
||||
}
|
||||
|
||||
if params.IncludeDeletedUsers {
|
||||
values.Add("include_deleted_users", "true")
|
||||
}
|
||||
|
||||
if params.Before > 0 {
|
||||
values.Add("before", strconv.FormatInt(params.Before, 10))
|
||||
}
|
||||
|
||||
if params.After > 0 {
|
||||
values.Add("after", strconv.FormatInt(params.After, 10))
|
||||
}
|
||||
|
||||
if params.IncludeContextMessages {
|
||||
values.Add("include_context_messages", "true")
|
||||
}
|
||||
|
||||
if params.Sort != "" {
|
||||
values.Add("sort", params.Sort)
|
||||
}
|
||||
|
||||
if params.SortDir != "" {
|
||||
values.Add("sort_dir", params.SortDir)
|
||||
}
|
||||
|
||||
if params.IncludeMessageBlocks {
|
||||
values.Add("include_message_blocks", "true")
|
||||
}
|
||||
|
||||
if params.Highlight {
|
||||
values.Add("highlight", "true")
|
||||
}
|
||||
|
||||
if len(params.TermClauses) > 0 {
|
||||
values.Add("term_clauses", strings.Join(params.TermClauses, ","))
|
||||
}
|
||||
|
||||
if params.Modifiers != "" {
|
||||
values.Add("modifiers", params.Modifiers)
|
||||
}
|
||||
|
||||
if params.IncludeArchivedChannels {
|
||||
values.Add("include_archived_channels", "true")
|
||||
}
|
||||
|
||||
if params.DisableSemanticSearch {
|
||||
values.Add("disable_semantic_search", "true")
|
||||
}
|
||||
|
||||
response := &AssistantSearchContextResponse{}
|
||||
|
||||
err := api.postMethod(ctx, "assistant.search.context", values, response)
|
||||
|
||||
+3
@@ -21,6 +21,9 @@ const (
|
||||
MBTTable MessageBlockType = "table"
|
||||
MBTTaskCard MessageBlockType = "task_card"
|
||||
MBTPlan MessageBlockType = "plan"
|
||||
MBTAlert MessageBlockType = "alert"
|
||||
MBTCard MessageBlockType = "card"
|
||||
MBTCarousel MessageBlockType = "carousel"
|
||||
)
|
||||
|
||||
// Block defines an interface all block types should implement
|
||||
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
package slack
|
||||
|
||||
// AlertLevel defines the severity for an AlertBlock.
|
||||
type AlertLevel string
|
||||
|
||||
const (
|
||||
AlertLevelDefault AlertLevel = "default"
|
||||
AlertLevelInfo AlertLevel = "info"
|
||||
AlertLevelWarning AlertLevel = "warning"
|
||||
AlertLevelError AlertLevel = "error"
|
||||
AlertLevelSuccess AlertLevel = "success"
|
||||
)
|
||||
|
||||
// AlertBlock defines a block of type alert used to surface a notification
|
||||
// message with an optional severity level.
|
||||
//
|
||||
// Surface: modal only. Slack rejects alert blocks sent via chat.postMessage
|
||||
// or the streaming APIs — use OpenView / UpdateView / PushView with a
|
||||
// ModalViewRequest whose Blocks include the alert.
|
||||
//
|
||||
// More Information: https://docs.slack.dev/reference/block-kit/blocks/alert-block/
|
||||
type AlertBlock struct {
|
||||
Type MessageBlockType `json:"type"`
|
||||
Text *TextBlockObject `json:"text"`
|
||||
Level AlertLevel `json:"level,omitempty"`
|
||||
BlockID string `json:"block_id,omitempty"`
|
||||
}
|
||||
|
||||
// BlockType returns the type of the block
|
||||
func (s AlertBlock) BlockType() MessageBlockType {
|
||||
return s.Type
|
||||
}
|
||||
|
||||
// ID returns the ID of the block
|
||||
func (s AlertBlock) ID() string {
|
||||
return s.BlockID
|
||||
}
|
||||
|
||||
// AlertBlockOption allows configuration of options for a new alert block
|
||||
type AlertBlockOption func(*AlertBlock)
|
||||
|
||||
// AlertBlockOptionLevel sets the severity level for the alert block
|
||||
func AlertBlockOptionLevel(level AlertLevel) AlertBlockOption {
|
||||
return func(block *AlertBlock) {
|
||||
block.Level = level
|
||||
}
|
||||
}
|
||||
|
||||
// AlertBlockOptionBlockID sets the block ID for the alert block
|
||||
func AlertBlockOptionBlockID(blockID string) AlertBlockOption {
|
||||
return func(block *AlertBlock) {
|
||||
block.BlockID = blockID
|
||||
}
|
||||
}
|
||||
|
||||
// NewAlertBlock returns a new instance of an alert block
|
||||
func NewAlertBlock(text *TextBlockObject, options ...AlertBlockOption) *AlertBlock {
|
||||
block := AlertBlock{
|
||||
Type: MBTAlert,
|
||||
Text: text,
|
||||
}
|
||||
|
||||
for _, option := range options {
|
||||
if option != nil {
|
||||
option(&block)
|
||||
}
|
||||
}
|
||||
|
||||
return &block
|
||||
}
|
||||
+90
@@ -0,0 +1,90 @@
|
||||
package slack
|
||||
|
||||
// CardBlock defines a block of type card used to display a rich, self-contained
|
||||
// piece of content with an optional hero image, icon, title, subtitle, body,
|
||||
// and action buttons. Cards can stand alone or be grouped inside a
|
||||
// CarouselBlock.
|
||||
//
|
||||
// More Information: https://docs.slack.dev/reference/block-kit/blocks/card-block/
|
||||
type CardBlock struct {
|
||||
Type MessageBlockType `json:"type"`
|
||||
BlockID string `json:"block_id,omitempty"`
|
||||
HeroImage *ImageBlockElement `json:"hero_image,omitempty"`
|
||||
Icon *ImageBlockElement `json:"icon,omitempty"`
|
||||
Title *TextBlockObject `json:"title,omitempty"`
|
||||
Subtitle *TextBlockObject `json:"subtitle,omitempty"`
|
||||
Body *TextBlockObject `json:"body,omitempty"`
|
||||
Actions *BlockElements `json:"actions,omitempty"`
|
||||
}
|
||||
|
||||
// BlockType returns the type of the block
|
||||
func (s CardBlock) BlockType() MessageBlockType {
|
||||
return s.Type
|
||||
}
|
||||
|
||||
// ID returns the ID of the block
|
||||
func (s CardBlock) ID() string {
|
||||
return s.BlockID
|
||||
}
|
||||
|
||||
// CardBlockOption allows configuration of options for a new card block
|
||||
type CardBlockOption func(*CardBlock)
|
||||
|
||||
// CardBlockOptionBlockID sets the block ID for the card block
|
||||
func CardBlockOptionBlockID(blockID string) CardBlockOption {
|
||||
return func(block *CardBlock) {
|
||||
block.BlockID = blockID
|
||||
}
|
||||
}
|
||||
|
||||
// NewCardBlock returns a new instance of a card block. Use the chainable
|
||||
// With* methods or provide options to populate its fields.
|
||||
func NewCardBlock(options ...CardBlockOption) *CardBlock {
|
||||
block := CardBlock{
|
||||
Type: MBTCard,
|
||||
}
|
||||
|
||||
for _, option := range options {
|
||||
if option != nil {
|
||||
option(&block)
|
||||
}
|
||||
}
|
||||
|
||||
return &block
|
||||
}
|
||||
|
||||
// WithTitle sets the title text for the CardBlock
|
||||
func (s *CardBlock) WithTitle(title *TextBlockObject) *CardBlock {
|
||||
s.Title = title
|
||||
return s
|
||||
}
|
||||
|
||||
// WithSubtitle sets the subtitle text for the CardBlock
|
||||
func (s *CardBlock) WithSubtitle(subtitle *TextBlockObject) *CardBlock {
|
||||
s.Subtitle = subtitle
|
||||
return s
|
||||
}
|
||||
|
||||
// WithBody sets the body text for the CardBlock
|
||||
func (s *CardBlock) WithBody(body *TextBlockObject) *CardBlock {
|
||||
s.Body = body
|
||||
return s
|
||||
}
|
||||
|
||||
// WithIcon sets the icon image for the CardBlock
|
||||
func (s *CardBlock) WithIcon(icon *ImageBlockElement) *CardBlock {
|
||||
s.Icon = icon
|
||||
return s
|
||||
}
|
||||
|
||||
// WithHeroImage sets the hero image for the CardBlock
|
||||
func (s *CardBlock) WithHeroImage(heroImage *ImageBlockElement) *CardBlock {
|
||||
s.HeroImage = heroImage
|
||||
return s
|
||||
}
|
||||
|
||||
// WithActions sets the action buttons displayed at the bottom of the card
|
||||
func (s *CardBlock) WithActions(elements ...BlockElement) *CardBlock {
|
||||
s.Actions = &BlockElements{ElementSet: elements}
|
||||
return s
|
||||
}
|
||||
+42
@@ -0,0 +1,42 @@
|
||||
package slack
|
||||
|
||||
// CarouselBlock defines a block of type carousel that displays a scrollable
|
||||
// list of cards. A carousel must contain between 1 and 10 cards.
|
||||
//
|
||||
// More Information: https://docs.slack.dev/reference/block-kit/blocks/carousel-block/
|
||||
type CarouselBlock struct {
|
||||
Type MessageBlockType `json:"type"`
|
||||
BlockID string `json:"block_id,omitempty"`
|
||||
Elements []*CardBlock `json:"elements"`
|
||||
}
|
||||
|
||||
// BlockType returns the type of the block
|
||||
func (s CarouselBlock) BlockType() MessageBlockType {
|
||||
return s.Type
|
||||
}
|
||||
|
||||
// ID returns the ID of the block
|
||||
func (s CarouselBlock) ID() string {
|
||||
return s.BlockID
|
||||
}
|
||||
|
||||
// NewCarouselBlock returns a new instance of a carousel block containing the
|
||||
// given cards.
|
||||
func NewCarouselBlock(cards ...*CardBlock) *CarouselBlock {
|
||||
return &CarouselBlock{
|
||||
Type: MBTCarousel,
|
||||
Elements: cards,
|
||||
}
|
||||
}
|
||||
|
||||
// WithBlockID sets the block ID for the CarouselBlock
|
||||
func (s *CarouselBlock) WithBlockID(blockID string) *CarouselBlock {
|
||||
s.BlockID = blockID
|
||||
return s
|
||||
}
|
||||
|
||||
// AddCard appends a card to the carousel
|
||||
func (s *CarouselBlock) AddCard(card *CardBlock) *CarouselBlock {
|
||||
s.Elements = append(s.Elements, card)
|
||||
return s
|
||||
}
|
||||
+6
@@ -81,8 +81,14 @@ func (b *Blocks) UnmarshalJSON(data []byte) error {
|
||||
block = &TableBlock{}
|
||||
case "task_card":
|
||||
block = &TaskCardBlock{}
|
||||
case "alert":
|
||||
block = &AlertBlock{}
|
||||
case "plan":
|
||||
block = &PlanBlock{}
|
||||
case "card":
|
||||
block = &CardBlock{}
|
||||
case "carousel":
|
||||
block = &CarouselBlock{}
|
||||
default:
|
||||
b := &UnknownBlock{raw: r}
|
||||
if err = json.Unmarshal(r, b); err != nil {
|
||||
|
||||
+9
-4
@@ -298,10 +298,14 @@ type RichTextSectionElement interface {
|
||||
}
|
||||
|
||||
type RichTextSectionTextStyle struct {
|
||||
Bold bool `json:"bold,omitempty"`
|
||||
Italic bool `json:"italic,omitempty"`
|
||||
Strike bool `json:"strike,omitempty"`
|
||||
Code bool `json:"code,omitempty"`
|
||||
Bold bool `json:"bold,omitempty"`
|
||||
Italic bool `json:"italic,omitempty"`
|
||||
Strike bool `json:"strike,omitempty"`
|
||||
Code bool `json:"code,omitempty"`
|
||||
Underline bool `json:"underline,omitempty"`
|
||||
Highlight bool `json:"highlight,omitempty"`
|
||||
ClientHighlight bool `json:"client_highlight,omitempty"`
|
||||
Unlink bool `json:"unlink,omitempty"`
|
||||
}
|
||||
|
||||
type RichTextSectionTextElement struct {
|
||||
@@ -420,6 +424,7 @@ func NewRichTextSectionTeamElement(teamID string, style *RichTextSectionTextStyl
|
||||
type RichTextSectionUserGroupElement struct {
|
||||
Type RichTextSectionElementType `json:"type"`
|
||||
UsergroupID string `json:"usergroup_id"`
|
||||
Style *RichTextSectionTextStyle `json:"style,omitempty"`
|
||||
}
|
||||
|
||||
func (r RichTextSectionUserGroupElement) RichTextSectionElementType() RichTextSectionElementType {
|
||||
|
||||
+25
@@ -644,6 +644,7 @@ func MsgOptionDeleteOriginal(responseURL string) MsgOption {
|
||||
// MsgOptionAsUser whether or not to send the message as the user.
|
||||
func MsgOptionAsUser(b bool) MsgOption {
|
||||
return func(config *sendConfig) error {
|
||||
//lint:ignore S1002 - we want to explicitly check against the constant
|
||||
if b != DEFAULT_MESSAGE_ASUSER {
|
||||
config.values.Set("as_user", "true")
|
||||
}
|
||||
@@ -911,6 +912,24 @@ func MsgOptionMarkdownText(text string) MsgOption {
|
||||
}
|
||||
}
|
||||
|
||||
// TaskDisplayMode controls how task_card / task_update chunks render in a
|
||||
// streamed message. Used with chat.startStream.
|
||||
type TaskDisplayMode string
|
||||
|
||||
const (
|
||||
TaskDisplayModeTimeline TaskDisplayMode = "timeline"
|
||||
TaskDisplayModePlan TaskDisplayMode = "plan"
|
||||
)
|
||||
|
||||
// MsgOptionTaskDisplayMode sets task_display_mode on chat.startStream,
|
||||
// controlling whether tasks render as a sequential timeline or a grouped plan.
|
||||
func MsgOptionTaskDisplayMode(mode TaskDisplayMode) MsgOption {
|
||||
return func(config *sendConfig) error {
|
||||
config.values.Set("task_display_mode", string(mode))
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// UnsafeMsgOptionEndpoint deliver the message to the specified endpoint.
|
||||
// NOTE: USE AT YOUR OWN RISK: No issues relating to the use of this Option
|
||||
// will be supported by the library, it is subject to change without notice that
|
||||
@@ -945,15 +964,19 @@ func MsgOptionPostMessageParameters(params PostMessageParameters) MsgOption {
|
||||
config.values.Set("link_names", "1")
|
||||
}
|
||||
|
||||
//lint:ignore S1002 - we want to explicitly check against the constant
|
||||
if params.UnfurlLinks != DEFAULT_MESSAGE_UNFURL_LINKS {
|
||||
config.values.Set("unfurl_links", "true")
|
||||
}
|
||||
|
||||
// I want to send a message with explicit `as_user` `true` and `unfurl_links` `false` in request.
|
||||
// Because setting `as_user` to `true` will change the default value for `unfurl_links` to `true` on Slack API side.
|
||||
//lint:ignore S1002 - we want to explicitly check against the constants
|
||||
if params.AsUser != DEFAULT_MESSAGE_ASUSER && params.UnfurlLinks == DEFAULT_MESSAGE_UNFURL_LINKS {
|
||||
config.values.Set("unfurl_links", "false")
|
||||
}
|
||||
|
||||
//lint:ignore S1002 - we want to explicitly check against the constant
|
||||
if params.UnfurlMedia != DEFAULT_MESSAGE_UNFURL_MEDIA {
|
||||
config.values.Set("unfurl_media", "false")
|
||||
}
|
||||
@@ -963,6 +986,7 @@ func MsgOptionPostMessageParameters(params PostMessageParameters) MsgOption {
|
||||
if params.IconEmoji != DEFAULT_MESSAGE_ICON_EMOJI {
|
||||
config.values.Set("icon_emoji", params.IconEmoji)
|
||||
}
|
||||
//lint:ignore S1002 - we want to explicitly check against the constant
|
||||
if params.Markdown != DEFAULT_MESSAGE_MARKDOWN {
|
||||
config.values.Set("mrkdwn", "false")
|
||||
}
|
||||
@@ -970,6 +994,7 @@ func MsgOptionPostMessageParameters(params PostMessageParameters) MsgOption {
|
||||
if params.ThreadTimestamp != DEFAULT_MESSAGE_THREAD_TIMESTAMP {
|
||||
config.values.Set("thread_ts", params.ThreadTimestamp)
|
||||
}
|
||||
//lint:ignore S1002 - we want to explicitly check against the constant
|
||||
if params.ReplyBroadcast != DEFAULT_MESSAGE_REPLY_BROADCAST {
|
||||
config.values.Set("reply_broadcast", "true")
|
||||
}
|
||||
|
||||
+95
@@ -0,0 +1,95 @@
|
||||
package slack
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
)
|
||||
|
||||
// StreamChunkType identifies a chunk in the chat.startStream / chat.appendStream
|
||||
// / chat.stopStream streaming-message protocol.
|
||||
//
|
||||
// More information: https://docs.slack.dev/reference/methods/chat.appendStream/
|
||||
type StreamChunkType string
|
||||
|
||||
const (
|
||||
StreamChunkMarkdownText StreamChunkType = "markdown_text"
|
||||
StreamChunkTaskUpdate StreamChunkType = "task_update"
|
||||
StreamChunkPlanUpdate StreamChunkType = "plan_update"
|
||||
StreamChunkBlocks StreamChunkType = "blocks"
|
||||
)
|
||||
|
||||
// StreamChunk represents a single chunk in the streaming-message chunks array.
|
||||
type StreamChunk interface {
|
||||
ChunkType() StreamChunkType
|
||||
}
|
||||
|
||||
// MarkdownTextChunk streams markdown-formatted text.
|
||||
type MarkdownTextChunk struct {
|
||||
Type StreamChunkType `json:"type"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
func (c MarkdownTextChunk) ChunkType() StreamChunkType { return c.Type }
|
||||
|
||||
// NewMarkdownTextChunk returns a markdown_text chunk.
|
||||
func NewMarkdownTextChunk(text string) MarkdownTextChunk {
|
||||
return MarkdownTextChunk{Type: StreamChunkMarkdownText, Text: text}
|
||||
}
|
||||
|
||||
// TaskUpdateChunk streams a task status update that renders as a task card.
|
||||
type TaskUpdateChunk struct {
|
||||
Type StreamChunkType `json:"type"`
|
||||
ID string `json:"id"`
|
||||
Title string `json:"title"`
|
||||
Status TaskCardStatus `json:"status,omitempty"`
|
||||
Details string `json:"details,omitempty"`
|
||||
Output string `json:"output,omitempty"`
|
||||
Sources []TaskCardSource `json:"sources,omitempty"`
|
||||
}
|
||||
|
||||
func (c TaskUpdateChunk) ChunkType() StreamChunkType { return c.Type }
|
||||
|
||||
// NewTaskUpdateChunk returns a task_update chunk with the given id and title.
|
||||
func NewTaskUpdateChunk(id, title string) TaskUpdateChunk {
|
||||
return TaskUpdateChunk{Type: StreamChunkTaskUpdate, ID: id, Title: title}
|
||||
}
|
||||
|
||||
// PlanUpdateChunk streams an update to the current plan's title.
|
||||
type PlanUpdateChunk struct {
|
||||
Type StreamChunkType `json:"type"`
|
||||
Title string `json:"title"`
|
||||
}
|
||||
|
||||
func (c PlanUpdateChunk) ChunkType() StreamChunkType { return c.Type }
|
||||
|
||||
// NewPlanUpdateChunk returns a plan_update chunk.
|
||||
func NewPlanUpdateChunk(title string) PlanUpdateChunk {
|
||||
return PlanUpdateChunk{Type: StreamChunkPlanUpdate, Title: title}
|
||||
}
|
||||
|
||||
// BlocksChunk streams a group of Block Kit blocks. Up to 50 blocks per chunk.
|
||||
type BlocksChunk struct {
|
||||
Type StreamChunkType `json:"type"`
|
||||
Blocks []Block `json:"blocks"`
|
||||
}
|
||||
|
||||
func (c BlocksChunk) ChunkType() StreamChunkType { return c.Type }
|
||||
|
||||
// NewBlocksChunk returns a blocks chunk containing the given blocks.
|
||||
func NewBlocksChunk(blocks ...Block) BlocksChunk {
|
||||
return BlocksChunk{Type: StreamChunkBlocks, Blocks: blocks}
|
||||
}
|
||||
|
||||
// MsgOptionChunks sets the `chunks` parameter for the streaming chat methods
|
||||
// (chat.startStream / chat.appendStream / chat.stopStream). It is the
|
||||
// transport for Block Kit agent-UI blocks (Alert, Card, Carousel, etc.) which
|
||||
// chat.postMessage rejects as "Unsupported block type".
|
||||
func MsgOptionChunks(chunks ...StreamChunk) MsgOption {
|
||||
return func(config *sendConfig) error {
|
||||
encoded, err := json.Marshal(chunks)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
config.values.Set("chunks", string(encoded))
|
||||
return nil
|
||||
}
|
||||
}
|
||||
+1
@@ -310,6 +310,7 @@ func (api *Client) GetFilesContext(ctx context.Context, params GetFilesParameter
|
||||
if params.Page != DEFAULT_FILES_PAGE {
|
||||
values.Add("page", strconv.Itoa(params.Page))
|
||||
}
|
||||
//lint:ignore S1002 - we want to explicitly check against the constant
|
||||
if params.ShowHidden != DEFAULT_FILES_SHOW_HIDDEN {
|
||||
values.Add("show_files_hidden_by_limit", strconv.FormatBool(params.ShowHidden))
|
||||
}
|
||||
|
||||
+4
-2
@@ -269,8 +269,10 @@ type OAuthConfig struct {
|
||||
|
||||
// OAuthScopes is a group of settings that describe permission scopes configuration
|
||||
type OAuthScopes struct {
|
||||
Bot []string `json:"bot,omitempty" yaml:"bot,omitempty"`
|
||||
User []string `json:"user,omitempty" yaml:"user,omitempty"`
|
||||
Bot []string `json:"bot,omitempty" yaml:"bot,omitempty"`
|
||||
User []string `json:"user,omitempty" yaml:"user,omitempty"`
|
||||
BotOptional []string `json:"bot_optional,omitempty" yaml:"bot_optional,omitempty"`
|
||||
UserOptional []string `json:"user_optional,omitempty" yaml:"user_optional,omitempty"`
|
||||
}
|
||||
|
||||
// ManifestResponse is the response returned by the API for apps.manifest.x endpoints
|
||||
|
||||
+1
@@ -1,3 +1,4 @@
|
||||
[tools]
|
||||
go = "1.25"
|
||||
golangci-lint = "2.10.1"
|
||||
"go:honnef.co/go/tools/cmd/staticcheck" = "2026.1"
|
||||
|
||||
+51
-9
@@ -2,6 +2,9 @@ package slack
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"net/url"
|
||||
)
|
||||
|
||||
@@ -80,7 +83,8 @@ type OpenIDConnectResponse struct {
|
||||
}
|
||||
|
||||
type oauthConfig struct {
|
||||
apiURL string
|
||||
apiURL string
|
||||
codeVerifier string
|
||||
}
|
||||
|
||||
// OAuthOption configures package-level OAuth functions.
|
||||
@@ -91,12 +95,22 @@ func OAuthOptionAPIURL(url string) OAuthOption {
|
||||
return func(c *oauthConfig) { c.apiURL = url }
|
||||
}
|
||||
|
||||
func resolveOAuthAPIURL(opts []OAuthOption) string {
|
||||
// OAuthOptionCodeVerifier sets the PKCE code_verifier for the OAuth token exchange.
|
||||
// Use this when your authorization request included a code_challenge.
|
||||
func OAuthOptionCodeVerifier(verifier string) OAuthOption {
|
||||
return func(c *oauthConfig) { c.codeVerifier = verifier }
|
||||
}
|
||||
|
||||
func resolveOAuthConfig(opts []OAuthOption) oauthConfig {
|
||||
c := oauthConfig{apiURL: APIURL}
|
||||
for _, o := range opts {
|
||||
o(&c)
|
||||
}
|
||||
return c.apiURL
|
||||
return c
|
||||
}
|
||||
|
||||
func resolveOAuthAPIURL(opts []OAuthOption) string {
|
||||
return resolveOAuthConfig(opts).apiURL
|
||||
}
|
||||
|
||||
// GetOAuthToken retrieves an AccessToken.
|
||||
@@ -160,16 +174,23 @@ func GetOAuthV2Response(client httpClient, clientID, clientSecret, code, redirec
|
||||
}
|
||||
|
||||
// GetOAuthV2ResponseContext with a context, gets a V2 OAuth access token response.
|
||||
// For PKCE flows, pass OAuthOptionCodeVerifier and an empty clientSecret.
|
||||
// Slack API docs: https://api.slack.com/methods/oauth.v2.access
|
||||
func GetOAuthV2ResponseContext(ctx context.Context, client httpClient, clientID, clientSecret, code, redirectURI string, opts ...OAuthOption) (resp *OAuthV2Response, err error) {
|
||||
cfg := resolveOAuthConfig(opts)
|
||||
values := url.Values{
|
||||
"client_id": {clientID},
|
||||
"client_secret": {clientSecret},
|
||||
"code": {code},
|
||||
"redirect_uri": {redirectURI},
|
||||
"client_id": {clientID},
|
||||
"code": {code},
|
||||
"redirect_uri": {redirectURI},
|
||||
}
|
||||
if clientSecret != "" {
|
||||
values.Set("client_secret", clientSecret)
|
||||
}
|
||||
if cfg.codeVerifier != "" {
|
||||
values.Set("code_verifier", cfg.codeVerifier)
|
||||
}
|
||||
response := &OAuthV2Response{}
|
||||
if _, err = postForm(ctx, client, resolveOAuthAPIURL(opts)+"oauth.v2.access", values, response, discard{}); err != nil {
|
||||
if _, err = postForm(ctx, client, cfg.apiURL+"oauth.v2.access", values, response, discard{}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return response, response.Err()
|
||||
@@ -182,14 +203,17 @@ func RefreshOAuthV2Token(client httpClient, clientID, clientSecret, refreshToken
|
||||
}
|
||||
|
||||
// RefreshOAuthV2TokenContext with a context, gets a V2 OAuth access token response.
|
||||
// For PKCE public clients, pass an empty clientSecret.
|
||||
// Slack API docs: https://api.slack.com/methods/oauth.v2.access
|
||||
func RefreshOAuthV2TokenContext(ctx context.Context, client httpClient, clientID, clientSecret, refreshToken string, opts ...OAuthOption) (resp *OAuthV2Response, err error) {
|
||||
values := url.Values{
|
||||
"client_id": {clientID},
|
||||
"client_secret": {clientSecret},
|
||||
"refresh_token": {refreshToken},
|
||||
"grant_type": {"refresh_token"},
|
||||
}
|
||||
if clientSecret != "" {
|
||||
values.Set("client_secret", clientSecret)
|
||||
}
|
||||
response := &OAuthV2Response{}
|
||||
if _, err = postForm(ctx, client, resolveOAuthAPIURL(opts)+"oauth.v2.access", values, response, discard{}); err != nil {
|
||||
return nil, err
|
||||
@@ -286,3 +310,21 @@ func GetOpenIDConnectTokenContext(ctx context.Context, client httpClient, client
|
||||
}
|
||||
return response, response.Err()
|
||||
}
|
||||
|
||||
// GenerateCodeVerifier creates a cryptographically random PKCE code verifier
|
||||
// string suitable for use with OAuth 2.0 PKCE flows. The returned string is
|
||||
// 43 characters of URL-safe base64 (no padding).
|
||||
func GenerateCodeVerifier() (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
// GenerateCodeChallenge creates a PKCE code challenge from a code verifier
|
||||
// using the S256 method (SHA-256 hash, base64url-encoded without padding).
|
||||
func GenerateCodeChallenge(verifier string) string {
|
||||
h := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
+6
-2
@@ -30,6 +30,10 @@ func unsafeSignatureVerifier(header http.Header, secret string) (_ SecretsVerifi
|
||||
bsignature []byte
|
||||
)
|
||||
|
||||
if secret == "" {
|
||||
return SecretsVerifier{}, ErrInvalidConfiguration
|
||||
}
|
||||
|
||||
signature := header.Get(hSignature)
|
||||
stimestamp := header.Get(hTimestamp)
|
||||
|
||||
@@ -42,7 +46,7 @@ func unsafeSignatureVerifier(header http.Header, secret string) (_ SecretsVerifi
|
||||
}
|
||||
|
||||
hash := hmac.New(sha256.New, []byte(secret))
|
||||
if _, err = hash.Write([]byte(fmt.Sprintf("v0:%s:", stimestamp))); err != nil {
|
||||
if _, err = fmt.Fprintf(hash, "v0:%s:", stimestamp); err != nil {
|
||||
return SecretsVerifier{}, err
|
||||
}
|
||||
|
||||
@@ -95,7 +99,7 @@ func (v SecretsVerifier) Ensure() error {
|
||||
if v.d != nil && v.d.Debug() {
|
||||
v.d.Debugln(fmt.Sprintf("Expected signing signature: %s, but computed: %s", hex.EncodeToString(v.signature), hex.EncodeToString(computed)))
|
||||
}
|
||||
return fmt.Errorf("Computed unexpected signature of: %s", hex.EncodeToString(computed))
|
||||
return fmt.Errorf("computed unexpected signature of: %s", hex.EncodeToString(computed))
|
||||
}
|
||||
|
||||
func abs64(n int64) int64 {
|
||||
|
||||
+3
-1
@@ -36,7 +36,9 @@ func (api *Client) StartSocketModeContext(ctx context.Context) (info *SocketMode
|
||||
// time significantly shorter (360 seconds).
|
||||
if api.debug {
|
||||
u, _ := url.Parse(response.SocketModeConnection.URL)
|
||||
u.Query().Add("debug_reconnects", "true")
|
||||
q := u.Query()
|
||||
q.Set("debug_reconnects", "true")
|
||||
u.RawQuery = q.Encode()
|
||||
response.SocketModeConnection.URL = u.String()
|
||||
}
|
||||
return &response.SocketModeConnection, response.SocketModeConnection.URL, response.Err()
|
||||
|
||||
+5
-1
@@ -167,7 +167,11 @@ type OrderError struct {
|
||||
}
|
||||
|
||||
func (oe *OrderError) Error() string {
|
||||
return fmt.Sprintf("acme: order %s status: %s", oe.OrderURL, oe.Status)
|
||||
str := fmt.Sprintf("acme: order %s status: %s", oe.OrderURL, oe.Status)
|
||||
if oe.Problem != nil {
|
||||
str += fmt.Sprintf("; problem: %s", oe.Problem)
|
||||
}
|
||||
return str
|
||||
}
|
||||
|
||||
// RateLimit reports whether err represents a rate limit error and
|
||||
|
||||
+13
-60
@@ -2,24 +2,17 @@
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
/*
|
||||
Package pbkdf2 implements the key derivation function PBKDF2 as defined in RFC
|
||||
2898 / PKCS #5 v2.0.
|
||||
|
||||
A key derivation function is useful when encrypting data based on a password
|
||||
or any other not-fully-random data. It uses a pseudorandom function to derive
|
||||
a secure encryption key based on the password.
|
||||
|
||||
While v2.0 of the standard defines only one pseudorandom function to use,
|
||||
HMAC-SHA1, the drafted v2.1 specification allows use of all five FIPS Approved
|
||||
Hash Functions SHA-1, SHA-224, SHA-256, SHA-384 and SHA-512 for HMAC. To
|
||||
choose, you can pass the `New` functions from the different SHA packages to
|
||||
pbkdf2.Key.
|
||||
*/
|
||||
// Package pbkdf2 implements the key derivation function PBKDF2 as defined in
|
||||
// RFC 8018 (PKCS #5 v2.1).
|
||||
//
|
||||
// This package is a wrapper for the PBKDF2 implementation in the
|
||||
// [crypto/pbkdf2] package. It is [frozen] and is not accepting new features.
|
||||
//
|
||||
// [frozen]: https://go.dev/wiki/Frozen
|
||||
package pbkdf2
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/pbkdf2"
|
||||
"hash"
|
||||
)
|
||||
|
||||
@@ -27,51 +20,11 @@ import (
|
||||
// []byte of length keylen that can be used as cryptographic key. The key is
|
||||
// derived based on the method described as PBKDF2 with the HMAC variant using
|
||||
// the supplied hash function.
|
||||
//
|
||||
// For example, to use a HMAC-SHA-1 based PBKDF2 key derivation function, you
|
||||
// can get a derived key for e.g. AES-256 (which needs a 32-byte key) by
|
||||
// doing:
|
||||
//
|
||||
// dk := pbkdf2.Key([]byte("some password"), salt, 4096, 32, sha1.New)
|
||||
//
|
||||
// Remember to get a good random salt. At least 8 bytes is recommended by the
|
||||
// RFC.
|
||||
//
|
||||
// Using a higher iteration count will increase the cost of an exhaustive
|
||||
// search but will also make derivation proportionally slower.
|
||||
func Key(password, salt []byte, iter, keyLen int, h func() hash.Hash) []byte {
|
||||
prf := hmac.New(h, password)
|
||||
hashLen := prf.Size()
|
||||
numBlocks := (keyLen + hashLen - 1) / hashLen
|
||||
|
||||
var buf [4]byte
|
||||
dk := make([]byte, 0, numBlocks*hashLen)
|
||||
U := make([]byte, hashLen)
|
||||
for block := 1; block <= numBlocks; block++ {
|
||||
// N.B.: || means concatenation, ^ means XOR
|
||||
// for each block T_i = U_1 ^ U_2 ^ ... ^ U_iter
|
||||
// U_1 = PRF(password, salt || uint(i))
|
||||
prf.Reset()
|
||||
prf.Write(salt)
|
||||
buf[0] = byte(block >> 24)
|
||||
buf[1] = byte(block >> 16)
|
||||
buf[2] = byte(block >> 8)
|
||||
buf[3] = byte(block)
|
||||
prf.Write(buf[:4])
|
||||
dk = prf.Sum(dk)
|
||||
T := dk[len(dk)-hashLen:]
|
||||
copy(U, T)
|
||||
|
||||
// U_n = PRF(password, U_(n-1))
|
||||
for n := 2; n <= iter; n++ {
|
||||
prf.Reset()
|
||||
prf.Write(U)
|
||||
U = U[:0]
|
||||
U = prf.Sum(U)
|
||||
for x := range U {
|
||||
T[x] ^= U[x]
|
||||
}
|
||||
}
|
||||
out, err := pbkdf2.Key(h, string(password), salt, iter, keyLen)
|
||||
if err != nil {
|
||||
// FIPS 140 enforcement, or an invalid key length.
|
||||
panic(err)
|
||||
}
|
||||
return dk[:keyLen]
|
||||
return out
|
||||
}
|
||||
|
||||
+193
@@ -0,0 +1,193 @@
|
||||
// Copyright 2014 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package riff implements the Resource Interchange File Format, used by media
|
||||
// formats such as AVI, WAVE and WEBP.
|
||||
//
|
||||
// A RIFF stream contains a sequence of chunks. Each chunk consists of an 8-byte
|
||||
// header (containing a 4-byte chunk type and a 4-byte chunk length), the chunk
|
||||
// data (presented as an io.Reader), and some padding bytes.
|
||||
//
|
||||
// A detailed description of the format is at
|
||||
// http://www.tactilemedia.com/info/MCI_Control_Info.html
|
||||
package riff // import "golang.org/x/image/riff"
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"io/ioutil"
|
||||
"math"
|
||||
)
|
||||
|
||||
var (
|
||||
errMissingPaddingByte = errors.New("riff: missing padding byte")
|
||||
errMissingRIFFChunkHeader = errors.New("riff: missing RIFF chunk header")
|
||||
errListSubchunkTooLong = errors.New("riff: list subchunk too long")
|
||||
errShortChunkData = errors.New("riff: short chunk data")
|
||||
errShortChunkHeader = errors.New("riff: short chunk header")
|
||||
errStaleReader = errors.New("riff: stale reader")
|
||||
)
|
||||
|
||||
// u32 decodes the first four bytes of b as a little-endian integer.
|
||||
func u32(b []byte) uint32 {
|
||||
return uint32(b[0]) | uint32(b[1])<<8 | uint32(b[2])<<16 | uint32(b[3])<<24
|
||||
}
|
||||
|
||||
const chunkHeaderSize = 8
|
||||
|
||||
// FourCC is a four character code.
|
||||
type FourCC [4]byte
|
||||
|
||||
// LIST is the "LIST" FourCC.
|
||||
var LIST = FourCC{'L', 'I', 'S', 'T'}
|
||||
|
||||
// NewReader returns the RIFF stream's form type, such as "AVI " or "WAVE", and
|
||||
// its chunks as a *Reader.
|
||||
func NewReader(r io.Reader) (formType FourCC, data *Reader, err error) {
|
||||
var buf [chunkHeaderSize]byte
|
||||
if _, err := io.ReadFull(r, buf[:]); err != nil {
|
||||
if err == io.EOF || err == io.ErrUnexpectedEOF {
|
||||
err = errMissingRIFFChunkHeader
|
||||
}
|
||||
return FourCC{}, nil, err
|
||||
}
|
||||
if buf[0] != 'R' || buf[1] != 'I' || buf[2] != 'F' || buf[3] != 'F' {
|
||||
return FourCC{}, nil, errMissingRIFFChunkHeader
|
||||
}
|
||||
return NewListReader(u32(buf[4:]), r)
|
||||
}
|
||||
|
||||
// NewListReader returns a LIST chunk's list type, such as "movi" or "wavl",
|
||||
// and its chunks as a *Reader.
|
||||
func NewListReader(chunkLen uint32, chunkData io.Reader) (listType FourCC, data *Reader, err error) {
|
||||
if chunkLen < 4 {
|
||||
return FourCC{}, nil, errShortChunkData
|
||||
}
|
||||
z := &Reader{r: chunkData}
|
||||
if _, err := io.ReadFull(chunkData, z.buf[:4]); err != nil {
|
||||
if err == io.EOF || err == io.ErrUnexpectedEOF {
|
||||
err = errShortChunkData
|
||||
}
|
||||
return FourCC{}, nil, err
|
||||
}
|
||||
z.totalLen = chunkLen - 4
|
||||
return FourCC{z.buf[0], z.buf[1], z.buf[2], z.buf[3]}, z, nil
|
||||
}
|
||||
|
||||
// Reader reads chunks from an underlying io.Reader.
|
||||
type Reader struct {
|
||||
r io.Reader
|
||||
err error
|
||||
|
||||
totalLen uint32
|
||||
chunkLen uint32
|
||||
|
||||
chunkReader *chunkReader
|
||||
buf [chunkHeaderSize]byte
|
||||
padded bool
|
||||
}
|
||||
|
||||
// Next returns the next chunk's ID, length and data. It returns io.EOF if there
|
||||
// are no more chunks. The io.Reader returned becomes stale after the next Next
|
||||
// call, and should no longer be used.
|
||||
//
|
||||
// It is valid to call Next even if all of the previous chunk's data has not
|
||||
// been read.
|
||||
func (z *Reader) Next() (chunkID FourCC, chunkLen uint32, chunkData io.Reader, err error) {
|
||||
if z.err != nil {
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
|
||||
// Drain the rest of the previous chunk.
|
||||
if z.chunkLen != 0 {
|
||||
want := z.chunkLen
|
||||
var got int64
|
||||
got, z.err = io.Copy(ioutil.Discard, z.chunkReader)
|
||||
if z.err == nil && uint32(got) != want {
|
||||
z.err = errShortChunkData
|
||||
}
|
||||
if z.err != nil {
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
}
|
||||
z.chunkReader = nil
|
||||
if z.padded {
|
||||
if z.totalLen == 0 {
|
||||
z.err = errListSubchunkTooLong
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
z.totalLen--
|
||||
_, z.err = io.ReadFull(z.r, z.buf[:1])
|
||||
if z.err != nil {
|
||||
if z.err == io.EOF {
|
||||
z.err = errMissingPaddingByte
|
||||
}
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
}
|
||||
|
||||
// We are done if we have no more data.
|
||||
if z.totalLen == 0 {
|
||||
z.err = io.EOF
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
|
||||
// Read the next chunk header.
|
||||
if z.totalLen < chunkHeaderSize {
|
||||
z.err = errShortChunkHeader
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
z.totalLen -= chunkHeaderSize
|
||||
if _, z.err = io.ReadFull(z.r, z.buf[:chunkHeaderSize]); z.err != nil {
|
||||
if z.err == io.EOF || z.err == io.ErrUnexpectedEOF {
|
||||
z.err = errShortChunkHeader
|
||||
}
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
chunkID = FourCC{z.buf[0], z.buf[1], z.buf[2], z.buf[3]}
|
||||
z.chunkLen = u32(z.buf[4:])
|
||||
if z.chunkLen > z.totalLen {
|
||||
z.err = errListSubchunkTooLong
|
||||
return FourCC{}, 0, nil, z.err
|
||||
}
|
||||
z.padded = z.chunkLen&1 == 1
|
||||
z.chunkReader = &chunkReader{z}
|
||||
return chunkID, z.chunkLen, z.chunkReader, nil
|
||||
}
|
||||
|
||||
type chunkReader struct {
|
||||
z *Reader
|
||||
}
|
||||
|
||||
func (c *chunkReader) Read(p []byte) (int, error) {
|
||||
if c != c.z.chunkReader {
|
||||
return 0, errStaleReader
|
||||
}
|
||||
z := c.z
|
||||
if z.err != nil {
|
||||
if z.err == io.EOF {
|
||||
return 0, errStaleReader
|
||||
}
|
||||
return 0, z.err
|
||||
}
|
||||
|
||||
n := int(z.chunkLen)
|
||||
if n == 0 {
|
||||
return 0, io.EOF
|
||||
}
|
||||
if n < 0 {
|
||||
// Converting uint32 to int overflowed.
|
||||
n = math.MaxInt32
|
||||
}
|
||||
if n > len(p) {
|
||||
n = len(p)
|
||||
}
|
||||
n, err := z.r.Read(p[:n])
|
||||
z.totalLen -= uint32(n)
|
||||
z.chunkLen -= uint32(n)
|
||||
if err != io.EOF {
|
||||
z.err = err
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
+403
@@ -0,0 +1,403 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package vp8 implements a decoder for the VP8 lossy image format.
|
||||
//
|
||||
// The VP8 specification is RFC 6386.
|
||||
package vp8 // import "golang.org/x/image/vp8"
|
||||
|
||||
// This file implements the top-level decoding algorithm.
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"image"
|
||||
"io"
|
||||
)
|
||||
|
||||
// limitReader wraps an io.Reader to read at most n bytes from it.
|
||||
type limitReader struct {
|
||||
r io.Reader
|
||||
n int
|
||||
}
|
||||
|
||||
// ReadFull reads exactly len(p) bytes into p.
|
||||
func (r *limitReader) ReadFull(p []byte) error {
|
||||
if len(p) > r.n {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
n, err := io.ReadFull(r.r, p)
|
||||
r.n -= n
|
||||
return err
|
||||
}
|
||||
|
||||
// FrameHeader is a frame header, as specified in section 9.1.
|
||||
type FrameHeader struct {
|
||||
KeyFrame bool
|
||||
VersionNumber uint8
|
||||
ShowFrame bool
|
||||
FirstPartitionLen uint32
|
||||
Width int
|
||||
Height int
|
||||
XScale uint8
|
||||
YScale uint8
|
||||
}
|
||||
|
||||
const (
|
||||
nSegment = 4
|
||||
nSegmentProb = 3
|
||||
)
|
||||
|
||||
// segmentHeader holds segment-related header information.
|
||||
type segmentHeader struct {
|
||||
useSegment bool
|
||||
updateMap bool
|
||||
relativeDelta bool
|
||||
quantizer [nSegment]int8
|
||||
filterStrength [nSegment]int8
|
||||
prob [nSegmentProb]uint8
|
||||
}
|
||||
|
||||
const (
|
||||
nRefLFDelta = 4
|
||||
nModeLFDelta = 4
|
||||
)
|
||||
|
||||
// filterHeader holds filter-related header information.
|
||||
type filterHeader struct {
|
||||
simple bool
|
||||
level int8
|
||||
sharpness uint8
|
||||
useLFDelta bool
|
||||
refLFDelta [nRefLFDelta]int8
|
||||
modeLFDelta [nModeLFDelta]int8
|
||||
perSegmentLevel [nSegment]int8
|
||||
}
|
||||
|
||||
// mb is the per-macroblock decode state. A decoder maintains mbw+1 of these
|
||||
// as it is decoding macroblocks left-to-right and top-to-bottom: mbw for the
|
||||
// macroblocks in the row above, and one for the macroblock to the left.
|
||||
type mb struct {
|
||||
// pred is the predictor mode for the 4 bottom or right 4x4 luma regions.
|
||||
pred [4]uint8
|
||||
// nzMask is a mask of 8 bits: 4 for the bottom or right 4x4 luma regions,
|
||||
// and 2 + 2 for the bottom or right 4x4 chroma regions. A 1 bit indicates
|
||||
// that region has non-zero coefficients.
|
||||
nzMask uint8
|
||||
// nzY16 is a 0/1 value that is 1 if the macroblock used Y16 prediction and
|
||||
// had non-zero coefficients.
|
||||
nzY16 uint8
|
||||
}
|
||||
|
||||
// Decoder decodes VP8 bitstreams into frames. Decoding one frame consists of
|
||||
// calling Init, DecodeFrameHeader and then DecodeFrame in that order.
|
||||
// A Decoder can be re-used to decode multiple frames.
|
||||
type Decoder struct {
|
||||
// r is the input bitsream.
|
||||
r limitReader
|
||||
// scratch is a scratch buffer.
|
||||
scratch [8]byte
|
||||
// img is the YCbCr image to decode into.
|
||||
img *image.YCbCr
|
||||
// mbw and mbh are the number of 16x16 macroblocks wide and high the image is.
|
||||
mbw, mbh int
|
||||
// frameHeader is the frame header. When decoding multiple frames,
|
||||
// frames that aren't key frames will inherit the Width, Height,
|
||||
// XScale and YScale of the most recent key frame.
|
||||
frameHeader FrameHeader
|
||||
// Other headers.
|
||||
segmentHeader segmentHeader
|
||||
filterHeader filterHeader
|
||||
// The image data is divided into a number of independent partitions.
|
||||
// There is 1 "first partition" and between 1 and 8 "other partitions"
|
||||
// for coefficient data.
|
||||
fp partition
|
||||
op [8]partition
|
||||
nOP int
|
||||
// Quantization factors.
|
||||
quant [nSegment]quant
|
||||
// DCT/WHT coefficient decoding probabilities.
|
||||
tokenProb [nPlane][nBand][nContext][nProb]uint8
|
||||
useSkipProb bool
|
||||
skipProb uint8
|
||||
// Loop filter parameters.
|
||||
filterParams [nSegment][2]filterParam
|
||||
perMBFilterParams []filterParam
|
||||
|
||||
// The eight fields below relate to the current macroblock being decoded.
|
||||
//
|
||||
// Segment-based adjustments.
|
||||
segment int
|
||||
// Per-macroblock state for the macroblock immediately left of and those
|
||||
// macroblocks immediately above the current macroblock.
|
||||
leftMB mb
|
||||
upMB []mb
|
||||
// Bitmasks for which 4x4 regions of coeff contain non-zero coefficients.
|
||||
nzDCMask, nzACMask uint32
|
||||
// Predictor modes.
|
||||
usePredY16 bool // The libwebp C code calls this !is_i4x4_.
|
||||
predY16 uint8
|
||||
predC8 uint8
|
||||
predY4 [4][4]uint8
|
||||
|
||||
// The two fields below form a workspace for reconstructing a macroblock.
|
||||
// Their specific sizes are documented in reconstruct.go.
|
||||
coeff [1*16*16 + 2*8*8 + 1*4*4]int16
|
||||
ybr [1 + 16 + 1 + 8][32]uint8
|
||||
}
|
||||
|
||||
// NewDecoder returns a new Decoder.
|
||||
func NewDecoder() *Decoder {
|
||||
return &Decoder{}
|
||||
}
|
||||
|
||||
// Init initializes the decoder to read at most n bytes from r.
|
||||
func (d *Decoder) Init(r io.Reader, n int) {
|
||||
d.r = limitReader{r, n}
|
||||
}
|
||||
|
||||
// DecodeFrameHeader decodes the frame header.
|
||||
func (d *Decoder) DecodeFrameHeader() (fh FrameHeader, err error) {
|
||||
// All frame headers are at least 3 bytes long.
|
||||
b := d.scratch[:3]
|
||||
if err = d.r.ReadFull(b); err != nil {
|
||||
return
|
||||
}
|
||||
d.frameHeader.KeyFrame = (b[0] & 1) == 0
|
||||
d.frameHeader.VersionNumber = (b[0] >> 1) & 7
|
||||
d.frameHeader.ShowFrame = (b[0]>>4)&1 == 1
|
||||
d.frameHeader.FirstPartitionLen = uint32(b[0])>>5 | uint32(b[1])<<3 | uint32(b[2])<<11
|
||||
if !d.frameHeader.KeyFrame {
|
||||
return d.frameHeader, nil
|
||||
}
|
||||
// Frame headers for key frames are an additional 7 bytes long.
|
||||
b = d.scratch[:7]
|
||||
if err = d.r.ReadFull(b); err != nil {
|
||||
return
|
||||
}
|
||||
// Check the magic sync code.
|
||||
if b[0] != 0x9d || b[1] != 0x01 || b[2] != 0x2a {
|
||||
err = errors.New("vp8: invalid format")
|
||||
return
|
||||
}
|
||||
d.frameHeader.Width = int(b[4]&0x3f)<<8 | int(b[3])
|
||||
d.frameHeader.Height = int(b[6]&0x3f)<<8 | int(b[5])
|
||||
d.frameHeader.XScale = b[4] >> 6
|
||||
d.frameHeader.YScale = b[6] >> 6
|
||||
d.mbw = (d.frameHeader.Width + 0x0f) >> 4
|
||||
d.mbh = (d.frameHeader.Height + 0x0f) >> 4
|
||||
d.segmentHeader = segmentHeader{
|
||||
prob: [3]uint8{0xff, 0xff, 0xff},
|
||||
}
|
||||
d.tokenProb = defaultTokenProb
|
||||
d.segment = 0
|
||||
return d.frameHeader, nil
|
||||
}
|
||||
|
||||
// ensureImg ensures that d.img is large enough to hold the decoded frame.
|
||||
func (d *Decoder) ensureImg() {
|
||||
if d.img != nil {
|
||||
p0, p1 := d.img.Rect.Min, d.img.Rect.Max
|
||||
if p0.X == 0 && p0.Y == 0 && p1.X >= 16*d.mbw && p1.Y >= 16*d.mbh {
|
||||
return
|
||||
}
|
||||
}
|
||||
m := image.NewYCbCr(image.Rect(0, 0, 16*d.mbw, 16*d.mbh), image.YCbCrSubsampleRatio420)
|
||||
d.img = m.SubImage(image.Rect(0, 0, d.frameHeader.Width, d.frameHeader.Height)).(*image.YCbCr)
|
||||
d.perMBFilterParams = make([]filterParam, d.mbw*d.mbh)
|
||||
d.upMB = make([]mb, d.mbw)
|
||||
}
|
||||
|
||||
// parseSegmentHeader parses the segment header, as specified in section 9.3.
|
||||
func (d *Decoder) parseSegmentHeader() {
|
||||
d.segmentHeader.useSegment = d.fp.readBit(uniformProb)
|
||||
if !d.segmentHeader.useSegment {
|
||||
d.segmentHeader.updateMap = false
|
||||
return
|
||||
}
|
||||
d.segmentHeader.updateMap = d.fp.readBit(uniformProb)
|
||||
if d.fp.readBit(uniformProb) {
|
||||
d.segmentHeader.relativeDelta = !d.fp.readBit(uniformProb)
|
||||
for i := range d.segmentHeader.quantizer {
|
||||
d.segmentHeader.quantizer[i] = int8(d.fp.readOptionalInt(uniformProb, 7))
|
||||
}
|
||||
for i := range d.segmentHeader.filterStrength {
|
||||
d.segmentHeader.filterStrength[i] = int8(d.fp.readOptionalInt(uniformProb, 6))
|
||||
}
|
||||
}
|
||||
if !d.segmentHeader.updateMap {
|
||||
return
|
||||
}
|
||||
for i := range d.segmentHeader.prob {
|
||||
if d.fp.readBit(uniformProb) {
|
||||
d.segmentHeader.prob[i] = uint8(d.fp.readUint(uniformProb, 8))
|
||||
} else {
|
||||
d.segmentHeader.prob[i] = 0xff
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// parseFilterHeader parses the filter header, as specified in section 9.4.
|
||||
func (d *Decoder) parseFilterHeader() {
|
||||
d.filterHeader.simple = d.fp.readBit(uniformProb)
|
||||
d.filterHeader.level = int8(d.fp.readUint(uniformProb, 6))
|
||||
d.filterHeader.sharpness = uint8(d.fp.readUint(uniformProb, 3))
|
||||
d.filterHeader.useLFDelta = d.fp.readBit(uniformProb)
|
||||
if d.filterHeader.useLFDelta && d.fp.readBit(uniformProb) {
|
||||
for i := range d.filterHeader.refLFDelta {
|
||||
d.filterHeader.refLFDelta[i] = int8(d.fp.readOptionalInt(uniformProb, 6))
|
||||
}
|
||||
for i := range d.filterHeader.modeLFDelta {
|
||||
d.filterHeader.modeLFDelta[i] = int8(d.fp.readOptionalInt(uniformProb, 6))
|
||||
}
|
||||
}
|
||||
if d.filterHeader.level == 0 {
|
||||
return
|
||||
}
|
||||
if d.segmentHeader.useSegment {
|
||||
for i := range d.filterHeader.perSegmentLevel {
|
||||
strength := d.segmentHeader.filterStrength[i]
|
||||
if d.segmentHeader.relativeDelta {
|
||||
strength += d.filterHeader.level
|
||||
}
|
||||
d.filterHeader.perSegmentLevel[i] = strength
|
||||
}
|
||||
} else {
|
||||
d.filterHeader.perSegmentLevel[0] = d.filterHeader.level
|
||||
}
|
||||
d.computeFilterParams()
|
||||
}
|
||||
|
||||
// parseOtherPartitions parses the other partitions, as specified in section 9.5.
|
||||
func (d *Decoder) parseOtherPartitions() error {
|
||||
const maxNOP = 1 << 3
|
||||
var partLens [maxNOP]int
|
||||
d.nOP = 1 << d.fp.readUint(uniformProb, 2)
|
||||
|
||||
// The final partition length is implied by the remaining chunk data
|
||||
// (d.r.n) and the other d.nOP-1 partition lengths. Those d.nOP-1 partition
|
||||
// lengths are stored as 24-bit uints, i.e. up to 16 MiB per partition.
|
||||
n := 3 * (d.nOP - 1)
|
||||
partLens[d.nOP-1] = d.r.n - n
|
||||
if partLens[d.nOP-1] < 0 {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
if n > 0 {
|
||||
buf := make([]byte, n)
|
||||
if err := d.r.ReadFull(buf); err != nil {
|
||||
return err
|
||||
}
|
||||
for i := 0; i < d.nOP-1; i++ {
|
||||
pl := int(buf[3*i+0]) | int(buf[3*i+1])<<8 | int(buf[3*i+2])<<16
|
||||
if pl > partLens[d.nOP-1] {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
partLens[i] = pl
|
||||
partLens[d.nOP-1] -= pl
|
||||
}
|
||||
}
|
||||
|
||||
// We check if the final partition length can also fit into a 24-bit uint.
|
||||
// Strictly speaking, this isn't part of the spec, but it guards against a
|
||||
// malicious WEBP image that is too large to ReadFull the encoded DCT
|
||||
// coefficients into memory, whether that's because the actual WEBP file is
|
||||
// too large, or whether its RIFF metadata lists too large a chunk.
|
||||
if 1<<24 <= partLens[d.nOP-1] {
|
||||
return errors.New("vp8: too much data to decode")
|
||||
}
|
||||
|
||||
buf := make([]byte, d.r.n)
|
||||
if err := d.r.ReadFull(buf); err != nil {
|
||||
return err
|
||||
}
|
||||
for i, pl := range partLens {
|
||||
if i == d.nOP {
|
||||
break
|
||||
}
|
||||
d.op[i].init(buf[:pl])
|
||||
buf = buf[pl:]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// parseOtherHeaders parses header information other than the frame header.
|
||||
func (d *Decoder) parseOtherHeaders() error {
|
||||
// Initialize and parse the first partition.
|
||||
firstPartition := make([]byte, d.frameHeader.FirstPartitionLen)
|
||||
if err := d.r.ReadFull(firstPartition); err != nil {
|
||||
return err
|
||||
}
|
||||
d.fp.init(firstPartition)
|
||||
if d.frameHeader.KeyFrame {
|
||||
// Read and ignore the color space and pixel clamp values. They are
|
||||
// specified in section 9.2, but are unimplemented.
|
||||
d.fp.readBit(uniformProb)
|
||||
d.fp.readBit(uniformProb)
|
||||
}
|
||||
d.parseSegmentHeader()
|
||||
d.parseFilterHeader()
|
||||
if err := d.parseOtherPartitions(); err != nil {
|
||||
return err
|
||||
}
|
||||
d.parseQuant()
|
||||
if !d.frameHeader.KeyFrame {
|
||||
// Golden and AltRef frames are specified in section 9.7.
|
||||
// TODO(nigeltao): implement. Note that they are only used for video, not still images.
|
||||
return errors.New("vp8: Golden / AltRef frames are not implemented")
|
||||
}
|
||||
// Read and ignore the refreshLastFrameBuffer bit, specified in section 9.8.
|
||||
// It applies only to video, and not still images.
|
||||
d.fp.readBit(uniformProb)
|
||||
d.parseTokenProb()
|
||||
d.useSkipProb = d.fp.readBit(uniformProb)
|
||||
if d.useSkipProb {
|
||||
d.skipProb = uint8(d.fp.readUint(uniformProb, 8))
|
||||
}
|
||||
if d.fp.unexpectedEOF {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DecodeFrame decodes the frame and returns it as an YCbCr image.
|
||||
// The image's contents are valid up until the next call to Decoder.Init.
|
||||
func (d *Decoder) DecodeFrame() (*image.YCbCr, error) {
|
||||
d.ensureImg()
|
||||
if err := d.parseOtherHeaders(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Reconstruct the rows.
|
||||
for mbx := 0; mbx < d.mbw; mbx++ {
|
||||
d.upMB[mbx] = mb{}
|
||||
}
|
||||
for mby := 0; mby < d.mbh; mby++ {
|
||||
d.leftMB = mb{}
|
||||
for mbx := 0; mbx < d.mbw; mbx++ {
|
||||
skip := d.reconstruct(mbx, mby)
|
||||
fs := d.filterParams[d.segment][btou(!d.usePredY16)]
|
||||
fs.inner = fs.inner || !skip
|
||||
d.perMBFilterParams[d.mbw*mby+mbx] = fs
|
||||
}
|
||||
}
|
||||
if d.fp.unexpectedEOF {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
for i := 0; i < d.nOP; i++ {
|
||||
if d.op[i].unexpectedEOF {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
}
|
||||
// Apply the loop filter.
|
||||
//
|
||||
// Even if we are using per-segment levels, section 15 says that "loop
|
||||
// filtering must be skipped entirely if loop_filter_level at either the
|
||||
// frame header level or macroblock override level is 0".
|
||||
if d.filterHeader.level != 0 {
|
||||
if d.filterHeader.simple {
|
||||
d.simpleFilter()
|
||||
} else {
|
||||
d.normalFilter()
|
||||
}
|
||||
}
|
||||
return d.img, nil
|
||||
}
|
||||
+273
@@ -0,0 +1,273 @@
|
||||
// Copyright 2014 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// filter2 modifies a 2-pixel wide or 2-pixel high band along an edge.
|
||||
func filter2(pix []byte, level, index, iStep, jStep int) {
|
||||
for n := 16; n > 0; n, index = n-1, index+iStep {
|
||||
p1 := int(pix[index-2*jStep])
|
||||
p0 := int(pix[index-1*jStep])
|
||||
q0 := int(pix[index+0*jStep])
|
||||
q1 := int(pix[index+1*jStep])
|
||||
if abs(p0-q0)<<1+abs(p1-q1)>>1 > level {
|
||||
continue
|
||||
}
|
||||
a := 3*(q0-p0) + clamp127(p1-q1)
|
||||
a1 := clamp15((a + 4) >> 3)
|
||||
a2 := clamp15((a + 3) >> 3)
|
||||
pix[index-1*jStep] = clamp255(p0 + a2)
|
||||
pix[index+0*jStep] = clamp255(q0 - a1)
|
||||
}
|
||||
}
|
||||
|
||||
// filter246 modifies a 2-, 4- or 6-pixel wide or high band along an edge.
|
||||
func filter246(pix []byte, n, level, ilevel, hlevel, index, iStep, jStep int, fourNotSix bool) {
|
||||
for ; n > 0; n, index = n-1, index+iStep {
|
||||
p3 := int(pix[index-4*jStep])
|
||||
p2 := int(pix[index-3*jStep])
|
||||
p1 := int(pix[index-2*jStep])
|
||||
p0 := int(pix[index-1*jStep])
|
||||
q0 := int(pix[index+0*jStep])
|
||||
q1 := int(pix[index+1*jStep])
|
||||
q2 := int(pix[index+2*jStep])
|
||||
q3 := int(pix[index+3*jStep])
|
||||
if abs(p0-q0)<<1+abs(p1-q1)>>1 > level {
|
||||
continue
|
||||
}
|
||||
if abs(p3-p2) > ilevel ||
|
||||
abs(p2-p1) > ilevel ||
|
||||
abs(p1-p0) > ilevel ||
|
||||
abs(q1-q0) > ilevel ||
|
||||
abs(q2-q1) > ilevel ||
|
||||
abs(q3-q2) > ilevel {
|
||||
continue
|
||||
}
|
||||
if abs(p1-p0) > hlevel || abs(q1-q0) > hlevel {
|
||||
// Filter 2 pixels.
|
||||
a := 3*(q0-p0) + clamp127(p1-q1)
|
||||
a1 := clamp15((a + 4) >> 3)
|
||||
a2 := clamp15((a + 3) >> 3)
|
||||
pix[index-1*jStep] = clamp255(p0 + a2)
|
||||
pix[index+0*jStep] = clamp255(q0 - a1)
|
||||
} else if fourNotSix {
|
||||
// Filter 4 pixels.
|
||||
a := 3 * (q0 - p0)
|
||||
a1 := clamp15((a + 4) >> 3)
|
||||
a2 := clamp15((a + 3) >> 3)
|
||||
a3 := (a1 + 1) >> 1
|
||||
pix[index-2*jStep] = clamp255(p1 + a3)
|
||||
pix[index-1*jStep] = clamp255(p0 + a2)
|
||||
pix[index+0*jStep] = clamp255(q0 - a1)
|
||||
pix[index+1*jStep] = clamp255(q1 - a3)
|
||||
} else {
|
||||
// Filter 6 pixels.
|
||||
a := clamp127(3*(q0-p0) + clamp127(p1-q1))
|
||||
a1 := (27*a + 63) >> 7
|
||||
a2 := (18*a + 63) >> 7
|
||||
a3 := (9*a + 63) >> 7
|
||||
pix[index-3*jStep] = clamp255(p2 + a3)
|
||||
pix[index-2*jStep] = clamp255(p1 + a2)
|
||||
pix[index-1*jStep] = clamp255(p0 + a1)
|
||||
pix[index+0*jStep] = clamp255(q0 - a1)
|
||||
pix[index+1*jStep] = clamp255(q1 - a2)
|
||||
pix[index+2*jStep] = clamp255(q2 - a3)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// simpleFilter implements the simple filter, as specified in section 15.2.
|
||||
func (d *Decoder) simpleFilter() {
|
||||
for mby := 0; mby < d.mbh; mby++ {
|
||||
for mbx := 0; mbx < d.mbw; mbx++ {
|
||||
f := d.perMBFilterParams[d.mbw*mby+mbx]
|
||||
if f.level == 0 {
|
||||
continue
|
||||
}
|
||||
l := int(f.level)
|
||||
yIndex := (mby*d.img.YStride + mbx) * 16
|
||||
if mbx > 0 {
|
||||
filter2(d.img.Y, l+4, yIndex, d.img.YStride, 1)
|
||||
}
|
||||
if f.inner {
|
||||
filter2(d.img.Y, l, yIndex+0x4, d.img.YStride, 1)
|
||||
filter2(d.img.Y, l, yIndex+0x8, d.img.YStride, 1)
|
||||
filter2(d.img.Y, l, yIndex+0xc, d.img.YStride, 1)
|
||||
}
|
||||
if mby > 0 {
|
||||
filter2(d.img.Y, l+4, yIndex, 1, d.img.YStride)
|
||||
}
|
||||
if f.inner {
|
||||
filter2(d.img.Y, l, yIndex+d.img.YStride*0x4, 1, d.img.YStride)
|
||||
filter2(d.img.Y, l, yIndex+d.img.YStride*0x8, 1, d.img.YStride)
|
||||
filter2(d.img.Y, l, yIndex+d.img.YStride*0xc, 1, d.img.YStride)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// normalFilter implements the normal filter, as specified in section 15.3.
|
||||
func (d *Decoder) normalFilter() {
|
||||
for mby := 0; mby < d.mbh; mby++ {
|
||||
for mbx := 0; mbx < d.mbw; mbx++ {
|
||||
f := d.perMBFilterParams[d.mbw*mby+mbx]
|
||||
if f.level == 0 {
|
||||
continue
|
||||
}
|
||||
l, il, hl := int(f.level), int(f.ilevel), int(f.hlevel)
|
||||
yIndex := (mby*d.img.YStride + mbx) * 16
|
||||
cIndex := (mby*d.img.CStride + mbx) * 8
|
||||
if mbx > 0 {
|
||||
filter246(d.img.Y, 16, l+4, il, hl, yIndex, d.img.YStride, 1, false)
|
||||
filter246(d.img.Cb, 8, l+4, il, hl, cIndex, d.img.CStride, 1, false)
|
||||
filter246(d.img.Cr, 8, l+4, il, hl, cIndex, d.img.CStride, 1, false)
|
||||
}
|
||||
if f.inner {
|
||||
filter246(d.img.Y, 16, l, il, hl, yIndex+0x4, d.img.YStride, 1, true)
|
||||
filter246(d.img.Y, 16, l, il, hl, yIndex+0x8, d.img.YStride, 1, true)
|
||||
filter246(d.img.Y, 16, l, il, hl, yIndex+0xc, d.img.YStride, 1, true)
|
||||
filter246(d.img.Cb, 8, l, il, hl, cIndex+0x4, d.img.CStride, 1, true)
|
||||
filter246(d.img.Cr, 8, l, il, hl, cIndex+0x4, d.img.CStride, 1, true)
|
||||
}
|
||||
if mby > 0 {
|
||||
filter246(d.img.Y, 16, l+4, il, hl, yIndex, 1, d.img.YStride, false)
|
||||
filter246(d.img.Cb, 8, l+4, il, hl, cIndex, 1, d.img.CStride, false)
|
||||
filter246(d.img.Cr, 8, l+4, il, hl, cIndex, 1, d.img.CStride, false)
|
||||
}
|
||||
if f.inner {
|
||||
filter246(d.img.Y, 16, l, il, hl, yIndex+d.img.YStride*0x4, 1, d.img.YStride, true)
|
||||
filter246(d.img.Y, 16, l, il, hl, yIndex+d.img.YStride*0x8, 1, d.img.YStride, true)
|
||||
filter246(d.img.Y, 16, l, il, hl, yIndex+d.img.YStride*0xc, 1, d.img.YStride, true)
|
||||
filter246(d.img.Cb, 8, l, il, hl, cIndex+d.img.CStride*0x4, 1, d.img.CStride, true)
|
||||
filter246(d.img.Cr, 8, l, il, hl, cIndex+d.img.CStride*0x4, 1, d.img.CStride, true)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// filterParam holds the loop filter parameters for a macroblock.
|
||||
type filterParam struct {
|
||||
// The first three fields are thresholds used by the loop filter to smooth
|
||||
// over the edges and interior of a macroblock. level is used by both the
|
||||
// simple and normal filters. The inner level and high edge variance level
|
||||
// are only used by the normal filter.
|
||||
level, ilevel, hlevel uint8
|
||||
// inner is whether the inner loop filter cannot be optimized out as a
|
||||
// no-op for this particular macroblock.
|
||||
inner bool
|
||||
}
|
||||
|
||||
// computeFilterParams computes the loop filter parameters, as specified in
|
||||
// section 15.4.
|
||||
func (d *Decoder) computeFilterParams() {
|
||||
for i := range d.filterParams {
|
||||
baseLevel := d.filterHeader.level
|
||||
if d.segmentHeader.useSegment {
|
||||
baseLevel = d.segmentHeader.filterStrength[i]
|
||||
if d.segmentHeader.relativeDelta {
|
||||
baseLevel += d.filterHeader.level
|
||||
}
|
||||
}
|
||||
|
||||
for j := range d.filterParams[i] {
|
||||
p := &d.filterParams[i][j]
|
||||
p.inner = j != 0
|
||||
level := baseLevel
|
||||
if d.filterHeader.useLFDelta {
|
||||
// The libwebp C code has a "TODO: only CURRENT is handled for now."
|
||||
level += d.filterHeader.refLFDelta[0]
|
||||
if j != 0 {
|
||||
level += d.filterHeader.modeLFDelta[0]
|
||||
}
|
||||
}
|
||||
if level <= 0 {
|
||||
p.level = 0
|
||||
continue
|
||||
}
|
||||
if level > 63 {
|
||||
level = 63
|
||||
}
|
||||
ilevel := level
|
||||
if d.filterHeader.sharpness > 0 {
|
||||
if d.filterHeader.sharpness > 4 {
|
||||
ilevel >>= 2
|
||||
} else {
|
||||
ilevel >>= 1
|
||||
}
|
||||
if x := int8(9 - d.filterHeader.sharpness); ilevel > x {
|
||||
ilevel = x
|
||||
}
|
||||
}
|
||||
if ilevel < 1 {
|
||||
ilevel = 1
|
||||
}
|
||||
p.ilevel = uint8(ilevel)
|
||||
p.level = uint8(2*level + ilevel)
|
||||
if d.frameHeader.KeyFrame {
|
||||
if level < 15 {
|
||||
p.hlevel = 0
|
||||
} else if level < 40 {
|
||||
p.hlevel = 1
|
||||
} else {
|
||||
p.hlevel = 2
|
||||
}
|
||||
} else {
|
||||
if level < 15 {
|
||||
p.hlevel = 0
|
||||
} else if level < 20 {
|
||||
p.hlevel = 1
|
||||
} else if level < 40 {
|
||||
p.hlevel = 2
|
||||
} else {
|
||||
p.hlevel = 3
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// intSize is either 32 or 64.
|
||||
const intSize = 32 << (^uint(0) >> 63)
|
||||
|
||||
func abs(x int) int {
|
||||
// m := -1 if x < 0. m := 0 otherwise.
|
||||
m := x >> (intSize - 1)
|
||||
|
||||
// In two's complement representation, the negative number
|
||||
// of any number (except the smallest one) can be computed
|
||||
// by flipping all the bits and add 1. This is faster than
|
||||
// code with a branch.
|
||||
// See Hacker's Delight, section 2-4.
|
||||
return (x ^ m) - m
|
||||
}
|
||||
|
||||
func clamp15(x int) int {
|
||||
if x < -16 {
|
||||
return -16
|
||||
}
|
||||
if x > 15 {
|
||||
return 15
|
||||
}
|
||||
return x
|
||||
}
|
||||
|
||||
func clamp127(x int) int {
|
||||
if x < -128 {
|
||||
return -128
|
||||
}
|
||||
if x > 127 {
|
||||
return 127
|
||||
}
|
||||
return x
|
||||
}
|
||||
|
||||
func clamp255(x int) uint8 {
|
||||
if x < 0 {
|
||||
return 0
|
||||
}
|
||||
if x > 255 {
|
||||
return 255
|
||||
}
|
||||
return uint8(x)
|
||||
}
|
||||
+98
@@ -0,0 +1,98 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// This file implements the inverse Discrete Cosine Transform and the inverse
|
||||
// Walsh Hadamard Transform (WHT), as specified in sections 14.3 and 14.4.
|
||||
|
||||
func clip8(i int32) uint8 {
|
||||
if i < 0 {
|
||||
return 0
|
||||
}
|
||||
if i > 255 {
|
||||
return 255
|
||||
}
|
||||
return uint8(i)
|
||||
}
|
||||
|
||||
func (z *Decoder) inverseDCT4(y, x, coeffBase int) {
|
||||
const (
|
||||
c1 = 85627 // 65536 * cos(pi/8) * sqrt(2).
|
||||
c2 = 35468 // 65536 * sin(pi/8) * sqrt(2).
|
||||
)
|
||||
var m [4][4]int32
|
||||
for i := 0; i < 4; i++ {
|
||||
a := int32(z.coeff[coeffBase+0]) + int32(z.coeff[coeffBase+8])
|
||||
b := int32(z.coeff[coeffBase+0]) - int32(z.coeff[coeffBase+8])
|
||||
c := (int32(z.coeff[coeffBase+4])*c2)>>16 - (int32(z.coeff[coeffBase+12])*c1)>>16
|
||||
d := (int32(z.coeff[coeffBase+4])*c1)>>16 + (int32(z.coeff[coeffBase+12])*c2)>>16
|
||||
m[i][0] = a + d
|
||||
m[i][1] = b + c
|
||||
m[i][2] = b - c
|
||||
m[i][3] = a - d
|
||||
coeffBase++
|
||||
}
|
||||
for j := 0; j < 4; j++ {
|
||||
dc := m[0][j] + 4
|
||||
a := dc + m[2][j]
|
||||
b := dc - m[2][j]
|
||||
c := (m[1][j]*c2)>>16 - (m[3][j]*c1)>>16
|
||||
d := (m[1][j]*c1)>>16 + (m[3][j]*c2)>>16
|
||||
z.ybr[y+j][x+0] = clip8(int32(z.ybr[y+j][x+0]) + (a+d)>>3)
|
||||
z.ybr[y+j][x+1] = clip8(int32(z.ybr[y+j][x+1]) + (b+c)>>3)
|
||||
z.ybr[y+j][x+2] = clip8(int32(z.ybr[y+j][x+2]) + (b-c)>>3)
|
||||
z.ybr[y+j][x+3] = clip8(int32(z.ybr[y+j][x+3]) + (a-d)>>3)
|
||||
}
|
||||
}
|
||||
|
||||
func (z *Decoder) inverseDCT4DCOnly(y, x, coeffBase int) {
|
||||
dc := (int32(z.coeff[coeffBase+0]) + 4) >> 3
|
||||
for j := 0; j < 4; j++ {
|
||||
for i := 0; i < 4; i++ {
|
||||
z.ybr[y+j][x+i] = clip8(int32(z.ybr[y+j][x+i]) + dc)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (z *Decoder) inverseDCT8(y, x, coeffBase int) {
|
||||
z.inverseDCT4(y+0, x+0, coeffBase+0*16)
|
||||
z.inverseDCT4(y+0, x+4, coeffBase+1*16)
|
||||
z.inverseDCT4(y+4, x+0, coeffBase+2*16)
|
||||
z.inverseDCT4(y+4, x+4, coeffBase+3*16)
|
||||
}
|
||||
|
||||
func (z *Decoder) inverseDCT8DCOnly(y, x, coeffBase int) {
|
||||
z.inverseDCT4DCOnly(y+0, x+0, coeffBase+0*16)
|
||||
z.inverseDCT4DCOnly(y+0, x+4, coeffBase+1*16)
|
||||
z.inverseDCT4DCOnly(y+4, x+0, coeffBase+2*16)
|
||||
z.inverseDCT4DCOnly(y+4, x+4, coeffBase+3*16)
|
||||
}
|
||||
|
||||
func (d *Decoder) inverseWHT16() {
|
||||
var m [16]int32
|
||||
for i := 0; i < 4; i++ {
|
||||
a0 := int32(d.coeff[384+0+i]) + int32(d.coeff[384+12+i])
|
||||
a1 := int32(d.coeff[384+4+i]) + int32(d.coeff[384+8+i])
|
||||
a2 := int32(d.coeff[384+4+i]) - int32(d.coeff[384+8+i])
|
||||
a3 := int32(d.coeff[384+0+i]) - int32(d.coeff[384+12+i])
|
||||
m[0+i] = a0 + a1
|
||||
m[8+i] = a0 - a1
|
||||
m[4+i] = a3 + a2
|
||||
m[12+i] = a3 - a2
|
||||
}
|
||||
out := 0
|
||||
for i := 0; i < 4; i++ {
|
||||
dc := m[0+i*4] + 3
|
||||
a0 := dc + m[3+i*4]
|
||||
a1 := m[1+i*4] + m[2+i*4]
|
||||
a2 := m[1+i*4] - m[2+i*4]
|
||||
a3 := dc - m[3+i*4]
|
||||
d.coeff[out+0] = int16((a0 + a1) >> 3)
|
||||
d.coeff[out+16] = int16((a3 + a2) >> 3)
|
||||
d.coeff[out+32] = int16((a0 - a1) >> 3)
|
||||
d.coeff[out+48] = int16((a3 - a2) >> 3)
|
||||
out += 64
|
||||
}
|
||||
}
|
||||
+129
@@ -0,0 +1,129 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// Each VP8 frame consists of between 2 and 9 bitstream partitions.
|
||||
// Each partition is byte-aligned and is independently arithmetic-encoded.
|
||||
//
|
||||
// This file implements decoding a partition's bitstream, as specified in
|
||||
// chapter 7. The implementation follows libwebp's approach instead of the
|
||||
// specification's reference C implementation. For example, we use a look-up
|
||||
// table instead of a for loop to recalibrate the encoded range.
|
||||
|
||||
var (
|
||||
lutShift = [127]uint8{
|
||||
7, 6, 6, 5, 5, 5, 5, 4, 4, 4, 4, 4, 4, 4, 4,
|
||||
3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3,
|
||||
2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
|
||||
2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
|
||||
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
|
||||
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
|
||||
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
|
||||
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
|
||||
}
|
||||
lutRangeM1 = [127]uint8{
|
||||
127,
|
||||
127, 191,
|
||||
127, 159, 191, 223,
|
||||
127, 143, 159, 175, 191, 207, 223, 239,
|
||||
127, 135, 143, 151, 159, 167, 175, 183, 191, 199, 207, 215, 223, 231, 239, 247,
|
||||
127, 131, 135, 139, 143, 147, 151, 155, 159, 163, 167, 171, 175, 179, 183, 187,
|
||||
191, 195, 199, 203, 207, 211, 215, 219, 223, 227, 231, 235, 239, 243, 247, 251,
|
||||
127, 129, 131, 133, 135, 137, 139, 141, 143, 145, 147, 149, 151, 153, 155, 157,
|
||||
159, 161, 163, 165, 167, 169, 171, 173, 175, 177, 179, 181, 183, 185, 187, 189,
|
||||
191, 193, 195, 197, 199, 201, 203, 205, 207, 209, 211, 213, 215, 217, 219, 221,
|
||||
223, 225, 227, 229, 231, 233, 235, 237, 239, 241, 243, 245, 247, 249, 251, 253,
|
||||
}
|
||||
)
|
||||
|
||||
// uniformProb represents a 50% probability that the next bit is 0.
|
||||
const uniformProb = 128
|
||||
|
||||
// partition holds arithmetic-coded bits.
|
||||
type partition struct {
|
||||
// buf is the input bytes.
|
||||
buf []byte
|
||||
// r is how many of buf's bytes have been consumed.
|
||||
r int
|
||||
// rangeM1 is range minus 1, where range is in the arithmetic coding sense,
|
||||
// not the Go language sense.
|
||||
rangeM1 uint32
|
||||
// bits and nBits hold those bits shifted out of buf but not yet consumed.
|
||||
bits uint32
|
||||
nBits uint8
|
||||
// unexpectedEOF tells whether we tried to read past buf.
|
||||
unexpectedEOF bool
|
||||
}
|
||||
|
||||
// init initializes the partition.
|
||||
func (p *partition) init(buf []byte) {
|
||||
p.buf = buf
|
||||
p.r = 0
|
||||
p.rangeM1 = 254
|
||||
p.bits = 0
|
||||
p.nBits = 0
|
||||
p.unexpectedEOF = false
|
||||
}
|
||||
|
||||
// readBit returns the next bit.
|
||||
func (p *partition) readBit(prob uint8) bool {
|
||||
if p.nBits < 8 {
|
||||
if p.r >= len(p.buf) {
|
||||
p.unexpectedEOF = true
|
||||
return false
|
||||
}
|
||||
// Expression split for 386 compiler.
|
||||
x := uint32(p.buf[p.r])
|
||||
p.bits |= x << (8 - p.nBits)
|
||||
p.r++
|
||||
p.nBits += 8
|
||||
}
|
||||
split := (p.rangeM1*uint32(prob))>>8 + 1
|
||||
bit := p.bits >= split<<8
|
||||
if bit {
|
||||
p.rangeM1 -= split
|
||||
p.bits -= split << 8
|
||||
} else {
|
||||
p.rangeM1 = split - 1
|
||||
}
|
||||
if p.rangeM1 < 127 {
|
||||
shift := lutShift[p.rangeM1]
|
||||
p.rangeM1 = uint32(lutRangeM1[p.rangeM1])
|
||||
p.bits <<= shift
|
||||
p.nBits -= shift
|
||||
}
|
||||
return bit
|
||||
}
|
||||
|
||||
// readUint returns the next n-bit unsigned integer.
|
||||
func (p *partition) readUint(prob, n uint8) uint32 {
|
||||
var u uint32
|
||||
for n > 0 {
|
||||
n--
|
||||
if p.readBit(prob) {
|
||||
u |= 1 << n
|
||||
}
|
||||
}
|
||||
return u
|
||||
}
|
||||
|
||||
// readInt returns the next n-bit signed integer.
|
||||
func (p *partition) readInt(prob, n uint8) int32 {
|
||||
u := p.readUint(prob, n)
|
||||
b := p.readBit(prob)
|
||||
if b {
|
||||
return -int32(u)
|
||||
}
|
||||
return int32(u)
|
||||
}
|
||||
|
||||
// readOptionalInt returns the next n-bit signed integer in an encoding
|
||||
// where the likely result is zero.
|
||||
func (p *partition) readOptionalInt(prob, n uint8) int32 {
|
||||
if !p.readBit(prob) {
|
||||
return 0
|
||||
}
|
||||
return p.readInt(prob, n)
|
||||
}
|
||||
+201
@@ -0,0 +1,201 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// This file implements parsing the predictor modes, as specified in chapter
|
||||
// 11.
|
||||
|
||||
func (d *Decoder) parsePredModeY16(mbx int) {
|
||||
var p uint8
|
||||
if !d.fp.readBit(156) {
|
||||
if !d.fp.readBit(163) {
|
||||
p = predDC
|
||||
} else {
|
||||
p = predVE
|
||||
}
|
||||
} else if !d.fp.readBit(128) {
|
||||
p = predHE
|
||||
} else {
|
||||
p = predTM
|
||||
}
|
||||
for i := 0; i < 4; i++ {
|
||||
d.upMB[mbx].pred[i] = p
|
||||
d.leftMB.pred[i] = p
|
||||
}
|
||||
d.predY16 = p
|
||||
}
|
||||
|
||||
func (d *Decoder) parsePredModeC8() {
|
||||
if !d.fp.readBit(142) {
|
||||
d.predC8 = predDC
|
||||
} else if !d.fp.readBit(114) {
|
||||
d.predC8 = predVE
|
||||
} else if !d.fp.readBit(183) {
|
||||
d.predC8 = predHE
|
||||
} else {
|
||||
d.predC8 = predTM
|
||||
}
|
||||
}
|
||||
|
||||
func (d *Decoder) parsePredModeY4(mbx int) {
|
||||
for j := 0; j < 4; j++ {
|
||||
p := d.leftMB.pred[j]
|
||||
for i := 0; i < 4; i++ {
|
||||
prob := &predProb[d.upMB[mbx].pred[i]][p]
|
||||
if !d.fp.readBit(prob[0]) {
|
||||
p = predDC
|
||||
} else if !d.fp.readBit(prob[1]) {
|
||||
p = predTM
|
||||
} else if !d.fp.readBit(prob[2]) {
|
||||
p = predVE
|
||||
} else if !d.fp.readBit(prob[3]) {
|
||||
if !d.fp.readBit(prob[4]) {
|
||||
p = predHE
|
||||
} else if !d.fp.readBit(prob[5]) {
|
||||
p = predRD
|
||||
} else {
|
||||
p = predVR
|
||||
}
|
||||
} else if !d.fp.readBit(prob[6]) {
|
||||
p = predLD
|
||||
} else if !d.fp.readBit(prob[7]) {
|
||||
p = predVL
|
||||
} else if !d.fp.readBit(prob[8]) {
|
||||
p = predHD
|
||||
} else {
|
||||
p = predHU
|
||||
}
|
||||
d.predY4[j][i] = p
|
||||
d.upMB[mbx].pred[i] = p
|
||||
}
|
||||
d.leftMB.pred[j] = p
|
||||
}
|
||||
}
|
||||
|
||||
// predProb are the probabilities to decode a 4x4 region's predictor mode given
|
||||
// the predictor modes of the regions above and left of it.
|
||||
// These values are specified in section 11.5.
|
||||
var predProb = [nPred][nPred][9]uint8{
|
||||
{
|
||||
{231, 120, 48, 89, 115, 113, 120, 152, 112},
|
||||
{152, 179, 64, 126, 170, 118, 46, 70, 95},
|
||||
{175, 69, 143, 80, 85, 82, 72, 155, 103},
|
||||
{56, 58, 10, 171, 218, 189, 17, 13, 152},
|
||||
{114, 26, 17, 163, 44, 195, 21, 10, 173},
|
||||
{121, 24, 80, 195, 26, 62, 44, 64, 85},
|
||||
{144, 71, 10, 38, 171, 213, 144, 34, 26},
|
||||
{170, 46, 55, 19, 136, 160, 33, 206, 71},
|
||||
{63, 20, 8, 114, 114, 208, 12, 9, 226},
|
||||
{81, 40, 11, 96, 182, 84, 29, 16, 36},
|
||||
},
|
||||
{
|
||||
{134, 183, 89, 137, 98, 101, 106, 165, 148},
|
||||
{72, 187, 100, 130, 157, 111, 32, 75, 80},
|
||||
{66, 102, 167, 99, 74, 62, 40, 234, 128},
|
||||
{41, 53, 9, 178, 241, 141, 26, 8, 107},
|
||||
{74, 43, 26, 146, 73, 166, 49, 23, 157},
|
||||
{65, 38, 105, 160, 51, 52, 31, 115, 128},
|
||||
{104, 79, 12, 27, 217, 255, 87, 17, 7},
|
||||
{87, 68, 71, 44, 114, 51, 15, 186, 23},
|
||||
{47, 41, 14, 110, 182, 183, 21, 17, 194},
|
||||
{66, 45, 25, 102, 197, 189, 23, 18, 22},
|
||||
},
|
||||
{
|
||||
{88, 88, 147, 150, 42, 46, 45, 196, 205},
|
||||
{43, 97, 183, 117, 85, 38, 35, 179, 61},
|
||||
{39, 53, 200, 87, 26, 21, 43, 232, 171},
|
||||
{56, 34, 51, 104, 114, 102, 29, 93, 77},
|
||||
{39, 28, 85, 171, 58, 165, 90, 98, 64},
|
||||
{34, 22, 116, 206, 23, 34, 43, 166, 73},
|
||||
{107, 54, 32, 26, 51, 1, 81, 43, 31},
|
||||
{68, 25, 106, 22, 64, 171, 36, 225, 114},
|
||||
{34, 19, 21, 102, 132, 188, 16, 76, 124},
|
||||
{62, 18, 78, 95, 85, 57, 50, 48, 51},
|
||||
},
|
||||
{
|
||||
{193, 101, 35, 159, 215, 111, 89, 46, 111},
|
||||
{60, 148, 31, 172, 219, 228, 21, 18, 111},
|
||||
{112, 113, 77, 85, 179, 255, 38, 120, 114},
|
||||
{40, 42, 1, 196, 245, 209, 10, 25, 109},
|
||||
{88, 43, 29, 140, 166, 213, 37, 43, 154},
|
||||
{61, 63, 30, 155, 67, 45, 68, 1, 209},
|
||||
{100, 80, 8, 43, 154, 1, 51, 26, 71},
|
||||
{142, 78, 78, 16, 255, 128, 34, 197, 171},
|
||||
{41, 40, 5, 102, 211, 183, 4, 1, 221},
|
||||
{51, 50, 17, 168, 209, 192, 23, 25, 82},
|
||||
},
|
||||
{
|
||||
{138, 31, 36, 171, 27, 166, 38, 44, 229},
|
||||
{67, 87, 58, 169, 82, 115, 26, 59, 179},
|
||||
{63, 59, 90, 180, 59, 166, 93, 73, 154},
|
||||
{40, 40, 21, 116, 143, 209, 34, 39, 175},
|
||||
{47, 15, 16, 183, 34, 223, 49, 45, 183},
|
||||
{46, 17, 33, 183, 6, 98, 15, 32, 183},
|
||||
{57, 46, 22, 24, 128, 1, 54, 17, 37},
|
||||
{65, 32, 73, 115, 28, 128, 23, 128, 205},
|
||||
{40, 3, 9, 115, 51, 192, 18, 6, 223},
|
||||
{87, 37, 9, 115, 59, 77, 64, 21, 47},
|
||||
},
|
||||
{
|
||||
{104, 55, 44, 218, 9, 54, 53, 130, 226},
|
||||
{64, 90, 70, 205, 40, 41, 23, 26, 57},
|
||||
{54, 57, 112, 184, 5, 41, 38, 166, 213},
|
||||
{30, 34, 26, 133, 152, 116, 10, 32, 134},
|
||||
{39, 19, 53, 221, 26, 114, 32, 73, 255},
|
||||
{31, 9, 65, 234, 2, 15, 1, 118, 73},
|
||||
{75, 32, 12, 51, 192, 255, 160, 43, 51},
|
||||
{88, 31, 35, 67, 102, 85, 55, 186, 85},
|
||||
{56, 21, 23, 111, 59, 205, 45, 37, 192},
|
||||
{55, 38, 70, 124, 73, 102, 1, 34, 98},
|
||||
},
|
||||
{
|
||||
{125, 98, 42, 88, 104, 85, 117, 175, 82},
|
||||
{95, 84, 53, 89, 128, 100, 113, 101, 45},
|
||||
{75, 79, 123, 47, 51, 128, 81, 171, 1},
|
||||
{57, 17, 5, 71, 102, 57, 53, 41, 49},
|
||||
{38, 33, 13, 121, 57, 73, 26, 1, 85},
|
||||
{41, 10, 67, 138, 77, 110, 90, 47, 114},
|
||||
{115, 21, 2, 10, 102, 255, 166, 23, 6},
|
||||
{101, 29, 16, 10, 85, 128, 101, 196, 26},
|
||||
{57, 18, 10, 102, 102, 213, 34, 20, 43},
|
||||
{117, 20, 15, 36, 163, 128, 68, 1, 26},
|
||||
},
|
||||
{
|
||||
{102, 61, 71, 37, 34, 53, 31, 243, 192},
|
||||
{69, 60, 71, 38, 73, 119, 28, 222, 37},
|
||||
{68, 45, 128, 34, 1, 47, 11, 245, 171},
|
||||
{62, 17, 19, 70, 146, 85, 55, 62, 70},
|
||||
{37, 43, 37, 154, 100, 163, 85, 160, 1},
|
||||
{63, 9, 92, 136, 28, 64, 32, 201, 85},
|
||||
{75, 15, 9, 9, 64, 255, 184, 119, 16},
|
||||
{86, 6, 28, 5, 64, 255, 25, 248, 1},
|
||||
{56, 8, 17, 132, 137, 255, 55, 116, 128},
|
||||
{58, 15, 20, 82, 135, 57, 26, 121, 40},
|
||||
},
|
||||
{
|
||||
{164, 50, 31, 137, 154, 133, 25, 35, 218},
|
||||
{51, 103, 44, 131, 131, 123, 31, 6, 158},
|
||||
{86, 40, 64, 135, 148, 224, 45, 183, 128},
|
||||
{22, 26, 17, 131, 240, 154, 14, 1, 209},
|
||||
{45, 16, 21, 91, 64, 222, 7, 1, 197},
|
||||
{56, 21, 39, 155, 60, 138, 23, 102, 213},
|
||||
{83, 12, 13, 54, 192, 255, 68, 47, 28},
|
||||
{85, 26, 85, 85, 128, 128, 32, 146, 171},
|
||||
{18, 11, 7, 63, 144, 171, 4, 4, 246},
|
||||
{35, 27, 10, 146, 174, 171, 12, 26, 128},
|
||||
},
|
||||
{
|
||||
{190, 80, 35, 99, 180, 80, 126, 54, 45},
|
||||
{85, 126, 47, 87, 176, 51, 41, 20, 32},
|
||||
{101, 75, 128, 139, 118, 146, 116, 128, 85},
|
||||
{56, 41, 15, 176, 236, 85, 37, 9, 62},
|
||||
{71, 30, 17, 119, 118, 255, 17, 18, 138},
|
||||
{101, 38, 60, 138, 55, 70, 43, 26, 142},
|
||||
{146, 36, 19, 30, 171, 255, 97, 27, 20},
|
||||
{138, 45, 61, 62, 219, 1, 81, 188, 64},
|
||||
{32, 41, 20, 117, 151, 142, 20, 21, 163},
|
||||
{112, 19, 12, 61, 195, 128, 48, 4, 24},
|
||||
},
|
||||
}
|
||||
+553
@@ -0,0 +1,553 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// This file implements the prediction functions, as specified in chapter 12.
|
||||
//
|
||||
// For each macroblock (of 1x16x16 luma and 2x8x8 chroma coefficients), the
|
||||
// luma values are either predicted as one large 16x16 region or 16 separate
|
||||
// 4x4 regions. The chroma values are always predicted as one 8x8 region.
|
||||
//
|
||||
// For 4x4 regions, the target block's predicted values (Xs) are a function of
|
||||
// its previously-decoded top and left border values, as well as a number of
|
||||
// pixels from the top-right:
|
||||
//
|
||||
// a b c d e f g h
|
||||
// p X X X X
|
||||
// q X X X X
|
||||
// r X X X X
|
||||
// s X X X X
|
||||
//
|
||||
// The predictor modes are:
|
||||
// - DC: all Xs = (b + c + d + e + p + q + r + s + 4) / 8.
|
||||
// - TM: the first X = (b + p - a), the second X = (c + p - a), and so on.
|
||||
// - VE: each X = the weighted average of its column's top value and that
|
||||
// value's neighbors, i.e. averages of abc, bcd, cde or def.
|
||||
// - HE: similar to VE except rows instead of columns, and the final row is
|
||||
// an average of r, s and s.
|
||||
// - RD, VR, LD, VL, HD, HU: these diagonal modes ("Right Down", "Vertical
|
||||
// Right", etc) are more complicated and are described in section 12.3.
|
||||
// All Xs are clipped to the range [0, 255].
|
||||
//
|
||||
// For 8x8 and 16x16 regions, the target block's predicted values are a
|
||||
// function of the top and left border values without the top-right overhang,
|
||||
// i.e. without the 8x8 or 16x16 equivalent of f, g and h. Furthermore:
|
||||
// - There are no diagonal predictor modes, only DC, TM, VE and HE.
|
||||
// - The DC mode has variants for macroblocks in the top row and/or left
|
||||
// column, i.e. for macroblocks with mby == 0 || mbx == 0.
|
||||
// - The VE and HE modes take only the column top or row left values; they do
|
||||
// not smooth that top/left value with its neighbors.
|
||||
|
||||
// nPred is the number of predictor modes, not including the Top/Left versions
|
||||
// of the DC predictor mode.
|
||||
const nPred = 10
|
||||
|
||||
const (
|
||||
predDC = iota
|
||||
predTM
|
||||
predVE
|
||||
predHE
|
||||
predRD
|
||||
predVR
|
||||
predLD
|
||||
predVL
|
||||
predHD
|
||||
predHU
|
||||
predDCTop
|
||||
predDCLeft
|
||||
predDCTopLeft
|
||||
)
|
||||
|
||||
func checkTopLeftPred(mbx, mby int, p uint8) uint8 {
|
||||
if p != predDC {
|
||||
return p
|
||||
}
|
||||
if mbx == 0 {
|
||||
if mby == 0 {
|
||||
return predDCTopLeft
|
||||
}
|
||||
return predDCLeft
|
||||
}
|
||||
if mby == 0 {
|
||||
return predDCTop
|
||||
}
|
||||
return predDC
|
||||
}
|
||||
|
||||
var predFunc4 = [...]func(*Decoder, int, int){
|
||||
predFunc4DC,
|
||||
predFunc4TM,
|
||||
predFunc4VE,
|
||||
predFunc4HE,
|
||||
predFunc4RD,
|
||||
predFunc4VR,
|
||||
predFunc4LD,
|
||||
predFunc4VL,
|
||||
predFunc4HD,
|
||||
predFunc4HU,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
}
|
||||
|
||||
var predFunc8 = [...]func(*Decoder, int, int){
|
||||
predFunc8DC,
|
||||
predFunc8TM,
|
||||
predFunc8VE,
|
||||
predFunc8HE,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
predFunc8DCTop,
|
||||
predFunc8DCLeft,
|
||||
predFunc8DCTopLeft,
|
||||
}
|
||||
|
||||
var predFunc16 = [...]func(*Decoder, int, int){
|
||||
predFunc16DC,
|
||||
predFunc16TM,
|
||||
predFunc16VE,
|
||||
predFunc16HE,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
predFunc16DCTop,
|
||||
predFunc16DCLeft,
|
||||
predFunc16DCTopLeft,
|
||||
}
|
||||
|
||||
func predFunc4DC(z *Decoder, y, x int) {
|
||||
sum := uint32(4)
|
||||
for i := 0; i < 4; i++ {
|
||||
sum += uint32(z.ybr[y-1][x+i])
|
||||
}
|
||||
for j := 0; j < 4; j++ {
|
||||
sum += uint32(z.ybr[y+j][x-1])
|
||||
}
|
||||
avg := uint8(sum / 8)
|
||||
for j := 0; j < 4; j++ {
|
||||
for i := 0; i < 4; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc4TM(z *Decoder, y, x int) {
|
||||
delta0 := -int32(z.ybr[y-1][x-1])
|
||||
for j := 0; j < 4; j++ {
|
||||
delta1 := delta0 + int32(z.ybr[y+j][x-1])
|
||||
for i := 0; i < 4; i++ {
|
||||
delta2 := delta1 + int32(z.ybr[y-1][x+i])
|
||||
z.ybr[y+j][x+i] = uint8(clip(delta2, 0, 255))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc4VE(z *Decoder, y, x int) {
|
||||
a := int32(z.ybr[y-1][x-1])
|
||||
b := int32(z.ybr[y-1][x+0])
|
||||
c := int32(z.ybr[y-1][x+1])
|
||||
d := int32(z.ybr[y-1][x+2])
|
||||
e := int32(z.ybr[y-1][x+3])
|
||||
f := int32(z.ybr[y-1][x+4])
|
||||
abc := uint8((a + 2*b + c + 2) / 4)
|
||||
bcd := uint8((b + 2*c + d + 2) / 4)
|
||||
cde := uint8((c + 2*d + e + 2) / 4)
|
||||
def := uint8((d + 2*e + f + 2) / 4)
|
||||
for j := 0; j < 4; j++ {
|
||||
z.ybr[y+j][x+0] = abc
|
||||
z.ybr[y+j][x+1] = bcd
|
||||
z.ybr[y+j][x+2] = cde
|
||||
z.ybr[y+j][x+3] = def
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc4HE(z *Decoder, y, x int) {
|
||||
s := int32(z.ybr[y+3][x-1])
|
||||
r := int32(z.ybr[y+2][x-1])
|
||||
q := int32(z.ybr[y+1][x-1])
|
||||
p := int32(z.ybr[y+0][x-1])
|
||||
a := int32(z.ybr[y-1][x-1])
|
||||
ssr := uint8((s + 2*s + r + 2) / 4)
|
||||
srq := uint8((s + 2*r + q + 2) / 4)
|
||||
rqp := uint8((r + 2*q + p + 2) / 4)
|
||||
apq := uint8((a + 2*p + q + 2) / 4)
|
||||
for i := 0; i < 4; i++ {
|
||||
z.ybr[y+0][x+i] = apq
|
||||
z.ybr[y+1][x+i] = rqp
|
||||
z.ybr[y+2][x+i] = srq
|
||||
z.ybr[y+3][x+i] = ssr
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc4RD(z *Decoder, y, x int) {
|
||||
s := int32(z.ybr[y+3][x-1])
|
||||
r := int32(z.ybr[y+2][x-1])
|
||||
q := int32(z.ybr[y+1][x-1])
|
||||
p := int32(z.ybr[y+0][x-1])
|
||||
a := int32(z.ybr[y-1][x-1])
|
||||
b := int32(z.ybr[y-1][x+0])
|
||||
c := int32(z.ybr[y-1][x+1])
|
||||
d := int32(z.ybr[y-1][x+2])
|
||||
e := int32(z.ybr[y-1][x+3])
|
||||
srq := uint8((s + 2*r + q + 2) / 4)
|
||||
rqp := uint8((r + 2*q + p + 2) / 4)
|
||||
qpa := uint8((q + 2*p + a + 2) / 4)
|
||||
pab := uint8((p + 2*a + b + 2) / 4)
|
||||
abc := uint8((a + 2*b + c + 2) / 4)
|
||||
bcd := uint8((b + 2*c + d + 2) / 4)
|
||||
cde := uint8((c + 2*d + e + 2) / 4)
|
||||
z.ybr[y+0][x+0] = pab
|
||||
z.ybr[y+0][x+1] = abc
|
||||
z.ybr[y+0][x+2] = bcd
|
||||
z.ybr[y+0][x+3] = cde
|
||||
z.ybr[y+1][x+0] = qpa
|
||||
z.ybr[y+1][x+1] = pab
|
||||
z.ybr[y+1][x+2] = abc
|
||||
z.ybr[y+1][x+3] = bcd
|
||||
z.ybr[y+2][x+0] = rqp
|
||||
z.ybr[y+2][x+1] = qpa
|
||||
z.ybr[y+2][x+2] = pab
|
||||
z.ybr[y+2][x+3] = abc
|
||||
z.ybr[y+3][x+0] = srq
|
||||
z.ybr[y+3][x+1] = rqp
|
||||
z.ybr[y+3][x+2] = qpa
|
||||
z.ybr[y+3][x+3] = pab
|
||||
}
|
||||
|
||||
func predFunc4VR(z *Decoder, y, x int) {
|
||||
r := int32(z.ybr[y+2][x-1])
|
||||
q := int32(z.ybr[y+1][x-1])
|
||||
p := int32(z.ybr[y+0][x-1])
|
||||
a := int32(z.ybr[y-1][x-1])
|
||||
b := int32(z.ybr[y-1][x+0])
|
||||
c := int32(z.ybr[y-1][x+1])
|
||||
d := int32(z.ybr[y-1][x+2])
|
||||
e := int32(z.ybr[y-1][x+3])
|
||||
ab := uint8((a + b + 1) / 2)
|
||||
bc := uint8((b + c + 1) / 2)
|
||||
cd := uint8((c + d + 1) / 2)
|
||||
de := uint8((d + e + 1) / 2)
|
||||
rqp := uint8((r + 2*q + p + 2) / 4)
|
||||
qpa := uint8((q + 2*p + a + 2) / 4)
|
||||
pab := uint8((p + 2*a + b + 2) / 4)
|
||||
abc := uint8((a + 2*b + c + 2) / 4)
|
||||
bcd := uint8((b + 2*c + d + 2) / 4)
|
||||
cde := uint8((c + 2*d + e + 2) / 4)
|
||||
z.ybr[y+0][x+0] = ab
|
||||
z.ybr[y+0][x+1] = bc
|
||||
z.ybr[y+0][x+2] = cd
|
||||
z.ybr[y+0][x+3] = de
|
||||
z.ybr[y+1][x+0] = pab
|
||||
z.ybr[y+1][x+1] = abc
|
||||
z.ybr[y+1][x+2] = bcd
|
||||
z.ybr[y+1][x+3] = cde
|
||||
z.ybr[y+2][x+0] = qpa
|
||||
z.ybr[y+2][x+1] = ab
|
||||
z.ybr[y+2][x+2] = bc
|
||||
z.ybr[y+2][x+3] = cd
|
||||
z.ybr[y+3][x+0] = rqp
|
||||
z.ybr[y+3][x+1] = pab
|
||||
z.ybr[y+3][x+2] = abc
|
||||
z.ybr[y+3][x+3] = bcd
|
||||
}
|
||||
|
||||
func predFunc4LD(z *Decoder, y, x int) {
|
||||
a := int32(z.ybr[y-1][x+0])
|
||||
b := int32(z.ybr[y-1][x+1])
|
||||
c := int32(z.ybr[y-1][x+2])
|
||||
d := int32(z.ybr[y-1][x+3])
|
||||
e := int32(z.ybr[y-1][x+4])
|
||||
f := int32(z.ybr[y-1][x+5])
|
||||
g := int32(z.ybr[y-1][x+6])
|
||||
h := int32(z.ybr[y-1][x+7])
|
||||
abc := uint8((a + 2*b + c + 2) / 4)
|
||||
bcd := uint8((b + 2*c + d + 2) / 4)
|
||||
cde := uint8((c + 2*d + e + 2) / 4)
|
||||
def := uint8((d + 2*e + f + 2) / 4)
|
||||
efg := uint8((e + 2*f + g + 2) / 4)
|
||||
fgh := uint8((f + 2*g + h + 2) / 4)
|
||||
ghh := uint8((g + 2*h + h + 2) / 4)
|
||||
z.ybr[y+0][x+0] = abc
|
||||
z.ybr[y+0][x+1] = bcd
|
||||
z.ybr[y+0][x+2] = cde
|
||||
z.ybr[y+0][x+3] = def
|
||||
z.ybr[y+1][x+0] = bcd
|
||||
z.ybr[y+1][x+1] = cde
|
||||
z.ybr[y+1][x+2] = def
|
||||
z.ybr[y+1][x+3] = efg
|
||||
z.ybr[y+2][x+0] = cde
|
||||
z.ybr[y+2][x+1] = def
|
||||
z.ybr[y+2][x+2] = efg
|
||||
z.ybr[y+2][x+3] = fgh
|
||||
z.ybr[y+3][x+0] = def
|
||||
z.ybr[y+3][x+1] = efg
|
||||
z.ybr[y+3][x+2] = fgh
|
||||
z.ybr[y+3][x+3] = ghh
|
||||
}
|
||||
|
||||
func predFunc4VL(z *Decoder, y, x int) {
|
||||
a := int32(z.ybr[y-1][x+0])
|
||||
b := int32(z.ybr[y-1][x+1])
|
||||
c := int32(z.ybr[y-1][x+2])
|
||||
d := int32(z.ybr[y-1][x+3])
|
||||
e := int32(z.ybr[y-1][x+4])
|
||||
f := int32(z.ybr[y-1][x+5])
|
||||
g := int32(z.ybr[y-1][x+6])
|
||||
h := int32(z.ybr[y-1][x+7])
|
||||
ab := uint8((a + b + 1) / 2)
|
||||
bc := uint8((b + c + 1) / 2)
|
||||
cd := uint8((c + d + 1) / 2)
|
||||
de := uint8((d + e + 1) / 2)
|
||||
abc := uint8((a + 2*b + c + 2) / 4)
|
||||
bcd := uint8((b + 2*c + d + 2) / 4)
|
||||
cde := uint8((c + 2*d + e + 2) / 4)
|
||||
def := uint8((d + 2*e + f + 2) / 4)
|
||||
efg := uint8((e + 2*f + g + 2) / 4)
|
||||
fgh := uint8((f + 2*g + h + 2) / 4)
|
||||
z.ybr[y+0][x+0] = ab
|
||||
z.ybr[y+0][x+1] = bc
|
||||
z.ybr[y+0][x+2] = cd
|
||||
z.ybr[y+0][x+3] = de
|
||||
z.ybr[y+1][x+0] = abc
|
||||
z.ybr[y+1][x+1] = bcd
|
||||
z.ybr[y+1][x+2] = cde
|
||||
z.ybr[y+1][x+3] = def
|
||||
z.ybr[y+2][x+0] = bc
|
||||
z.ybr[y+2][x+1] = cd
|
||||
z.ybr[y+2][x+2] = de
|
||||
z.ybr[y+2][x+3] = efg
|
||||
z.ybr[y+3][x+0] = bcd
|
||||
z.ybr[y+3][x+1] = cde
|
||||
z.ybr[y+3][x+2] = def
|
||||
z.ybr[y+3][x+3] = fgh
|
||||
}
|
||||
|
||||
func predFunc4HD(z *Decoder, y, x int) {
|
||||
s := int32(z.ybr[y+3][x-1])
|
||||
r := int32(z.ybr[y+2][x-1])
|
||||
q := int32(z.ybr[y+1][x-1])
|
||||
p := int32(z.ybr[y+0][x-1])
|
||||
a := int32(z.ybr[y-1][x-1])
|
||||
b := int32(z.ybr[y-1][x+0])
|
||||
c := int32(z.ybr[y-1][x+1])
|
||||
d := int32(z.ybr[y-1][x+2])
|
||||
sr := uint8((s + r + 1) / 2)
|
||||
rq := uint8((r + q + 1) / 2)
|
||||
qp := uint8((q + p + 1) / 2)
|
||||
pa := uint8((p + a + 1) / 2)
|
||||
srq := uint8((s + 2*r + q + 2) / 4)
|
||||
rqp := uint8((r + 2*q + p + 2) / 4)
|
||||
qpa := uint8((q + 2*p + a + 2) / 4)
|
||||
pab := uint8((p + 2*a + b + 2) / 4)
|
||||
abc := uint8((a + 2*b + c + 2) / 4)
|
||||
bcd := uint8((b + 2*c + d + 2) / 4)
|
||||
z.ybr[y+0][x+0] = pa
|
||||
z.ybr[y+0][x+1] = pab
|
||||
z.ybr[y+0][x+2] = abc
|
||||
z.ybr[y+0][x+3] = bcd
|
||||
z.ybr[y+1][x+0] = qp
|
||||
z.ybr[y+1][x+1] = qpa
|
||||
z.ybr[y+1][x+2] = pa
|
||||
z.ybr[y+1][x+3] = pab
|
||||
z.ybr[y+2][x+0] = rq
|
||||
z.ybr[y+2][x+1] = rqp
|
||||
z.ybr[y+2][x+2] = qp
|
||||
z.ybr[y+2][x+3] = qpa
|
||||
z.ybr[y+3][x+0] = sr
|
||||
z.ybr[y+3][x+1] = srq
|
||||
z.ybr[y+3][x+2] = rq
|
||||
z.ybr[y+3][x+3] = rqp
|
||||
}
|
||||
|
||||
func predFunc4HU(z *Decoder, y, x int) {
|
||||
s := int32(z.ybr[y+3][x-1])
|
||||
r := int32(z.ybr[y+2][x-1])
|
||||
q := int32(z.ybr[y+1][x-1])
|
||||
p := int32(z.ybr[y+0][x-1])
|
||||
pq := uint8((p + q + 1) / 2)
|
||||
qr := uint8((q + r + 1) / 2)
|
||||
rs := uint8((r + s + 1) / 2)
|
||||
pqr := uint8((p + 2*q + r + 2) / 4)
|
||||
qrs := uint8((q + 2*r + s + 2) / 4)
|
||||
rss := uint8((r + 2*s + s + 2) / 4)
|
||||
sss := uint8(s)
|
||||
z.ybr[y+0][x+0] = pq
|
||||
z.ybr[y+0][x+1] = pqr
|
||||
z.ybr[y+0][x+2] = qr
|
||||
z.ybr[y+0][x+3] = qrs
|
||||
z.ybr[y+1][x+0] = qr
|
||||
z.ybr[y+1][x+1] = qrs
|
||||
z.ybr[y+1][x+2] = rs
|
||||
z.ybr[y+1][x+3] = rss
|
||||
z.ybr[y+2][x+0] = rs
|
||||
z.ybr[y+2][x+1] = rss
|
||||
z.ybr[y+2][x+2] = sss
|
||||
z.ybr[y+2][x+3] = sss
|
||||
z.ybr[y+3][x+0] = sss
|
||||
z.ybr[y+3][x+1] = sss
|
||||
z.ybr[y+3][x+2] = sss
|
||||
z.ybr[y+3][x+3] = sss
|
||||
}
|
||||
|
||||
func predFunc8DC(z *Decoder, y, x int) {
|
||||
sum := uint32(8)
|
||||
for i := 0; i < 8; i++ {
|
||||
sum += uint32(z.ybr[y-1][x+i])
|
||||
}
|
||||
for j := 0; j < 8; j++ {
|
||||
sum += uint32(z.ybr[y+j][x-1])
|
||||
}
|
||||
avg := uint8(sum / 16)
|
||||
for j := 0; j < 8; j++ {
|
||||
for i := 0; i < 8; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc8TM(z *Decoder, y, x int) {
|
||||
delta0 := -int32(z.ybr[y-1][x-1])
|
||||
for j := 0; j < 8; j++ {
|
||||
delta1 := delta0 + int32(z.ybr[y+j][x-1])
|
||||
for i := 0; i < 8; i++ {
|
||||
delta2 := delta1 + int32(z.ybr[y-1][x+i])
|
||||
z.ybr[y+j][x+i] = uint8(clip(delta2, 0, 255))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc8VE(z *Decoder, y, x int) {
|
||||
for j := 0; j < 8; j++ {
|
||||
for i := 0; i < 8; i++ {
|
||||
z.ybr[y+j][x+i] = z.ybr[y-1][x+i]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc8HE(z *Decoder, y, x int) {
|
||||
for j := 0; j < 8; j++ {
|
||||
for i := 0; i < 8; i++ {
|
||||
z.ybr[y+j][x+i] = z.ybr[y+j][x-1]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc8DCTop(z *Decoder, y, x int) {
|
||||
sum := uint32(4)
|
||||
for j := 0; j < 8; j++ {
|
||||
sum += uint32(z.ybr[y+j][x-1])
|
||||
}
|
||||
avg := uint8(sum / 8)
|
||||
for j := 0; j < 8; j++ {
|
||||
for i := 0; i < 8; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc8DCLeft(z *Decoder, y, x int) {
|
||||
sum := uint32(4)
|
||||
for i := 0; i < 8; i++ {
|
||||
sum += uint32(z.ybr[y-1][x+i])
|
||||
}
|
||||
avg := uint8(sum / 8)
|
||||
for j := 0; j < 8; j++ {
|
||||
for i := 0; i < 8; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc8DCTopLeft(z *Decoder, y, x int) {
|
||||
for j := 0; j < 8; j++ {
|
||||
for i := 0; i < 8; i++ {
|
||||
z.ybr[y+j][x+i] = 0x80
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16DC(z *Decoder, y, x int) {
|
||||
sum := uint32(16)
|
||||
for i := 0; i < 16; i++ {
|
||||
sum += uint32(z.ybr[y-1][x+i])
|
||||
}
|
||||
for j := 0; j < 16; j++ {
|
||||
sum += uint32(z.ybr[y+j][x-1])
|
||||
}
|
||||
avg := uint8(sum / 32)
|
||||
for j := 0; j < 16; j++ {
|
||||
for i := 0; i < 16; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16TM(z *Decoder, y, x int) {
|
||||
delta0 := -int32(z.ybr[y-1][x-1])
|
||||
for j := 0; j < 16; j++ {
|
||||
delta1 := delta0 + int32(z.ybr[y+j][x-1])
|
||||
for i := 0; i < 16; i++ {
|
||||
delta2 := delta1 + int32(z.ybr[y-1][x+i])
|
||||
z.ybr[y+j][x+i] = uint8(clip(delta2, 0, 255))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16VE(z *Decoder, y, x int) {
|
||||
for j := 0; j < 16; j++ {
|
||||
for i := 0; i < 16; i++ {
|
||||
z.ybr[y+j][x+i] = z.ybr[y-1][x+i]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16HE(z *Decoder, y, x int) {
|
||||
for j := 0; j < 16; j++ {
|
||||
for i := 0; i < 16; i++ {
|
||||
z.ybr[y+j][x+i] = z.ybr[y+j][x-1]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16DCTop(z *Decoder, y, x int) {
|
||||
sum := uint32(8)
|
||||
for j := 0; j < 16; j++ {
|
||||
sum += uint32(z.ybr[y+j][x-1])
|
||||
}
|
||||
avg := uint8(sum / 16)
|
||||
for j := 0; j < 16; j++ {
|
||||
for i := 0; i < 16; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16DCLeft(z *Decoder, y, x int) {
|
||||
sum := uint32(8)
|
||||
for i := 0; i < 16; i++ {
|
||||
sum += uint32(z.ybr[y-1][x+i])
|
||||
}
|
||||
avg := uint8(sum / 16)
|
||||
for j := 0; j < 16; j++ {
|
||||
for i := 0; i < 16; i++ {
|
||||
z.ybr[y+j][x+i] = avg
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func predFunc16DCTopLeft(z *Decoder, y, x int) {
|
||||
for j := 0; j < 16; j++ {
|
||||
for i := 0; i < 16; i++ {
|
||||
z.ybr[y+j][x+i] = 0x80
|
||||
}
|
||||
}
|
||||
}
|
||||
+98
@@ -0,0 +1,98 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// This file implements parsing the quantization factors.
|
||||
|
||||
// quant are DC/AC quantization factors.
|
||||
type quant struct {
|
||||
y1 [2]uint16
|
||||
y2 [2]uint16
|
||||
uv [2]uint16
|
||||
}
|
||||
|
||||
// clip clips x to the range [min, max] inclusive.
|
||||
func clip(x, min, max int32) int32 {
|
||||
if x < min {
|
||||
return min
|
||||
}
|
||||
if x > max {
|
||||
return max
|
||||
}
|
||||
return x
|
||||
}
|
||||
|
||||
// parseQuant parses the quantization factors, as specified in section 9.6.
|
||||
func (d *Decoder) parseQuant() {
|
||||
baseQ0 := d.fp.readUint(uniformProb, 7)
|
||||
dqy1DC := d.fp.readOptionalInt(uniformProb, 4)
|
||||
const dqy1AC = 0
|
||||
dqy2DC := d.fp.readOptionalInt(uniformProb, 4)
|
||||
dqy2AC := d.fp.readOptionalInt(uniformProb, 4)
|
||||
dquvDC := d.fp.readOptionalInt(uniformProb, 4)
|
||||
dquvAC := d.fp.readOptionalInt(uniformProb, 4)
|
||||
for i := 0; i < nSegment; i++ {
|
||||
q := int32(baseQ0)
|
||||
if d.segmentHeader.useSegment {
|
||||
if d.segmentHeader.relativeDelta {
|
||||
q += int32(d.segmentHeader.quantizer[i])
|
||||
} else {
|
||||
q = int32(d.segmentHeader.quantizer[i])
|
||||
}
|
||||
}
|
||||
d.quant[i].y1[0] = dequantTableDC[clip(q+dqy1DC, 0, 127)]
|
||||
d.quant[i].y1[1] = dequantTableAC[clip(q+dqy1AC, 0, 127)]
|
||||
d.quant[i].y2[0] = dequantTableDC[clip(q+dqy2DC, 0, 127)] * 2
|
||||
d.quant[i].y2[1] = dequantTableAC[clip(q+dqy2AC, 0, 127)] * 155 / 100
|
||||
if d.quant[i].y2[1] < 8 {
|
||||
d.quant[i].y2[1] = 8
|
||||
}
|
||||
// The 117 is not a typo. The dequant_init function in the spec's Reference
|
||||
// Decoder Source Code (http://tools.ietf.org/html/rfc6386#section-9.6 Page 145)
|
||||
// says to clamp the LHS value at 132, which is equal to dequantTableDC[117].
|
||||
d.quant[i].uv[0] = dequantTableDC[clip(q+dquvDC, 0, 117)]
|
||||
d.quant[i].uv[1] = dequantTableAC[clip(q+dquvAC, 0, 127)]
|
||||
}
|
||||
}
|
||||
|
||||
// The dequantization tables are specified in section 14.1.
|
||||
var (
|
||||
dequantTableDC = [128]uint16{
|
||||
4, 5, 6, 7, 8, 9, 10, 10,
|
||||
11, 12, 13, 14, 15, 16, 17, 17,
|
||||
18, 19, 20, 20, 21, 21, 22, 22,
|
||||
23, 23, 24, 25, 25, 26, 27, 28,
|
||||
29, 30, 31, 32, 33, 34, 35, 36,
|
||||
37, 37, 38, 39, 40, 41, 42, 43,
|
||||
44, 45, 46, 46, 47, 48, 49, 50,
|
||||
51, 52, 53, 54, 55, 56, 57, 58,
|
||||
59, 60, 61, 62, 63, 64, 65, 66,
|
||||
67, 68, 69, 70, 71, 72, 73, 74,
|
||||
75, 76, 76, 77, 78, 79, 80, 81,
|
||||
82, 83, 84, 85, 86, 87, 88, 89,
|
||||
91, 93, 95, 96, 98, 100, 101, 102,
|
||||
104, 106, 108, 110, 112, 114, 116, 118,
|
||||
122, 124, 126, 128, 130, 132, 134, 136,
|
||||
138, 140, 143, 145, 148, 151, 154, 157,
|
||||
}
|
||||
dequantTableAC = [128]uint16{
|
||||
4, 5, 6, 7, 8, 9, 10, 11,
|
||||
12, 13, 14, 15, 16, 17, 18, 19,
|
||||
20, 21, 22, 23, 24, 25, 26, 27,
|
||||
28, 29, 30, 31, 32, 33, 34, 35,
|
||||
36, 37, 38, 39, 40, 41, 42, 43,
|
||||
44, 45, 46, 47, 48, 49, 50, 51,
|
||||
52, 53, 54, 55, 56, 57, 58, 60,
|
||||
62, 64, 66, 68, 70, 72, 74, 76,
|
||||
78, 80, 82, 84, 86, 88, 90, 92,
|
||||
94, 96, 98, 100, 102, 104, 106, 108,
|
||||
110, 112, 114, 116, 119, 122, 125, 128,
|
||||
131, 134, 137, 140, 143, 146, 149, 152,
|
||||
155, 158, 161, 164, 167, 170, 173, 177,
|
||||
181, 185, 189, 193, 197, 201, 205, 209,
|
||||
213, 217, 221, 225, 229, 234, 239, 245,
|
||||
249, 254, 259, 264, 269, 274, 279, 284,
|
||||
}
|
||||
)
|
||||
+442
@@ -0,0 +1,442 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// This file implements decoding DCT/WHT residual coefficients and
|
||||
// reconstructing YCbCr data equal to predicted values plus residuals.
|
||||
//
|
||||
// There are 1*16*16 + 2*8*8 + 1*4*4 coefficients per macroblock:
|
||||
// - 1*16*16 luma DCT coefficients,
|
||||
// - 2*8*8 chroma DCT coefficients, and
|
||||
// - 1*4*4 luma WHT coefficients.
|
||||
// Coefficients are read in lots of 16, and the later coefficients in each lot
|
||||
// are often zero.
|
||||
//
|
||||
// The YCbCr data consists of 1*16*16 luma values and 2*8*8 chroma values,
|
||||
// plus previously decoded values along the top and left borders. The combined
|
||||
// values are laid out as a [1+16+1+8][32]uint8 so that vertically adjacent
|
||||
// samples are 32 bytes apart. In detail, the layout is:
|
||||
//
|
||||
// 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1
|
||||
// . . . . . . . a b b b b b b b b b b b b b b b b c c c c . . . . 0
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 1
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 2
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 3
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y c c c c . . . . 4
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 5
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 6
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 7
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y c c c c . . . . 8
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 9
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 10
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 11
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y c c c c . . . . 12
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 13
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 14
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 15
|
||||
// . . . . . . . d Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y Y . . . . . . . . 16
|
||||
// . . . . . . . e f f f f f f f f . . . . . . . g h h h h h h h h 17
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 18
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 19
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 20
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 21
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 22
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 23
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 24
|
||||
// . . . . . . . i B B B B B B B B . . . . . . . j R R R R R R R R 25
|
||||
//
|
||||
// Y, B and R are the reconstructed luma (Y) and chroma (B, R) values.
|
||||
// The Y values are predicted (either as one 16x16 region or 16 4x4 regions)
|
||||
// based on the row above's Y values (some combination of {abc} or {dYC}) and
|
||||
// the column left's Y values (either {ad} or {bY}). Similarly, B and R values
|
||||
// are predicted on the row above and column left of their respective 8x8
|
||||
// region: {efi} for B, {ghj} for R.
|
||||
//
|
||||
// For uppermost macroblocks (i.e. those with mby == 0), the {abcefgh} values
|
||||
// are initialized to 0x81. Otherwise, they are copied from the bottom row of
|
||||
// the macroblock above. The {c} values are then duplicated from row 0 to rows
|
||||
// 4, 8 and 12 of the ybr workspace.
|
||||
// Similarly, for leftmost macroblocks (i.e. those with mbx == 0), the {adeigj}
|
||||
// values are initialized to 0x7f. Otherwise, they are copied from the right
|
||||
// column of the macroblock to the left.
|
||||
// For the top-left macroblock (with mby == 0 && mbx == 0), {aeg} is 0x81.
|
||||
//
|
||||
// When moving from one macroblock to the next horizontally, the {adeigj}
|
||||
// values can simply be copied from the workspace to itself, shifted by 8 or
|
||||
// 16 columns. When moving from one macroblock to the next vertically,
|
||||
// filtering can occur and hence the row values have to be copied from the
|
||||
// post-filtered image instead of the pre-filtered workspace.
|
||||
|
||||
const (
|
||||
bCoeffBase = 1*16*16 + 0*8*8
|
||||
rCoeffBase = 1*16*16 + 1*8*8
|
||||
whtCoeffBase = 1*16*16 + 2*8*8
|
||||
)
|
||||
|
||||
const (
|
||||
ybrYX = 8
|
||||
ybrYY = 1
|
||||
ybrBX = 8
|
||||
ybrBY = 18
|
||||
ybrRX = 24
|
||||
ybrRY = 18
|
||||
)
|
||||
|
||||
// prepareYBR prepares the {abcdefghij} elements of ybr.
|
||||
func (d *Decoder) prepareYBR(mbx, mby int) {
|
||||
if mbx == 0 {
|
||||
for y := 0; y < 17; y++ {
|
||||
d.ybr[y][7] = 0x81
|
||||
}
|
||||
for y := 17; y < 26; y++ {
|
||||
d.ybr[y][7] = 0x81
|
||||
d.ybr[y][23] = 0x81
|
||||
}
|
||||
} else {
|
||||
for y := 0; y < 17; y++ {
|
||||
d.ybr[y][7] = d.ybr[y][7+16]
|
||||
}
|
||||
for y := 17; y < 26; y++ {
|
||||
d.ybr[y][7] = d.ybr[y][15]
|
||||
d.ybr[y][23] = d.ybr[y][31]
|
||||
}
|
||||
}
|
||||
if mby == 0 {
|
||||
for x := 7; x < 28; x++ {
|
||||
d.ybr[0][x] = 0x7f
|
||||
}
|
||||
for x := 7; x < 16; x++ {
|
||||
d.ybr[17][x] = 0x7f
|
||||
}
|
||||
for x := 23; x < 32; x++ {
|
||||
d.ybr[17][x] = 0x7f
|
||||
}
|
||||
} else {
|
||||
for i := 0; i < 16; i++ {
|
||||
d.ybr[0][8+i] = d.img.Y[(16*mby-1)*d.img.YStride+16*mbx+i]
|
||||
}
|
||||
for i := 0; i < 8; i++ {
|
||||
d.ybr[17][8+i] = d.img.Cb[(8*mby-1)*d.img.CStride+8*mbx+i]
|
||||
}
|
||||
for i := 0; i < 8; i++ {
|
||||
d.ybr[17][24+i] = d.img.Cr[(8*mby-1)*d.img.CStride+8*mbx+i]
|
||||
}
|
||||
if mbx == d.mbw-1 {
|
||||
for i := 16; i < 20; i++ {
|
||||
d.ybr[0][8+i] = d.img.Y[(16*mby-1)*d.img.YStride+16*mbx+15]
|
||||
}
|
||||
} else {
|
||||
for i := 16; i < 20; i++ {
|
||||
d.ybr[0][8+i] = d.img.Y[(16*mby-1)*d.img.YStride+16*mbx+i]
|
||||
}
|
||||
}
|
||||
}
|
||||
for y := 4; y < 16; y += 4 {
|
||||
d.ybr[y][24] = d.ybr[0][24]
|
||||
d.ybr[y][25] = d.ybr[0][25]
|
||||
d.ybr[y][26] = d.ybr[0][26]
|
||||
d.ybr[y][27] = d.ybr[0][27]
|
||||
}
|
||||
}
|
||||
|
||||
// btou converts a bool to a 0/1 value.
|
||||
func btou(b bool) uint8 {
|
||||
if b {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// pack packs four 0/1 values into four bits of a uint32.
|
||||
func pack(x [4]uint8, shift int) uint32 {
|
||||
u := uint32(x[0])<<0 | uint32(x[1])<<1 | uint32(x[2])<<2 | uint32(x[3])<<3
|
||||
return u << uint(shift)
|
||||
}
|
||||
|
||||
// unpack unpacks four 0/1 values from a four-bit value.
|
||||
var unpack = [16][4]uint8{
|
||||
{0, 0, 0, 0},
|
||||
{1, 0, 0, 0},
|
||||
{0, 1, 0, 0},
|
||||
{1, 1, 0, 0},
|
||||
{0, 0, 1, 0},
|
||||
{1, 0, 1, 0},
|
||||
{0, 1, 1, 0},
|
||||
{1, 1, 1, 0},
|
||||
{0, 0, 0, 1},
|
||||
{1, 0, 0, 1},
|
||||
{0, 1, 0, 1},
|
||||
{1, 1, 0, 1},
|
||||
{0, 0, 1, 1},
|
||||
{1, 0, 1, 1},
|
||||
{0, 1, 1, 1},
|
||||
{1, 1, 1, 1},
|
||||
}
|
||||
|
||||
var (
|
||||
// The mapping from 4x4 region position to band is specified in section 13.3.
|
||||
bands = [17]uint8{0, 1, 2, 3, 6, 4, 5, 6, 6, 6, 6, 6, 6, 6, 6, 7, 0}
|
||||
// Category probabilities are specified in section 13.2.
|
||||
// Decoding categories 1 and 2 are done inline.
|
||||
cat3456 = [4][12]uint8{
|
||||
{173, 148, 140, 0, 0, 0, 0, 0, 0, 0, 0, 0},
|
||||
{176, 155, 140, 135, 0, 0, 0, 0, 0, 0, 0, 0},
|
||||
{180, 157, 141, 134, 130, 0, 0, 0, 0, 0, 0, 0},
|
||||
{254, 254, 243, 230, 196, 177, 153, 140, 133, 130, 129, 0},
|
||||
}
|
||||
// The zigzag order is:
|
||||
// 0 1 5 6
|
||||
// 2 4 7 12
|
||||
// 3 8 11 13
|
||||
// 9 10 14 15
|
||||
zigzag = [16]uint8{0, 1, 4, 8, 5, 2, 3, 6, 9, 12, 13, 10, 7, 11, 14, 15}
|
||||
)
|
||||
|
||||
// parseResiduals4 parses a 4x4 region of residual coefficients, as specified
|
||||
// in section 13.3, and returns a 0/1 value indicating whether there was at
|
||||
// least one non-zero coefficient.
|
||||
// r is the partition to read bits from.
|
||||
// plane and context describe which token probability table to use. context is
|
||||
// either 0, 1 or 2, and equals how many of the macroblock left and macroblock
|
||||
// above have non-zero coefficients.
|
||||
// quant are the DC/AC quantization factors.
|
||||
// skipFirstCoeff is whether the DC coefficient has already been parsed.
|
||||
// coeffBase is the base index of d.coeff to write to.
|
||||
func (d *Decoder) parseResiduals4(r *partition, plane int, context uint8, quant [2]uint16, skipFirstCoeff bool, coeffBase int) uint8 {
|
||||
prob, n := &d.tokenProb[plane], 0
|
||||
if skipFirstCoeff {
|
||||
n = 1
|
||||
}
|
||||
p := prob[bands[n]][context]
|
||||
if !r.readBit(p[0]) {
|
||||
return 0
|
||||
}
|
||||
for n != 16 {
|
||||
n++
|
||||
if !r.readBit(p[1]) {
|
||||
p = prob[bands[n]][0]
|
||||
continue
|
||||
}
|
||||
var v uint32
|
||||
if !r.readBit(p[2]) {
|
||||
v = 1
|
||||
p = prob[bands[n]][1]
|
||||
} else {
|
||||
if !r.readBit(p[3]) {
|
||||
if !r.readBit(p[4]) {
|
||||
v = 2
|
||||
} else {
|
||||
v = 3 + r.readUint(p[5], 1)
|
||||
}
|
||||
} else if !r.readBit(p[6]) {
|
||||
if !r.readBit(p[7]) {
|
||||
// Category 1.
|
||||
v = 5 + r.readUint(159, 1)
|
||||
} else {
|
||||
// Category 2.
|
||||
v = 7 + 2*r.readUint(165, 1) + r.readUint(145, 1)
|
||||
}
|
||||
} else {
|
||||
// Categories 3, 4, 5 or 6.
|
||||
b1 := r.readUint(p[8], 1)
|
||||
b0 := r.readUint(p[9+b1], 1)
|
||||
cat := 2*b1 + b0
|
||||
tab := &cat3456[cat]
|
||||
v = 0
|
||||
for i := 0; tab[i] != 0; i++ {
|
||||
v *= 2
|
||||
v += r.readUint(tab[i], 1)
|
||||
}
|
||||
v += 3 + (8 << cat)
|
||||
}
|
||||
p = prob[bands[n]][2]
|
||||
}
|
||||
z := zigzag[n-1]
|
||||
c := int32(v) * int32(quant[btou(z > 0)])
|
||||
if r.readBit(uniformProb) {
|
||||
c = -c
|
||||
}
|
||||
d.coeff[coeffBase+int(z)] = int16(c)
|
||||
if n == 16 || !r.readBit(p[0]) {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
// parseResiduals parses the residuals and returns whether inner loop filtering
|
||||
// should be skipped for this macroblock.
|
||||
func (d *Decoder) parseResiduals(mbx, mby int) (skip bool) {
|
||||
partition := &d.op[mby&(d.nOP-1)]
|
||||
plane := planeY1SansY2
|
||||
quant := &d.quant[d.segment]
|
||||
|
||||
// Parse the DC coefficient of each 4x4 luma region.
|
||||
if d.usePredY16 {
|
||||
nz := d.parseResiduals4(partition, planeY2, d.leftMB.nzY16+d.upMB[mbx].nzY16, quant.y2, false, whtCoeffBase)
|
||||
d.leftMB.nzY16 = nz
|
||||
d.upMB[mbx].nzY16 = nz
|
||||
d.inverseWHT16()
|
||||
plane = planeY1WithY2
|
||||
}
|
||||
|
||||
var (
|
||||
nzDC, nzAC [4]uint8
|
||||
nzDCMask, nzACMask uint32
|
||||
coeffBase int
|
||||
)
|
||||
|
||||
// Parse the luma coefficients.
|
||||
lnz := unpack[d.leftMB.nzMask&0x0f]
|
||||
unz := unpack[d.upMB[mbx].nzMask&0x0f]
|
||||
for y := 0; y < 4; y++ {
|
||||
nz := lnz[y]
|
||||
for x := 0; x < 4; x++ {
|
||||
nz = d.parseResiduals4(partition, plane, nz+unz[x], quant.y1, d.usePredY16, coeffBase)
|
||||
unz[x] = nz
|
||||
nzAC[x] = nz
|
||||
nzDC[x] = btou(d.coeff[coeffBase] != 0)
|
||||
coeffBase += 16
|
||||
}
|
||||
lnz[y] = nz
|
||||
nzDCMask |= pack(nzDC, y*4)
|
||||
nzACMask |= pack(nzAC, y*4)
|
||||
}
|
||||
lnzMask := pack(lnz, 0)
|
||||
unzMask := pack(unz, 0)
|
||||
|
||||
// Parse the chroma coefficients.
|
||||
lnz = unpack[d.leftMB.nzMask>>4]
|
||||
unz = unpack[d.upMB[mbx].nzMask>>4]
|
||||
for c := 0; c < 4; c += 2 {
|
||||
for y := 0; y < 2; y++ {
|
||||
nz := lnz[y+c]
|
||||
for x := 0; x < 2; x++ {
|
||||
nz = d.parseResiduals4(partition, planeUV, nz+unz[x+c], quant.uv, false, coeffBase)
|
||||
unz[x+c] = nz
|
||||
nzAC[y*2+x] = nz
|
||||
nzDC[y*2+x] = btou(d.coeff[coeffBase] != 0)
|
||||
coeffBase += 16
|
||||
}
|
||||
lnz[y+c] = nz
|
||||
}
|
||||
nzDCMask |= pack(nzDC, 16+c*2)
|
||||
nzACMask |= pack(nzAC, 16+c*2)
|
||||
}
|
||||
lnzMask |= pack(lnz, 4)
|
||||
unzMask |= pack(unz, 4)
|
||||
|
||||
// Save decoder state.
|
||||
d.leftMB.nzMask = uint8(lnzMask)
|
||||
d.upMB[mbx].nzMask = uint8(unzMask)
|
||||
d.nzDCMask = nzDCMask
|
||||
d.nzACMask = nzACMask
|
||||
|
||||
// Section 15.1 of the spec says that "Steps 2 and 4 [of the loop filter]
|
||||
// are skipped... [if] there is no DCT coefficient coded for the whole
|
||||
// macroblock."
|
||||
return nzDCMask == 0 && nzACMask == 0
|
||||
}
|
||||
|
||||
// reconstructMacroblock applies the predictor functions and adds the inverse-
|
||||
// DCT transformed residuals to recover the YCbCr data.
|
||||
func (d *Decoder) reconstructMacroblock(mbx, mby int) {
|
||||
if d.usePredY16 {
|
||||
p := checkTopLeftPred(mbx, mby, d.predY16)
|
||||
predFunc16[p](d, 1, 8)
|
||||
for j := 0; j < 4; j++ {
|
||||
for i := 0; i < 4; i++ {
|
||||
n := 4*j + i
|
||||
y := 4*j + 1
|
||||
x := 4*i + 8
|
||||
mask := uint32(1) << uint(n)
|
||||
if d.nzACMask&mask != 0 {
|
||||
d.inverseDCT4(y, x, 16*n)
|
||||
} else if d.nzDCMask&mask != 0 {
|
||||
d.inverseDCT4DCOnly(y, x, 16*n)
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for j := 0; j < 4; j++ {
|
||||
for i := 0; i < 4; i++ {
|
||||
n := 4*j + i
|
||||
y := 4*j + 1
|
||||
x := 4*i + 8
|
||||
predFunc4[d.predY4[j][i]](d, y, x)
|
||||
mask := uint32(1) << uint(n)
|
||||
if d.nzACMask&mask != 0 {
|
||||
d.inverseDCT4(y, x, 16*n)
|
||||
} else if d.nzDCMask&mask != 0 {
|
||||
d.inverseDCT4DCOnly(y, x, 16*n)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p := checkTopLeftPred(mbx, mby, d.predC8)
|
||||
predFunc8[p](d, ybrBY, ybrBX)
|
||||
if d.nzACMask&0x0f0000 != 0 {
|
||||
d.inverseDCT8(ybrBY, ybrBX, bCoeffBase)
|
||||
} else if d.nzDCMask&0x0f0000 != 0 {
|
||||
d.inverseDCT8DCOnly(ybrBY, ybrBX, bCoeffBase)
|
||||
}
|
||||
predFunc8[p](d, ybrRY, ybrRX)
|
||||
if d.nzACMask&0xf00000 != 0 {
|
||||
d.inverseDCT8(ybrRY, ybrRX, rCoeffBase)
|
||||
} else if d.nzDCMask&0xf00000 != 0 {
|
||||
d.inverseDCT8DCOnly(ybrRY, ybrRX, rCoeffBase)
|
||||
}
|
||||
}
|
||||
|
||||
// reconstruct reconstructs one macroblock and returns whether inner loop
|
||||
// filtering should be skipped for it.
|
||||
func (d *Decoder) reconstruct(mbx, mby int) (skip bool) {
|
||||
if d.segmentHeader.updateMap {
|
||||
if !d.fp.readBit(d.segmentHeader.prob[0]) {
|
||||
d.segment = int(d.fp.readUint(d.segmentHeader.prob[1], 1))
|
||||
} else {
|
||||
d.segment = int(d.fp.readUint(d.segmentHeader.prob[2], 1)) + 2
|
||||
}
|
||||
}
|
||||
if d.useSkipProb {
|
||||
skip = d.fp.readBit(d.skipProb)
|
||||
}
|
||||
// Prepare the workspace.
|
||||
for i := range d.coeff {
|
||||
d.coeff[i] = 0
|
||||
}
|
||||
d.prepareYBR(mbx, mby)
|
||||
// Parse the predictor modes.
|
||||
d.usePredY16 = d.fp.readBit(145)
|
||||
if d.usePredY16 {
|
||||
d.parsePredModeY16(mbx)
|
||||
} else {
|
||||
d.parsePredModeY4(mbx)
|
||||
}
|
||||
d.parsePredModeC8()
|
||||
// Parse the residuals.
|
||||
if !skip {
|
||||
skip = d.parseResiduals(mbx, mby)
|
||||
} else {
|
||||
if d.usePredY16 {
|
||||
d.leftMB.nzY16 = 0
|
||||
d.upMB[mbx].nzY16 = 0
|
||||
}
|
||||
d.leftMB.nzMask = 0
|
||||
d.upMB[mbx].nzMask = 0
|
||||
d.nzDCMask = 0
|
||||
d.nzACMask = 0
|
||||
}
|
||||
// Reconstruct the YCbCr data and copy it to the image.
|
||||
d.reconstructMacroblock(mbx, mby)
|
||||
for i, y := (mby*d.img.YStride+mbx)*16, 0; y < 16; i, y = i+d.img.YStride, y+1 {
|
||||
copy(d.img.Y[i:i+16], d.ybr[ybrYY+y][ybrYX:ybrYX+16])
|
||||
}
|
||||
for i, y := (mby*d.img.CStride+mbx)*8, 0; y < 8; i, y = i+d.img.CStride, y+1 {
|
||||
copy(d.img.Cb[i:i+8], d.ybr[ybrBY+y][ybrBX:ybrBX+8])
|
||||
copy(d.img.Cr[i:i+8], d.ybr[ybrRY+y][ybrRX:ybrRX+8])
|
||||
}
|
||||
return skip
|
||||
}
|
||||
+381
@@ -0,0 +1,381 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8
|
||||
|
||||
// This file contains token probabilities for decoding DCT/WHT coefficients, as
|
||||
// specified in chapter 13.
|
||||
|
||||
func (d *Decoder) parseTokenProb() {
|
||||
for i := range d.tokenProb {
|
||||
for j := range d.tokenProb[i] {
|
||||
for k := range d.tokenProb[i][j] {
|
||||
for l := range d.tokenProb[i][j][k] {
|
||||
if d.fp.readBit(tokenProbUpdateProb[i][j][k][l]) {
|
||||
d.tokenProb[i][j][k][l] = uint8(d.fp.readUint(uniformProb, 8))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The plane enumeration is specified in section 13.3.
|
||||
const (
|
||||
planeY1WithY2 = iota
|
||||
planeY2
|
||||
planeUV
|
||||
planeY1SansY2
|
||||
nPlane
|
||||
)
|
||||
|
||||
const (
|
||||
nBand = 8
|
||||
nContext = 3
|
||||
nProb = 11
|
||||
)
|
||||
|
||||
// Token probability update probabilities are specified in section 13.4.
|
||||
var tokenProbUpdateProb = [nPlane][nBand][nContext][nProb]uint8{
|
||||
{
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{176, 246, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{223, 241, 252, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{249, 253, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 244, 252, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{234, 254, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{253, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 246, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{239, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 248, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{251, 255, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{251, 254, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 253, 255, 254, 255, 255, 255, 255, 255, 255},
|
||||
{250, 255, 254, 255, 254, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
},
|
||||
{
|
||||
{
|
||||
{217, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{225, 252, 241, 253, 255, 255, 254, 255, 255, 255, 255},
|
||||
{234, 250, 241, 250, 253, 255, 253, 254, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{223, 254, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{238, 253, 254, 254, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 248, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{249, 254, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 253, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{247, 254, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{252, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{253, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{250, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
},
|
||||
{
|
||||
{
|
||||
{186, 251, 250, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{234, 251, 244, 254, 255, 255, 255, 255, 255, 255, 255},
|
||||
{251, 251, 243, 253, 254, 255, 254, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{236, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{251, 253, 253, 254, 254, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 254, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 254, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
},
|
||||
{
|
||||
{
|
||||
{248, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{250, 254, 252, 254, 255, 255, 255, 255, 255, 255, 255},
|
||||
{248, 254, 249, 253, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 253, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{246, 253, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{252, 254, 251, 254, 254, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 254, 252, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{248, 254, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{253, 255, 254, 254, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 251, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{245, 251, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{253, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 251, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{252, 253, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 254, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 252, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{249, 255, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 254, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 253, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{250, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
{
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{254, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
{255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Default token probabilities are specified in section 13.5.
|
||||
var defaultTokenProb = [nPlane][nBand][nContext][nProb]uint8{
|
||||
{
|
||||
{
|
||||
{128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{253, 136, 254, 255, 228, 219, 128, 128, 128, 128, 128},
|
||||
{189, 129, 242, 255, 227, 213, 255, 219, 128, 128, 128},
|
||||
{106, 126, 227, 252, 214, 209, 255, 255, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 98, 248, 255, 236, 226, 255, 255, 128, 128, 128},
|
||||
{181, 133, 238, 254, 221, 234, 255, 154, 128, 128, 128},
|
||||
{78, 134, 202, 247, 198, 180, 255, 219, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 185, 249, 255, 243, 255, 128, 128, 128, 128, 128},
|
||||
{184, 150, 247, 255, 236, 224, 128, 128, 128, 128, 128},
|
||||
{77, 110, 216, 255, 236, 230, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 101, 251, 255, 241, 255, 128, 128, 128, 128, 128},
|
||||
{170, 139, 241, 252, 236, 209, 255, 255, 128, 128, 128},
|
||||
{37, 116, 196, 243, 228, 255, 255, 255, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 204, 254, 255, 245, 255, 128, 128, 128, 128, 128},
|
||||
{207, 160, 250, 255, 238, 128, 128, 128, 128, 128, 128},
|
||||
{102, 103, 231, 255, 211, 171, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 152, 252, 255, 240, 255, 128, 128, 128, 128, 128},
|
||||
{177, 135, 243, 255, 234, 225, 128, 128, 128, 128, 128},
|
||||
{80, 129, 211, 255, 194, 224, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 1, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{246, 1, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{255, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
},
|
||||
{
|
||||
{
|
||||
{198, 35, 237, 223, 193, 187, 162, 160, 145, 155, 62},
|
||||
{131, 45, 198, 221, 172, 176, 220, 157, 252, 221, 1},
|
||||
{68, 47, 146, 208, 149, 167, 221, 162, 255, 223, 128},
|
||||
},
|
||||
{
|
||||
{1, 149, 241, 255, 221, 224, 255, 255, 128, 128, 128},
|
||||
{184, 141, 234, 253, 222, 220, 255, 199, 128, 128, 128},
|
||||
{81, 99, 181, 242, 176, 190, 249, 202, 255, 255, 128},
|
||||
},
|
||||
{
|
||||
{1, 129, 232, 253, 214, 197, 242, 196, 255, 255, 128},
|
||||
{99, 121, 210, 250, 201, 198, 255, 202, 128, 128, 128},
|
||||
{23, 91, 163, 242, 170, 187, 247, 210, 255, 255, 128},
|
||||
},
|
||||
{
|
||||
{1, 200, 246, 255, 234, 255, 128, 128, 128, 128, 128},
|
||||
{109, 178, 241, 255, 231, 245, 255, 255, 128, 128, 128},
|
||||
{44, 130, 201, 253, 205, 192, 255, 255, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 132, 239, 251, 219, 209, 255, 165, 128, 128, 128},
|
||||
{94, 136, 225, 251, 218, 190, 255, 255, 128, 128, 128},
|
||||
{22, 100, 174, 245, 186, 161, 255, 199, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 182, 249, 255, 232, 235, 128, 128, 128, 128, 128},
|
||||
{124, 143, 241, 255, 227, 234, 128, 128, 128, 128, 128},
|
||||
{35, 77, 181, 251, 193, 211, 255, 205, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 157, 247, 255, 236, 231, 255, 255, 128, 128, 128},
|
||||
{121, 141, 235, 255, 225, 227, 255, 255, 128, 128, 128},
|
||||
{45, 99, 188, 251, 195, 217, 255, 224, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 1, 251, 255, 213, 255, 128, 128, 128, 128, 128},
|
||||
{203, 1, 248, 255, 255, 128, 128, 128, 128, 128, 128},
|
||||
{137, 1, 177, 255, 224, 255, 128, 128, 128, 128, 128},
|
||||
},
|
||||
},
|
||||
{
|
||||
{
|
||||
{253, 9, 248, 251, 207, 208, 255, 192, 128, 128, 128},
|
||||
{175, 13, 224, 243, 193, 185, 249, 198, 255, 255, 128},
|
||||
{73, 17, 171, 221, 161, 179, 236, 167, 255, 234, 128},
|
||||
},
|
||||
{
|
||||
{1, 95, 247, 253, 212, 183, 255, 255, 128, 128, 128},
|
||||
{239, 90, 244, 250, 211, 209, 255, 255, 128, 128, 128},
|
||||
{155, 77, 195, 248, 188, 195, 255, 255, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 24, 239, 251, 218, 219, 255, 205, 128, 128, 128},
|
||||
{201, 51, 219, 255, 196, 186, 128, 128, 128, 128, 128},
|
||||
{69, 46, 190, 239, 201, 218, 255, 228, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 191, 251, 255, 255, 128, 128, 128, 128, 128, 128},
|
||||
{223, 165, 249, 255, 213, 255, 128, 128, 128, 128, 128},
|
||||
{141, 124, 248, 255, 255, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 16, 248, 255, 255, 128, 128, 128, 128, 128, 128},
|
||||
{190, 36, 230, 255, 236, 255, 128, 128, 128, 128, 128},
|
||||
{149, 1, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 226, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{247, 192, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{240, 128, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 134, 252, 255, 255, 128, 128, 128, 128, 128, 128},
|
||||
{213, 62, 250, 255, 255, 128, 128, 128, 128, 128, 128},
|
||||
{55, 93, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
},
|
||||
{
|
||||
{
|
||||
{202, 24, 213, 235, 186, 191, 220, 160, 240, 175, 255},
|
||||
{126, 38, 182, 232, 169, 184, 228, 174, 255, 187, 128},
|
||||
{61, 46, 138, 219, 151, 178, 240, 170, 255, 216, 128},
|
||||
},
|
||||
{
|
||||
{1, 112, 230, 250, 199, 191, 247, 159, 255, 255, 128},
|
||||
{166, 109, 228, 252, 211, 215, 255, 174, 128, 128, 128},
|
||||
{39, 77, 162, 232, 172, 180, 245, 178, 255, 255, 128},
|
||||
},
|
||||
{
|
||||
{1, 52, 220, 246, 198, 199, 249, 220, 255, 255, 128},
|
||||
{124, 74, 191, 243, 183, 193, 250, 221, 255, 255, 128},
|
||||
{24, 71, 130, 219, 154, 170, 243, 182, 255, 255, 128},
|
||||
},
|
||||
{
|
||||
{1, 182, 225, 249, 219, 240, 255, 224, 128, 128, 128},
|
||||
{149, 150, 226, 252, 216, 205, 255, 171, 128, 128, 128},
|
||||
{28, 108, 170, 242, 183, 194, 254, 223, 255, 255, 128},
|
||||
},
|
||||
{
|
||||
{1, 81, 230, 252, 204, 203, 255, 192, 128, 128, 128},
|
||||
{123, 102, 209, 247, 188, 196, 255, 233, 128, 128, 128},
|
||||
{20, 95, 153, 243, 164, 173, 255, 203, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 222, 248, 255, 216, 213, 128, 128, 128, 128, 128},
|
||||
{168, 175, 246, 252, 235, 205, 255, 255, 128, 128, 128},
|
||||
{47, 116, 215, 255, 211, 212, 255, 255, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 121, 236, 253, 212, 214, 255, 255, 128, 128, 128},
|
||||
{141, 84, 213, 252, 201, 202, 255, 219, 128, 128, 128},
|
||||
{42, 80, 160, 240, 162, 185, 255, 205, 128, 128, 128},
|
||||
},
|
||||
{
|
||||
{1, 1, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{244, 1, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
{238, 1, 255, 128, 128, 128, 128, 128, 128, 128, 128},
|
||||
},
|
||||
},
|
||||
}
|
||||
+603
@@ -0,0 +1,603 @@
|
||||
// Copyright 2014 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package vp8l implements a decoder for the VP8L lossless image format.
|
||||
//
|
||||
// The VP8L specification is at:
|
||||
// https://developers.google.com/speed/webp/docs/riff_container
|
||||
package vp8l // import "golang.org/x/image/vp8l"
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"image"
|
||||
"image/color"
|
||||
"io"
|
||||
)
|
||||
|
||||
var (
|
||||
errInvalidCodeLengths = errors.New("vp8l: invalid code lengths")
|
||||
errInvalidHuffmanTree = errors.New("vp8l: invalid Huffman tree")
|
||||
)
|
||||
|
||||
// colorCacheMultiplier is the multiplier used for the color cache hash
|
||||
// function, specified in section 4.2.3.
|
||||
const colorCacheMultiplier = 0x1e35a7bd
|
||||
|
||||
// distanceMapTable is the look-up table for distanceMap.
|
||||
var distanceMapTable = [120]uint8{
|
||||
0x18, 0x07, 0x17, 0x19, 0x28, 0x06, 0x27, 0x29, 0x16, 0x1a,
|
||||
0x26, 0x2a, 0x38, 0x05, 0x37, 0x39, 0x15, 0x1b, 0x36, 0x3a,
|
||||
0x25, 0x2b, 0x48, 0x04, 0x47, 0x49, 0x14, 0x1c, 0x35, 0x3b,
|
||||
0x46, 0x4a, 0x24, 0x2c, 0x58, 0x45, 0x4b, 0x34, 0x3c, 0x03,
|
||||
0x57, 0x59, 0x13, 0x1d, 0x56, 0x5a, 0x23, 0x2d, 0x44, 0x4c,
|
||||
0x55, 0x5b, 0x33, 0x3d, 0x68, 0x02, 0x67, 0x69, 0x12, 0x1e,
|
||||
0x66, 0x6a, 0x22, 0x2e, 0x54, 0x5c, 0x43, 0x4d, 0x65, 0x6b,
|
||||
0x32, 0x3e, 0x78, 0x01, 0x77, 0x79, 0x53, 0x5d, 0x11, 0x1f,
|
||||
0x64, 0x6c, 0x42, 0x4e, 0x76, 0x7a, 0x21, 0x2f, 0x75, 0x7b,
|
||||
0x31, 0x3f, 0x63, 0x6d, 0x52, 0x5e, 0x00, 0x74, 0x7c, 0x41,
|
||||
0x4f, 0x10, 0x20, 0x62, 0x6e, 0x30, 0x73, 0x7d, 0x51, 0x5f,
|
||||
0x40, 0x72, 0x7e, 0x61, 0x6f, 0x50, 0x71, 0x7f, 0x60, 0x70,
|
||||
}
|
||||
|
||||
// distanceMap maps a LZ77 backwards reference distance to a two-dimensional
|
||||
// pixel offset, specified in section 4.2.2.
|
||||
func distanceMap(w int32, code uint32) int32 {
|
||||
if int32(code) > int32(len(distanceMapTable)) {
|
||||
return int32(code) - int32(len(distanceMapTable))
|
||||
}
|
||||
distCode := int32(distanceMapTable[code-1])
|
||||
yOffset := distCode >> 4
|
||||
xOffset := 8 - distCode&0xf
|
||||
if d := yOffset*w + xOffset; d >= 1 {
|
||||
return d
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
// decoder holds the bit-stream for a VP8L image.
|
||||
type decoder struct {
|
||||
r io.ByteReader
|
||||
bits uint32
|
||||
nBits uint32
|
||||
}
|
||||
|
||||
// read reads the next n bits from the decoder's bit-stream.
|
||||
func (d *decoder) read(n uint32) (uint32, error) {
|
||||
for d.nBits < n {
|
||||
c, err := d.r.ReadByte()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
d.bits |= uint32(c) << d.nBits
|
||||
d.nBits += 8
|
||||
}
|
||||
u := d.bits & (1<<n - 1)
|
||||
d.bits >>= n
|
||||
d.nBits -= n
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// decodeTransform decodes the next transform and the width of the image after
|
||||
// transformation (or equivalently, before inverse transformation), specified
|
||||
// in section 3.
|
||||
func (d *decoder) decodeTransform(w int32, h int32) (t transform, newWidth int32, err error) {
|
||||
t.oldWidth = w
|
||||
t.transformType, err = d.read(2)
|
||||
if err != nil {
|
||||
return transform{}, 0, err
|
||||
}
|
||||
switch t.transformType {
|
||||
case transformTypePredictor, transformTypeCrossColor:
|
||||
t.bits, err = d.read(3)
|
||||
if err != nil {
|
||||
return transform{}, 0, err
|
||||
}
|
||||
t.bits += 2
|
||||
t.pix, err = d.decodePix(nTiles(w, t.bits), nTiles(h, t.bits), 0, false)
|
||||
if err != nil {
|
||||
return transform{}, 0, err
|
||||
}
|
||||
case transformTypeSubtractGreen:
|
||||
// No-op.
|
||||
case transformTypeColorIndexing:
|
||||
nColors, err := d.read(8)
|
||||
if err != nil {
|
||||
return transform{}, 0, err
|
||||
}
|
||||
nColors++
|
||||
t.bits = 0
|
||||
switch {
|
||||
case nColors <= 2:
|
||||
t.bits = 3
|
||||
case nColors <= 4:
|
||||
t.bits = 2
|
||||
case nColors <= 16:
|
||||
t.bits = 1
|
||||
}
|
||||
w = nTiles(w, t.bits)
|
||||
pix, err := d.decodePix(int32(nColors), 1, 4*256, false)
|
||||
if err != nil {
|
||||
return transform{}, 0, err
|
||||
}
|
||||
for p := 4; p < len(pix); p += 4 {
|
||||
pix[p+0] += pix[p-4]
|
||||
pix[p+1] += pix[p-3]
|
||||
pix[p+2] += pix[p-2]
|
||||
pix[p+3] += pix[p-1]
|
||||
}
|
||||
// The spec says that "if the index is equal or larger than color_table_size,
|
||||
// the argb color value should be set to 0x00000000 (transparent black)."
|
||||
// We re-slice up to 256 4-byte pixels.
|
||||
t.pix = pix[:4*256]
|
||||
}
|
||||
return t, w, nil
|
||||
}
|
||||
|
||||
// repeatsCodeLength is the minimum code length for repeated codes.
|
||||
const repeatsCodeLength = 16
|
||||
|
||||
// These magic numbers are specified at the end of section 5.2.2.
|
||||
// The 3-length arrays apply to code lengths >= repeatsCodeLength.
|
||||
var (
|
||||
codeLengthCodeOrder = [19]uint8{
|
||||
17, 18, 0, 1, 2, 3, 4, 5, 16, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
|
||||
}
|
||||
repeatBits = [3]uint8{2, 3, 7}
|
||||
repeatOffsets = [3]uint8{3, 3, 11}
|
||||
)
|
||||
|
||||
// decodeCodeLengths decodes a Huffman tree's code lengths which are themselves
|
||||
// encoded via a Huffman tree, specified in section 5.2.2.
|
||||
func (d *decoder) decodeCodeLengths(dst []uint32, codeLengthCodeLengths []uint32) error {
|
||||
h := hTree{}
|
||||
if err := h.build(codeLengthCodeLengths); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
maxSymbol := len(dst)
|
||||
useLength, err := d.read(1)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if useLength != 0 {
|
||||
n, err := d.read(3)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n = 2 + 2*n
|
||||
ms, err := d.read(n)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
maxSymbol = int(ms) + 2
|
||||
if maxSymbol > len(dst) {
|
||||
return errInvalidCodeLengths
|
||||
}
|
||||
}
|
||||
|
||||
// The spec says that "if code 16 [meaning repeat] is used before
|
||||
// a non-zero value has been emitted, a value of 8 is repeated."
|
||||
prevCodeLength := uint32(8)
|
||||
|
||||
for symbol := 0; symbol < len(dst); {
|
||||
if maxSymbol == 0 {
|
||||
break
|
||||
}
|
||||
maxSymbol--
|
||||
codeLength, err := h.next(d)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if codeLength < repeatsCodeLength {
|
||||
dst[symbol] = codeLength
|
||||
symbol++
|
||||
if codeLength != 0 {
|
||||
prevCodeLength = codeLength
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
repeat, err := d.read(uint32(repeatBits[codeLength-repeatsCodeLength]))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
repeat += uint32(repeatOffsets[codeLength-repeatsCodeLength])
|
||||
if symbol+int(repeat) > len(dst) {
|
||||
return errInvalidCodeLengths
|
||||
}
|
||||
// A code length of 16 repeats the previous non-zero code.
|
||||
// A code length of 17 or 18 repeats zeroes.
|
||||
cl := uint32(0)
|
||||
if codeLength == 16 {
|
||||
cl = prevCodeLength
|
||||
}
|
||||
for ; repeat > 0; repeat-- {
|
||||
dst[symbol] = cl
|
||||
symbol++
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// decodeHuffmanTree decodes a Huffman tree into h.
|
||||
func (d *decoder) decodeHuffmanTree(h *hTree, alphabetSize uint32) error {
|
||||
useSimple, err := d.read(1)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if useSimple != 0 {
|
||||
nSymbols, err := d.read(1)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
nSymbols++
|
||||
firstSymbolLengthCode, err := d.read(1)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
firstSymbolLengthCode = 7*firstSymbolLengthCode + 1
|
||||
var symbols [2]uint32
|
||||
symbols[0], err = d.read(firstSymbolLengthCode)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if nSymbols == 2 {
|
||||
symbols[1], err = d.read(8)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return h.buildSimple(nSymbols, symbols, alphabetSize)
|
||||
}
|
||||
|
||||
nCodes, err := d.read(4)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
nCodes += 4
|
||||
if int(nCodes) > len(codeLengthCodeOrder) {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
codeLengthCodeLengths := [len(codeLengthCodeOrder)]uint32{}
|
||||
for i := uint32(0); i < nCodes; i++ {
|
||||
codeLengthCodeLengths[codeLengthCodeOrder[i]], err = d.read(3)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
codeLengths := make([]uint32, alphabetSize)
|
||||
if err = d.decodeCodeLengths(codeLengths, codeLengthCodeLengths[:]); err != nil {
|
||||
return err
|
||||
}
|
||||
return h.build(codeLengths)
|
||||
}
|
||||
|
||||
const (
|
||||
huffGreen = 0
|
||||
huffRed = 1
|
||||
huffBlue = 2
|
||||
huffAlpha = 3
|
||||
huffDistance = 4
|
||||
nHuff = 5
|
||||
)
|
||||
|
||||
// hGroup is an array of 5 Huffman trees.
|
||||
type hGroup [nHuff]hTree
|
||||
|
||||
// decodeHuffmanGroups decodes the one or more hGroups used to decode the pixel
|
||||
// data. If one hGroup is used for the entire image, then hPix and hBits will
|
||||
// be zero. If more than one hGroup is used, then hPix contains the meta-image
|
||||
// that maps tiles to hGroup index, and hBits contains the log-2 tile size.
|
||||
func (d *decoder) decodeHuffmanGroups(w int32, h int32, topLevel bool, ccBits uint32) (
|
||||
hGroups []hGroup, hPix []byte, hBits uint32, err error) {
|
||||
|
||||
maxHGroupIndex := 0
|
||||
if topLevel {
|
||||
useMeta, err := d.read(1)
|
||||
if err != nil {
|
||||
return nil, nil, 0, err
|
||||
}
|
||||
if useMeta != 0 {
|
||||
hBits, err = d.read(3)
|
||||
if err != nil {
|
||||
return nil, nil, 0, err
|
||||
}
|
||||
hBits += 2
|
||||
hPix, err = d.decodePix(nTiles(w, hBits), nTiles(h, hBits), 0, false)
|
||||
if err != nil {
|
||||
return nil, nil, 0, err
|
||||
}
|
||||
for p := 0; p < len(hPix); p += 4 {
|
||||
i := int(hPix[p])<<8 | int(hPix[p+1])
|
||||
if maxHGroupIndex < i {
|
||||
maxHGroupIndex = i
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
hGroups = make([]hGroup, maxHGroupIndex+1)
|
||||
for i := range hGroups {
|
||||
for j, alphabetSize := range alphabetSizes {
|
||||
if j == 0 && ccBits > 0 {
|
||||
alphabetSize += 1 << ccBits
|
||||
}
|
||||
if err := d.decodeHuffmanTree(&hGroups[i][j], alphabetSize); err != nil {
|
||||
return nil, nil, 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
return hGroups, hPix, hBits, nil
|
||||
}
|
||||
|
||||
const (
|
||||
nLiteralCodes = 256
|
||||
nLengthCodes = 24
|
||||
nDistanceCodes = 40
|
||||
)
|
||||
|
||||
var alphabetSizes = [nHuff]uint32{
|
||||
nLiteralCodes + nLengthCodes,
|
||||
nLiteralCodes,
|
||||
nLiteralCodes,
|
||||
nLiteralCodes,
|
||||
nDistanceCodes,
|
||||
}
|
||||
|
||||
// decodePix decodes pixel data, specified in section 5.2.2.
|
||||
func (d *decoder) decodePix(w int32, h int32, minCap int32, topLevel bool) ([]byte, error) {
|
||||
// Decode the color cache parameters.
|
||||
ccBits, ccShift, ccEntries := uint32(0), uint32(0), ([]uint32)(nil)
|
||||
useColorCache, err := d.read(1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if useColorCache != 0 {
|
||||
ccBits, err = d.read(4)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if ccBits < 1 || 11 < ccBits {
|
||||
return nil, errors.New("vp8l: invalid color cache parameters")
|
||||
}
|
||||
ccShift = 32 - ccBits
|
||||
ccEntries = make([]uint32, 1<<ccBits)
|
||||
}
|
||||
|
||||
// Decode the Huffman groups.
|
||||
hGroups, hPix, hBits, err := d.decodeHuffmanGroups(w, h, topLevel, ccBits)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hMask, tilesPerRow := int32(0), int32(0)
|
||||
if hBits != 0 {
|
||||
hMask, tilesPerRow = 1<<hBits-1, nTiles(w, hBits)
|
||||
}
|
||||
|
||||
// Decode the pixels.
|
||||
if minCap < 4*w*h {
|
||||
minCap = 4 * w * h
|
||||
}
|
||||
pix := make([]byte, 4*w*h, minCap)
|
||||
p, cachedP := 0, 0
|
||||
x, y := int32(0), int32(0)
|
||||
hg, lookupHG := &hGroups[0], hMask != 0
|
||||
for p < len(pix) {
|
||||
if lookupHG {
|
||||
i := 4 * (tilesPerRow*(y>>hBits) + (x >> hBits))
|
||||
hg = &hGroups[uint32(hPix[i])<<8|uint32(hPix[i+1])]
|
||||
}
|
||||
|
||||
green, err := hg[huffGreen].next(d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch {
|
||||
case green < nLiteralCodes:
|
||||
// We have a literal pixel.
|
||||
red, err := hg[huffRed].next(d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
blue, err := hg[huffBlue].next(d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
alpha, err := hg[huffAlpha].next(d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pix[p+0] = uint8(red)
|
||||
pix[p+1] = uint8(green)
|
||||
pix[p+2] = uint8(blue)
|
||||
pix[p+3] = uint8(alpha)
|
||||
p += 4
|
||||
|
||||
x++
|
||||
if x == w {
|
||||
x, y = 0, y+1
|
||||
}
|
||||
lookupHG = hMask != 0 && x&hMask == 0
|
||||
|
||||
case green < nLiteralCodes+nLengthCodes:
|
||||
// We have a LZ77 backwards reference.
|
||||
length, err := d.lz77Param(green - nLiteralCodes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
distSym, err := hg[huffDistance].next(d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
distCode, err := d.lz77Param(distSym)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dist := distanceMap(w, distCode)
|
||||
pEnd := p + 4*int(length)
|
||||
q := p - 4*int(dist)
|
||||
qEnd := pEnd - 4*int(dist)
|
||||
if p < 0 || len(pix) < pEnd || q < 0 || len(pix) < qEnd {
|
||||
return nil, errors.New("vp8l: invalid LZ77 parameters")
|
||||
}
|
||||
for ; p < pEnd; p, q = p+1, q+1 {
|
||||
pix[p] = pix[q]
|
||||
}
|
||||
|
||||
x += int32(length)
|
||||
for x >= w {
|
||||
x, y = x-w, y+1
|
||||
}
|
||||
lookupHG = hMask != 0
|
||||
|
||||
default:
|
||||
// We have a color cache lookup. First, insert previous pixels
|
||||
// into the cache. Note that VP8L assumes ARGB order, but the
|
||||
// Go image.RGBA type is in RGBA order.
|
||||
for ; cachedP < p; cachedP += 4 {
|
||||
argb := uint32(pix[cachedP+0])<<16 |
|
||||
uint32(pix[cachedP+1])<<8 |
|
||||
uint32(pix[cachedP+2])<<0 |
|
||||
uint32(pix[cachedP+3])<<24
|
||||
ccEntries[(argb*colorCacheMultiplier)>>ccShift] = argb
|
||||
}
|
||||
green -= nLiteralCodes + nLengthCodes
|
||||
if int(green) >= len(ccEntries) {
|
||||
return nil, errors.New("vp8l: invalid color cache index")
|
||||
}
|
||||
argb := ccEntries[green]
|
||||
pix[p+0] = uint8(argb >> 16)
|
||||
pix[p+1] = uint8(argb >> 8)
|
||||
pix[p+2] = uint8(argb >> 0)
|
||||
pix[p+3] = uint8(argb >> 24)
|
||||
p += 4
|
||||
|
||||
x++
|
||||
if x == w {
|
||||
x, y = 0, y+1
|
||||
}
|
||||
lookupHG = hMask != 0 && x&hMask == 0
|
||||
}
|
||||
}
|
||||
return pix, nil
|
||||
}
|
||||
|
||||
// lz77Param returns the next LZ77 parameter: a length or a distance, specified
|
||||
// in section 4.2.2.
|
||||
func (d *decoder) lz77Param(symbol uint32) (uint32, error) {
|
||||
if symbol < 4 {
|
||||
return symbol + 1, nil
|
||||
}
|
||||
extraBits := (symbol - 2) >> 1
|
||||
offset := (2 + symbol&1) << extraBits
|
||||
n, err := d.read(extraBits)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return offset + n + 1, nil
|
||||
}
|
||||
|
||||
// decodeHeader decodes the VP8L header from r.
|
||||
func decodeHeader(r io.Reader) (d *decoder, w int32, h int32, err error) {
|
||||
rr, ok := r.(io.ByteReader)
|
||||
if !ok {
|
||||
rr = bufio.NewReader(r)
|
||||
}
|
||||
d = &decoder{r: rr}
|
||||
magic, err := d.read(8)
|
||||
if err != nil {
|
||||
return nil, 0, 0, err
|
||||
}
|
||||
if magic != 0x2f {
|
||||
return nil, 0, 0, errors.New("vp8l: invalid header")
|
||||
}
|
||||
width, err := d.read(14)
|
||||
if err != nil {
|
||||
return nil, 0, 0, err
|
||||
}
|
||||
width++
|
||||
height, err := d.read(14)
|
||||
if err != nil {
|
||||
return nil, 0, 0, err
|
||||
}
|
||||
height++
|
||||
_, err = d.read(1) // Read and ignore the hasAlpha hint.
|
||||
if err != nil {
|
||||
return nil, 0, 0, err
|
||||
}
|
||||
version, err := d.read(3)
|
||||
if err != nil {
|
||||
return nil, 0, 0, err
|
||||
}
|
||||
if version != 0 {
|
||||
return nil, 0, 0, errors.New("vp8l: invalid version")
|
||||
}
|
||||
return d, int32(width), int32(height), nil
|
||||
}
|
||||
|
||||
// DecodeConfig decodes the color model and dimensions of a VP8L image from r.
|
||||
func DecodeConfig(r io.Reader) (image.Config, error) {
|
||||
_, w, h, err := decodeHeader(r)
|
||||
if err != nil {
|
||||
return image.Config{}, err
|
||||
}
|
||||
return image.Config{
|
||||
ColorModel: color.NRGBAModel,
|
||||
Width: int(w),
|
||||
Height: int(h),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Decode decodes a VP8L image from r.
|
||||
func Decode(r io.Reader) (image.Image, error) {
|
||||
d, w, h, err := decodeHeader(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Decode the transforms.
|
||||
var (
|
||||
nTransforms int
|
||||
transforms [nTransformTypes]transform
|
||||
transformsSeen [nTransformTypes]bool
|
||||
originalW = w
|
||||
)
|
||||
for {
|
||||
more, err := d.read(1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if more == 0 {
|
||||
break
|
||||
}
|
||||
var t transform
|
||||
t, w, err = d.decodeTransform(w, h)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if transformsSeen[t.transformType] {
|
||||
return nil, errors.New("vp8l: repeated transform")
|
||||
}
|
||||
transformsSeen[t.transformType] = true
|
||||
transforms[nTransforms] = t
|
||||
nTransforms++
|
||||
}
|
||||
// Decode the transformed pixels.
|
||||
pix, err := d.decodePix(w, h, 0, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Apply the inverse transformations.
|
||||
for i := nTransforms - 1; i >= 0; i-- {
|
||||
t := &transforms[i]
|
||||
pix = inverseTransforms[t.transformType](t, pix, h)
|
||||
}
|
||||
return &image.NRGBA{
|
||||
Pix: pix,
|
||||
Stride: 4 * int(originalW),
|
||||
Rect: image.Rect(0, 0, int(originalW), int(h)),
|
||||
}, nil
|
||||
}
|
||||
+245
@@ -0,0 +1,245 @@
|
||||
// Copyright 2014 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8l
|
||||
|
||||
import (
|
||||
"io"
|
||||
)
|
||||
|
||||
// reverseBits reverses the bits in a byte.
|
||||
var reverseBits = [256]uint8{
|
||||
0x00, 0x80, 0x40, 0xc0, 0x20, 0xa0, 0x60, 0xe0, 0x10, 0x90, 0x50, 0xd0, 0x30, 0xb0, 0x70, 0xf0,
|
||||
0x08, 0x88, 0x48, 0xc8, 0x28, 0xa8, 0x68, 0xe8, 0x18, 0x98, 0x58, 0xd8, 0x38, 0xb8, 0x78, 0xf8,
|
||||
0x04, 0x84, 0x44, 0xc4, 0x24, 0xa4, 0x64, 0xe4, 0x14, 0x94, 0x54, 0xd4, 0x34, 0xb4, 0x74, 0xf4,
|
||||
0x0c, 0x8c, 0x4c, 0xcc, 0x2c, 0xac, 0x6c, 0xec, 0x1c, 0x9c, 0x5c, 0xdc, 0x3c, 0xbc, 0x7c, 0xfc,
|
||||
0x02, 0x82, 0x42, 0xc2, 0x22, 0xa2, 0x62, 0xe2, 0x12, 0x92, 0x52, 0xd2, 0x32, 0xb2, 0x72, 0xf2,
|
||||
0x0a, 0x8a, 0x4a, 0xca, 0x2a, 0xaa, 0x6a, 0xea, 0x1a, 0x9a, 0x5a, 0xda, 0x3a, 0xba, 0x7a, 0xfa,
|
||||
0x06, 0x86, 0x46, 0xc6, 0x26, 0xa6, 0x66, 0xe6, 0x16, 0x96, 0x56, 0xd6, 0x36, 0xb6, 0x76, 0xf6,
|
||||
0x0e, 0x8e, 0x4e, 0xce, 0x2e, 0xae, 0x6e, 0xee, 0x1e, 0x9e, 0x5e, 0xde, 0x3e, 0xbe, 0x7e, 0xfe,
|
||||
0x01, 0x81, 0x41, 0xc1, 0x21, 0xa1, 0x61, 0xe1, 0x11, 0x91, 0x51, 0xd1, 0x31, 0xb1, 0x71, 0xf1,
|
||||
0x09, 0x89, 0x49, 0xc9, 0x29, 0xa9, 0x69, 0xe9, 0x19, 0x99, 0x59, 0xd9, 0x39, 0xb9, 0x79, 0xf9,
|
||||
0x05, 0x85, 0x45, 0xc5, 0x25, 0xa5, 0x65, 0xe5, 0x15, 0x95, 0x55, 0xd5, 0x35, 0xb5, 0x75, 0xf5,
|
||||
0x0d, 0x8d, 0x4d, 0xcd, 0x2d, 0xad, 0x6d, 0xed, 0x1d, 0x9d, 0x5d, 0xdd, 0x3d, 0xbd, 0x7d, 0xfd,
|
||||
0x03, 0x83, 0x43, 0xc3, 0x23, 0xa3, 0x63, 0xe3, 0x13, 0x93, 0x53, 0xd3, 0x33, 0xb3, 0x73, 0xf3,
|
||||
0x0b, 0x8b, 0x4b, 0xcb, 0x2b, 0xab, 0x6b, 0xeb, 0x1b, 0x9b, 0x5b, 0xdb, 0x3b, 0xbb, 0x7b, 0xfb,
|
||||
0x07, 0x87, 0x47, 0xc7, 0x27, 0xa7, 0x67, 0xe7, 0x17, 0x97, 0x57, 0xd7, 0x37, 0xb7, 0x77, 0xf7,
|
||||
0x0f, 0x8f, 0x4f, 0xcf, 0x2f, 0xaf, 0x6f, 0xef, 0x1f, 0x9f, 0x5f, 0xdf, 0x3f, 0xbf, 0x7f, 0xff,
|
||||
}
|
||||
|
||||
// hNode is a node in a Huffman tree.
|
||||
type hNode struct {
|
||||
// symbol is the symbol held by this node.
|
||||
symbol uint32
|
||||
// children, if positive, is the hTree.nodes index of the first of
|
||||
// this node's two children. Zero means an uninitialized node,
|
||||
// and -1 means a leaf node.
|
||||
children int32
|
||||
}
|
||||
|
||||
const leafNode = -1
|
||||
|
||||
// lutSize is the log-2 size of an hTree's look-up table.
|
||||
const lutSize, lutMask = 7, 1<<7 - 1
|
||||
|
||||
// hTree is a Huffman tree.
|
||||
type hTree struct {
|
||||
// nodes are the nodes of the Huffman tree. During construction,
|
||||
// len(nodes) grows from 1 up to cap(nodes) by steps of two.
|
||||
// After construction, len(nodes) == cap(nodes), and both equal
|
||||
// 2*theNumberOfSymbols - 1.
|
||||
nodes []hNode
|
||||
// lut is a look-up table for walking the nodes. The x in lut[x] is
|
||||
// the next lutSize bits in the bit-stream. The low 8 bits of lut[x]
|
||||
// equals 1 plus the number of bits in the next code, or 0 if the
|
||||
// next code requires more than lutSize bits. The high 24 bits are:
|
||||
// - the symbol, if the code requires lutSize or fewer bits, or
|
||||
// - the hTree.nodes index to start the tree traversal from, if
|
||||
// the next code requires more than lutSize bits.
|
||||
lut [1 << lutSize]uint32
|
||||
}
|
||||
|
||||
// insert inserts into the hTree a symbol whose encoding is the least
|
||||
// significant codeLength bits of code.
|
||||
func (h *hTree) insert(symbol uint32, code uint32, codeLength uint32) error {
|
||||
if symbol > 0xffff || codeLength > 0xfe {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
baseCode := uint32(0)
|
||||
if codeLength > lutSize {
|
||||
baseCode = uint32(reverseBits[(code>>(codeLength-lutSize))&0xff]) >> (8 - lutSize)
|
||||
} else {
|
||||
baseCode = uint32(reverseBits[code&0xff]) >> (8 - codeLength)
|
||||
for i := 0; i < 1<<(lutSize-codeLength); i++ {
|
||||
h.lut[baseCode|uint32(i)<<codeLength] = symbol<<8 | (codeLength + 1)
|
||||
}
|
||||
}
|
||||
|
||||
n := uint32(0)
|
||||
for jump := lutSize; codeLength > 0; {
|
||||
codeLength--
|
||||
if int(n) > len(h.nodes) {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
switch h.nodes[n].children {
|
||||
case leafNode:
|
||||
return errInvalidHuffmanTree
|
||||
case 0:
|
||||
if len(h.nodes) == cap(h.nodes) {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
// Create two empty child nodes.
|
||||
h.nodes[n].children = int32(len(h.nodes))
|
||||
h.nodes = h.nodes[:len(h.nodes)+2]
|
||||
}
|
||||
n = uint32(h.nodes[n].children) + 1&(code>>codeLength)
|
||||
jump--
|
||||
if jump == 0 && h.lut[baseCode] == 0 {
|
||||
h.lut[baseCode] = n << 8
|
||||
}
|
||||
}
|
||||
|
||||
switch h.nodes[n].children {
|
||||
case leafNode:
|
||||
// No-op.
|
||||
case 0:
|
||||
// Turn the uninitialized node into a leaf.
|
||||
h.nodes[n].children = leafNode
|
||||
default:
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
h.nodes[n].symbol = symbol
|
||||
return nil
|
||||
}
|
||||
|
||||
// codeLengthsToCodes returns the canonical Huffman codes implied by the
|
||||
// sequence of code lengths.
|
||||
func codeLengthsToCodes(codeLengths []uint32) ([]uint32, error) {
|
||||
maxCodeLength := uint32(0)
|
||||
for _, cl := range codeLengths {
|
||||
if maxCodeLength < cl {
|
||||
maxCodeLength = cl
|
||||
}
|
||||
}
|
||||
const maxAllowedCodeLength = 15
|
||||
if len(codeLengths) == 0 || maxCodeLength > maxAllowedCodeLength {
|
||||
return nil, errInvalidHuffmanTree
|
||||
}
|
||||
histogram := [maxAllowedCodeLength + 1]uint32{}
|
||||
for _, cl := range codeLengths {
|
||||
histogram[cl]++
|
||||
}
|
||||
currCode, nextCodes := uint32(0), [maxAllowedCodeLength + 1]uint32{}
|
||||
for cl := 1; cl < len(nextCodes); cl++ {
|
||||
currCode = (currCode + histogram[cl-1]) << 1
|
||||
nextCodes[cl] = currCode
|
||||
}
|
||||
codes := make([]uint32, len(codeLengths))
|
||||
for symbol, cl := range codeLengths {
|
||||
if cl > 0 {
|
||||
codes[symbol] = nextCodes[cl]
|
||||
nextCodes[cl]++
|
||||
}
|
||||
}
|
||||
return codes, nil
|
||||
}
|
||||
|
||||
// build builds a canonical Huffman tree from the given code lengths.
|
||||
func (h *hTree) build(codeLengths []uint32) error {
|
||||
// Calculate the number of symbols.
|
||||
var nSymbols, lastSymbol uint32
|
||||
for symbol, cl := range codeLengths {
|
||||
if cl != 0 {
|
||||
nSymbols++
|
||||
lastSymbol = uint32(symbol)
|
||||
}
|
||||
}
|
||||
if nSymbols == 0 {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
h.nodes = make([]hNode, 1, 2*nSymbols-1)
|
||||
// Handle the trivial case.
|
||||
if nSymbols == 1 {
|
||||
if len(codeLengths) <= int(lastSymbol) {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
return h.insert(lastSymbol, 0, 0)
|
||||
}
|
||||
// Handle the non-trivial case.
|
||||
codes, err := codeLengthsToCodes(codeLengths)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for symbol, cl := range codeLengths {
|
||||
if cl > 0 {
|
||||
if err := h.insert(uint32(symbol), codes[symbol], cl); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildSimple builds a Huffman tree with 1 or 2 symbols.
|
||||
func (h *hTree) buildSimple(nSymbols uint32, symbols [2]uint32, alphabetSize uint32) error {
|
||||
h.nodes = make([]hNode, 1, 2*nSymbols-1)
|
||||
for i := uint32(0); i < nSymbols; i++ {
|
||||
if symbols[i] >= alphabetSize {
|
||||
return errInvalidHuffmanTree
|
||||
}
|
||||
if err := h.insert(symbols[i], i, nSymbols-1); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// next returns the next Huffman-encoded symbol from the bit-stream d.
|
||||
func (h *hTree) next(d *decoder) (uint32, error) {
|
||||
var n uint32
|
||||
// Read enough bits so that we can use the look-up table.
|
||||
if d.nBits < lutSize {
|
||||
c, err := d.r.ReadByte()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
// There are no more bytes of data, but we may still be able
|
||||
// to read the next symbol out of the previously read bits.
|
||||
goto slowPath
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
d.bits |= uint32(c) << d.nBits
|
||||
d.nBits += 8
|
||||
}
|
||||
// Use the look-up table.
|
||||
n = h.lut[d.bits&lutMask]
|
||||
if b := n & 0xff; b != 0 {
|
||||
b--
|
||||
d.bits >>= b
|
||||
d.nBits -= b
|
||||
return n >> 8, nil
|
||||
}
|
||||
n >>= 8
|
||||
d.bits >>= lutSize
|
||||
d.nBits -= lutSize
|
||||
|
||||
slowPath:
|
||||
for h.nodes[n].children != leafNode {
|
||||
if d.nBits == 0 {
|
||||
c, err := d.r.ReadByte()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
d.bits = uint32(c)
|
||||
d.nBits = 8
|
||||
}
|
||||
n = uint32(h.nodes[n].children) + 1&d.bits
|
||||
d.bits >>= 1
|
||||
d.nBits--
|
||||
}
|
||||
return h.nodes[n].symbol, nil
|
||||
}
|
||||
+299
@@ -0,0 +1,299 @@
|
||||
// Copyright 2014 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package vp8l
|
||||
|
||||
// This file deals with image transforms, specified in section 3.
|
||||
|
||||
// nTiles returns the number of tiles needed to cover size pixels, where each
|
||||
// tile's side is 1<<bits pixels long.
|
||||
func nTiles(size int32, bits uint32) int32 {
|
||||
return (size + 1<<bits - 1) >> bits
|
||||
}
|
||||
|
||||
const (
|
||||
transformTypePredictor = 0
|
||||
transformTypeCrossColor = 1
|
||||
transformTypeSubtractGreen = 2
|
||||
transformTypeColorIndexing = 3
|
||||
nTransformTypes = 4
|
||||
)
|
||||
|
||||
// transform holds the parameters for an invertible transform.
|
||||
type transform struct {
|
||||
// transformType is the type of the transform.
|
||||
transformType uint32
|
||||
// oldWidth is the width of the image before transformation (or
|
||||
// equivalently, after inverse transformation). The color-indexing
|
||||
// transform can reduce the width. For example, a 50-pixel-wide
|
||||
// image that only needs 4 bits (half a byte) per color index can
|
||||
// be transformed into a 25-pixel-wide image.
|
||||
oldWidth int32
|
||||
// bits is the log-2 size of the transform's tiles, for the predictor
|
||||
// and cross-color transforms. 8>>bits is the number of bits per
|
||||
// color index, for the color-index transform.
|
||||
bits uint32
|
||||
// pix is the tile values, for the predictor and cross-color
|
||||
// transforms, and the color palette, for the color-index transform.
|
||||
pix []byte
|
||||
}
|
||||
|
||||
var inverseTransforms = [nTransformTypes]func(*transform, []byte, int32) []byte{
|
||||
transformTypePredictor: inversePredictor,
|
||||
transformTypeCrossColor: inverseCrossColor,
|
||||
transformTypeSubtractGreen: inverseSubtractGreen,
|
||||
transformTypeColorIndexing: inverseColorIndexing,
|
||||
}
|
||||
|
||||
func inversePredictor(t *transform, pix []byte, h int32) []byte {
|
||||
if t.oldWidth == 0 || h == 0 {
|
||||
return pix
|
||||
}
|
||||
// The first pixel's predictor is mode 0 (opaque black).
|
||||
pix[3] += 0xff
|
||||
p, mask := int32(4), int32(1)<<t.bits-1
|
||||
for x := int32(1); x < t.oldWidth; x++ {
|
||||
// The rest of the first row's predictor is mode 1 (L).
|
||||
pix[p+0] += pix[p-4]
|
||||
pix[p+1] += pix[p-3]
|
||||
pix[p+2] += pix[p-2]
|
||||
pix[p+3] += pix[p-1]
|
||||
p += 4
|
||||
}
|
||||
top, tilesPerRow := 0, nTiles(t.oldWidth, t.bits)
|
||||
for y := int32(1); y < h; y++ {
|
||||
// The first column's predictor is mode 2 (T).
|
||||
pix[p+0] += pix[top+0]
|
||||
pix[p+1] += pix[top+1]
|
||||
pix[p+2] += pix[top+2]
|
||||
pix[p+3] += pix[top+3]
|
||||
p, top = p+4, top+4
|
||||
|
||||
q := 4 * (y >> t.bits) * tilesPerRow
|
||||
predictorMode := t.pix[q+1] & 0x0f
|
||||
q += 4
|
||||
for x := int32(1); x < t.oldWidth; x++ {
|
||||
if x&mask == 0 {
|
||||
predictorMode = t.pix[q+1] & 0x0f
|
||||
q += 4
|
||||
}
|
||||
switch predictorMode {
|
||||
case 0: // Opaque black.
|
||||
pix[p+3] += 0xff
|
||||
|
||||
case 1: // L.
|
||||
pix[p+0] += pix[p-4]
|
||||
pix[p+1] += pix[p-3]
|
||||
pix[p+2] += pix[p-2]
|
||||
pix[p+3] += pix[p-1]
|
||||
|
||||
case 2: // T.
|
||||
pix[p+0] += pix[top+0]
|
||||
pix[p+1] += pix[top+1]
|
||||
pix[p+2] += pix[top+2]
|
||||
pix[p+3] += pix[top+3]
|
||||
|
||||
case 3: // TR.
|
||||
pix[p+0] += pix[top+4]
|
||||
pix[p+1] += pix[top+5]
|
||||
pix[p+2] += pix[top+6]
|
||||
pix[p+3] += pix[top+7]
|
||||
|
||||
case 4: // TL.
|
||||
pix[p+0] += pix[top-4]
|
||||
pix[p+1] += pix[top-3]
|
||||
pix[p+2] += pix[top-2]
|
||||
pix[p+3] += pix[top-1]
|
||||
|
||||
case 5: // Average2(Average2(L, TR), T).
|
||||
pix[p+0] += avg2(avg2(pix[p-4], pix[top+4]), pix[top+0])
|
||||
pix[p+1] += avg2(avg2(pix[p-3], pix[top+5]), pix[top+1])
|
||||
pix[p+2] += avg2(avg2(pix[p-2], pix[top+6]), pix[top+2])
|
||||
pix[p+3] += avg2(avg2(pix[p-1], pix[top+7]), pix[top+3])
|
||||
|
||||
case 6: // Average2(L, TL).
|
||||
pix[p+0] += avg2(pix[p-4], pix[top-4])
|
||||
pix[p+1] += avg2(pix[p-3], pix[top-3])
|
||||
pix[p+2] += avg2(pix[p-2], pix[top-2])
|
||||
pix[p+3] += avg2(pix[p-1], pix[top-1])
|
||||
|
||||
case 7: // Average2(L, T).
|
||||
pix[p+0] += avg2(pix[p-4], pix[top+0])
|
||||
pix[p+1] += avg2(pix[p-3], pix[top+1])
|
||||
pix[p+2] += avg2(pix[p-2], pix[top+2])
|
||||
pix[p+3] += avg2(pix[p-1], pix[top+3])
|
||||
|
||||
case 8: // Average2(TL, T).
|
||||
pix[p+0] += avg2(pix[top-4], pix[top+0])
|
||||
pix[p+1] += avg2(pix[top-3], pix[top+1])
|
||||
pix[p+2] += avg2(pix[top-2], pix[top+2])
|
||||
pix[p+3] += avg2(pix[top-1], pix[top+3])
|
||||
|
||||
case 9: // Average2(T, TR).
|
||||
pix[p+0] += avg2(pix[top+0], pix[top+4])
|
||||
pix[p+1] += avg2(pix[top+1], pix[top+5])
|
||||
pix[p+2] += avg2(pix[top+2], pix[top+6])
|
||||
pix[p+3] += avg2(pix[top+3], pix[top+7])
|
||||
|
||||
case 10: // Average2(Average2(L, TL), Average2(T, TR)).
|
||||
pix[p+0] += avg2(avg2(pix[p-4], pix[top-4]), avg2(pix[top+0], pix[top+4]))
|
||||
pix[p+1] += avg2(avg2(pix[p-3], pix[top-3]), avg2(pix[top+1], pix[top+5]))
|
||||
pix[p+2] += avg2(avg2(pix[p-2], pix[top-2]), avg2(pix[top+2], pix[top+6]))
|
||||
pix[p+3] += avg2(avg2(pix[p-1], pix[top-1]), avg2(pix[top+3], pix[top+7]))
|
||||
|
||||
case 11: // Select(L, T, TL).
|
||||
l0 := int32(pix[p-4])
|
||||
l1 := int32(pix[p-3])
|
||||
l2 := int32(pix[p-2])
|
||||
l3 := int32(pix[p-1])
|
||||
c0 := int32(pix[top-4])
|
||||
c1 := int32(pix[top-3])
|
||||
c2 := int32(pix[top-2])
|
||||
c3 := int32(pix[top-1])
|
||||
t0 := int32(pix[top+0])
|
||||
t1 := int32(pix[top+1])
|
||||
t2 := int32(pix[top+2])
|
||||
t3 := int32(pix[top+3])
|
||||
l := abs(c0-t0) + abs(c1-t1) + abs(c2-t2) + abs(c3-t3)
|
||||
t := abs(c0-l0) + abs(c1-l1) + abs(c2-l2) + abs(c3-l3)
|
||||
if l < t {
|
||||
pix[p+0] += uint8(l0)
|
||||
pix[p+1] += uint8(l1)
|
||||
pix[p+2] += uint8(l2)
|
||||
pix[p+3] += uint8(l3)
|
||||
} else {
|
||||
pix[p+0] += uint8(t0)
|
||||
pix[p+1] += uint8(t1)
|
||||
pix[p+2] += uint8(t2)
|
||||
pix[p+3] += uint8(t3)
|
||||
}
|
||||
|
||||
case 12: // ClampAddSubtractFull(L, T, TL).
|
||||
pix[p+0] += clampAddSubtractFull(pix[p-4], pix[top+0], pix[top-4])
|
||||
pix[p+1] += clampAddSubtractFull(pix[p-3], pix[top+1], pix[top-3])
|
||||
pix[p+2] += clampAddSubtractFull(pix[p-2], pix[top+2], pix[top-2])
|
||||
pix[p+3] += clampAddSubtractFull(pix[p-1], pix[top+3], pix[top-1])
|
||||
|
||||
case 13: // ClampAddSubtractHalf(Average2(L, T), TL).
|
||||
pix[p+0] += clampAddSubtractHalf(avg2(pix[p-4], pix[top+0]), pix[top-4])
|
||||
pix[p+1] += clampAddSubtractHalf(avg2(pix[p-3], pix[top+1]), pix[top-3])
|
||||
pix[p+2] += clampAddSubtractHalf(avg2(pix[p-2], pix[top+2]), pix[top-2])
|
||||
pix[p+3] += clampAddSubtractHalf(avg2(pix[p-1], pix[top+3]), pix[top-1])
|
||||
}
|
||||
p, top = p+4, top+4
|
||||
}
|
||||
}
|
||||
return pix
|
||||
}
|
||||
|
||||
func inverseCrossColor(t *transform, pix []byte, h int32) []byte {
|
||||
var greenToRed, greenToBlue, redToBlue int32
|
||||
p, mask, tilesPerRow := int32(0), int32(1)<<t.bits-1, nTiles(t.oldWidth, t.bits)
|
||||
for y := int32(0); y < h; y++ {
|
||||
q := 4 * (y >> t.bits) * tilesPerRow
|
||||
for x := int32(0); x < t.oldWidth; x++ {
|
||||
if x&mask == 0 {
|
||||
redToBlue = int32(int8(t.pix[q+0]))
|
||||
greenToBlue = int32(int8(t.pix[q+1]))
|
||||
greenToRed = int32(int8(t.pix[q+2]))
|
||||
q += 4
|
||||
}
|
||||
red := pix[p+0]
|
||||
green := pix[p+1]
|
||||
blue := pix[p+2]
|
||||
red += uint8(uint32(greenToRed*int32(int8(green))) >> 5)
|
||||
blue += uint8(uint32(greenToBlue*int32(int8(green))) >> 5)
|
||||
blue += uint8(uint32(redToBlue*int32(int8(red))) >> 5)
|
||||
pix[p+0] = red
|
||||
pix[p+2] = blue
|
||||
p += 4
|
||||
}
|
||||
}
|
||||
return pix
|
||||
}
|
||||
|
||||
func inverseSubtractGreen(t *transform, pix []byte, h int32) []byte {
|
||||
for p := 0; p < len(pix); p += 4 {
|
||||
green := pix[p+1]
|
||||
pix[p+0] += green
|
||||
pix[p+2] += green
|
||||
}
|
||||
return pix
|
||||
}
|
||||
|
||||
func inverseColorIndexing(t *transform, pix []byte, h int32) []byte {
|
||||
if t.bits == 0 {
|
||||
for p := 0; p < len(pix); p += 4 {
|
||||
i := 4 * uint32(pix[p+1])
|
||||
pix[p+0] = t.pix[i+0]
|
||||
pix[p+1] = t.pix[i+1]
|
||||
pix[p+2] = t.pix[i+2]
|
||||
pix[p+3] = t.pix[i+3]
|
||||
}
|
||||
return pix
|
||||
}
|
||||
|
||||
vMask, xMask, bitsPerPixel := uint32(0), int32(0), uint32(8>>t.bits)
|
||||
switch t.bits {
|
||||
case 1:
|
||||
vMask, xMask = 0x0f, 0x01
|
||||
case 2:
|
||||
vMask, xMask = 0x03, 0x03
|
||||
case 3:
|
||||
vMask, xMask = 0x01, 0x07
|
||||
}
|
||||
|
||||
d, p, v, dst := 0, 0, uint32(0), make([]byte, 4*t.oldWidth*h)
|
||||
for y := int32(0); y < h; y++ {
|
||||
for x := int32(0); x < t.oldWidth; x++ {
|
||||
if x&xMask == 0 {
|
||||
v = uint32(pix[p+1])
|
||||
p += 4
|
||||
}
|
||||
|
||||
i := 4 * (v & vMask)
|
||||
dst[d+0] = t.pix[i+0]
|
||||
dst[d+1] = t.pix[i+1]
|
||||
dst[d+2] = t.pix[i+2]
|
||||
dst[d+3] = t.pix[i+3]
|
||||
d += 4
|
||||
|
||||
v >>= bitsPerPixel
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
func abs(x int32) int32 {
|
||||
if x < 0 {
|
||||
return -x
|
||||
}
|
||||
return x
|
||||
}
|
||||
|
||||
func avg2(a, b uint8) uint8 {
|
||||
return uint8((int32(a) + int32(b)) / 2)
|
||||
}
|
||||
|
||||
func clampAddSubtractFull(a, b, c uint8) uint8 {
|
||||
x := int32(a) + int32(b) - int32(c)
|
||||
if x < 0 {
|
||||
return 0
|
||||
}
|
||||
if x > 255 {
|
||||
return 255
|
||||
}
|
||||
return uint8(x)
|
||||
}
|
||||
|
||||
func clampAddSubtractHalf(a, b uint8) uint8 {
|
||||
x := int32(a) + (int32(a)-int32(b))/2
|
||||
if x < 0 {
|
||||
return 0
|
||||
}
|
||||
if x > 255 {
|
||||
return 255
|
||||
}
|
||||
return uint8(x)
|
||||
}
|
||||
+282
@@ -0,0 +1,282 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package webp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"image"
|
||||
"image/color"
|
||||
"io"
|
||||
|
||||
"golang.org/x/image/riff"
|
||||
"golang.org/x/image/vp8"
|
||||
"golang.org/x/image/vp8l"
|
||||
)
|
||||
|
||||
var errInvalidFormat = errors.New("webp: invalid format")
|
||||
|
||||
var (
|
||||
fccALPH = riff.FourCC{'A', 'L', 'P', 'H'}
|
||||
fccVP8 = riff.FourCC{'V', 'P', '8', ' '}
|
||||
fccVP8L = riff.FourCC{'V', 'P', '8', 'L'}
|
||||
fccVP8X = riff.FourCC{'V', 'P', '8', 'X'}
|
||||
fccWEBP = riff.FourCC{'W', 'E', 'B', 'P'}
|
||||
)
|
||||
|
||||
func decode(r io.Reader, configOnly bool) (image.Image, image.Config, error) {
|
||||
formType, riffReader, err := riff.NewReader(r)
|
||||
if err != nil {
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
if formType != fccWEBP {
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
|
||||
var (
|
||||
alpha []byte
|
||||
alphaStride int
|
||||
wantAlpha bool
|
||||
seenVP8X bool
|
||||
widthMinusOne uint32
|
||||
heightMinusOne uint32
|
||||
buf [10]byte
|
||||
)
|
||||
for {
|
||||
chunkID, chunkLen, chunkData, err := riffReader.Next()
|
||||
if err == io.EOF {
|
||||
err = errInvalidFormat
|
||||
}
|
||||
if err != nil {
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
|
||||
switch chunkID {
|
||||
case fccALPH:
|
||||
if !wantAlpha {
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
wantAlpha = false
|
||||
// Read the Pre-processing | Filter | Compression byte.
|
||||
if _, err := io.ReadFull(chunkData, buf[:1]); err != nil {
|
||||
if err == io.EOF {
|
||||
err = errInvalidFormat
|
||||
}
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
alpha, alphaStride, err = readAlpha(chunkData, widthMinusOne, heightMinusOne, buf[0]&0x03)
|
||||
if err != nil {
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
unfilterAlpha(alpha, alphaStride, (buf[0]>>2)&0x03)
|
||||
|
||||
case fccVP8:
|
||||
if wantAlpha || int32(chunkLen) < 0 {
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
d := vp8.NewDecoder()
|
||||
d.Init(chunkData, int(chunkLen))
|
||||
fh, err := d.DecodeFrameHeader()
|
||||
if err != nil {
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
if configOnly {
|
||||
return nil, image.Config{
|
||||
ColorModel: color.YCbCrModel,
|
||||
Width: fh.Width,
|
||||
Height: fh.Height,
|
||||
}, nil
|
||||
}
|
||||
m, err := d.DecodeFrame()
|
||||
if err != nil {
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
if alpha != nil {
|
||||
return &image.NYCbCrA{
|
||||
YCbCr: *m,
|
||||
A: alpha,
|
||||
AStride: alphaStride,
|
||||
}, image.Config{}, nil
|
||||
}
|
||||
return m, image.Config{}, nil
|
||||
|
||||
case fccVP8L:
|
||||
if alpha != nil {
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
if configOnly {
|
||||
c, err := vp8l.DecodeConfig(chunkData)
|
||||
return nil, c, err
|
||||
}
|
||||
m, err := vp8l.Decode(chunkData)
|
||||
return m, image.Config{}, err
|
||||
|
||||
case fccVP8X:
|
||||
if seenVP8X {
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
seenVP8X = true
|
||||
if chunkLen != 10 {
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
if _, err := io.ReadFull(chunkData, buf[:10]); err != nil {
|
||||
return nil, image.Config{}, err
|
||||
}
|
||||
const (
|
||||
animationBit = 1 << 1
|
||||
xmpMetadataBit = 1 << 2
|
||||
exifMetadataBit = 1 << 3
|
||||
alphaBit = 1 << 4
|
||||
iccProfileBit = 1 << 5
|
||||
)
|
||||
wantAlpha = (buf[0] & alphaBit) != 0
|
||||
widthMinusOne = uint32(buf[4]) | uint32(buf[5])<<8 | uint32(buf[6])<<16
|
||||
heightMinusOne = uint32(buf[7]) | uint32(buf[8])<<8 | uint32(buf[9])<<16
|
||||
if uint64(widthMinusOne+1)*uint64(heightMinusOne+1) > 1<<32-1 {
|
||||
// The product of _Canvas Width_ and _Canvas Height_ MUST be
|
||||
// at most 2^32 - 1.
|
||||
// https://www.rfc-editor.org/rfc/rfc9649.html#section-2.7-12
|
||||
return nil, image.Config{}, errInvalidFormat
|
||||
}
|
||||
if configOnly {
|
||||
if wantAlpha {
|
||||
return nil, image.Config{
|
||||
ColorModel: color.NYCbCrAModel,
|
||||
Width: int(widthMinusOne) + 1,
|
||||
Height: int(heightMinusOne) + 1,
|
||||
}, nil
|
||||
}
|
||||
return nil, image.Config{
|
||||
ColorModel: color.YCbCrModel,
|
||||
Width: int(widthMinusOne) + 1,
|
||||
Height: int(heightMinusOne) + 1,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func readAlpha(chunkData io.Reader, widthMinusOne, heightMinusOne uint32, compression byte) (
|
||||
alpha []byte, alphaStride int, err error) {
|
||||
|
||||
switch compression {
|
||||
case 0:
|
||||
w := int(widthMinusOne) + 1
|
||||
h := int(heightMinusOne) + 1
|
||||
alpha = make([]byte, w*h)
|
||||
if _, err := io.ReadFull(chunkData, alpha); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return alpha, w, nil
|
||||
|
||||
case 1:
|
||||
// Read the VP8L-compressed alpha values. First, synthesize a 5-byte VP8L header:
|
||||
// a 1-byte magic number, a 14-bit widthMinusOne, a 14-bit heightMinusOne,
|
||||
// a 1-bit (ignored, zero) alphaIsUsed and a 3-bit (zero) version.
|
||||
// TODO(nigeltao): be more efficient than decoding an *image.NRGBA just to
|
||||
// extract the green values to a separately allocated []byte. Fixing this
|
||||
// will require changes to the vp8l package's API.
|
||||
if widthMinusOne > 0x3fff || heightMinusOne > 0x3fff {
|
||||
return nil, 0, errors.New("webp: invalid format")
|
||||
}
|
||||
alphaImage, err := vp8l.Decode(io.MultiReader(
|
||||
bytes.NewReader([]byte{
|
||||
0x2f, // VP8L magic number.
|
||||
uint8(widthMinusOne),
|
||||
uint8(widthMinusOne>>8) | uint8(heightMinusOne<<6),
|
||||
uint8(heightMinusOne >> 2),
|
||||
uint8(heightMinusOne >> 10),
|
||||
}),
|
||||
chunkData,
|
||||
))
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
// The green values of the inner NRGBA image are the alpha values of the
|
||||
// outer NYCbCrA image.
|
||||
pix := alphaImage.(*image.NRGBA).Pix
|
||||
alpha = make([]byte, len(pix)/4)
|
||||
for i := range alpha {
|
||||
alpha[i] = pix[4*i+1]
|
||||
}
|
||||
return alpha, int(widthMinusOne) + 1, nil
|
||||
}
|
||||
return nil, 0, errInvalidFormat
|
||||
}
|
||||
|
||||
func unfilterAlpha(alpha []byte, alphaStride int, filter byte) {
|
||||
if len(alpha) == 0 || alphaStride == 0 {
|
||||
return
|
||||
}
|
||||
switch filter {
|
||||
case 1: // Horizontal filter.
|
||||
for i := 1; i < alphaStride; i++ {
|
||||
alpha[i] += alpha[i-1]
|
||||
}
|
||||
for i := alphaStride; i < len(alpha); i += alphaStride {
|
||||
// The first column is equivalent to the vertical filter.
|
||||
alpha[i] += alpha[i-alphaStride]
|
||||
|
||||
for j := 1; j < alphaStride; j++ {
|
||||
alpha[i+j] += alpha[i+j-1]
|
||||
}
|
||||
}
|
||||
|
||||
case 2: // Vertical filter.
|
||||
// The first row is equivalent to the horizontal filter.
|
||||
for i := 1; i < alphaStride; i++ {
|
||||
alpha[i] += alpha[i-1]
|
||||
}
|
||||
|
||||
for i := alphaStride; i < len(alpha); i++ {
|
||||
alpha[i] += alpha[i-alphaStride]
|
||||
}
|
||||
|
||||
case 3: // Gradient filter.
|
||||
// The first row is equivalent to the horizontal filter.
|
||||
for i := 1; i < alphaStride; i++ {
|
||||
alpha[i] += alpha[i-1]
|
||||
}
|
||||
|
||||
for i := alphaStride; i < len(alpha); i += alphaStride {
|
||||
// The first column is equivalent to the vertical filter.
|
||||
alpha[i] += alpha[i-alphaStride]
|
||||
|
||||
// The interior is predicted on the three top/left pixels.
|
||||
for j := 1; j < alphaStride; j++ {
|
||||
c := int(alpha[i+j-alphaStride-1])
|
||||
b := int(alpha[i+j-alphaStride])
|
||||
a := int(alpha[i+j-1])
|
||||
x := a + b - c
|
||||
if x < 0 {
|
||||
x = 0
|
||||
} else if x > 255 {
|
||||
x = 255
|
||||
}
|
||||
alpha[i+j] += uint8(x)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Decode reads a WEBP image from r and returns it as an image.Image.
|
||||
func Decode(r io.Reader) (image.Image, error) {
|
||||
m, _, err := decode(r, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// DecodeConfig returns the color model and dimensions of a WEBP image without
|
||||
// decoding the entire image.
|
||||
func DecodeConfig(r io.Reader) (image.Config, error) {
|
||||
_, c, err := decode(r, true)
|
||||
return c, err
|
||||
}
|
||||
|
||||
func init() {
|
||||
image.RegisterFormat("webp", "RIFF????WEBPVP8", Decode, DecodeConfig)
|
||||
}
|
||||
+9
@@ -0,0 +1,9 @@
|
||||
// Copyright 2016 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package webp implements a decoder for WEBP images.
|
||||
//
|
||||
// WEBP is defined at:
|
||||
// https://developers.google.com/speed/webp/docs/riff_container
|
||||
package webp // import "golang.org/x/image/webp"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user