feat: add configurable SMTP HELO hostname (#2146)
* feat: add configurable SMTP HELO hostname Allow the SMTP HELO/EHLO hostname to be configured separately from the SMTP server hostname. This is useful when the SMTP server requires clients to identify themselves with a fully qualified hostname different from the server address. * chore: remove vendored dependency changes * Bump go-pkgz/notify to v1.4.0 and document SMTP_HELO_HOST The HELOHost field lands in go-pkgz/notify v1.4.0, so the branch needs the bump to compile; v1.3.0 in master has no such field. The example module is tidied alongside, as any change to backend/go.mod requires. Documents the parameter in the parameters table and, separately, in the email setup page: what it does, that leaving it unset keeps the previous `localhost` greeting, and the case it exists for, a relay refusing the greeting under Postfix `reject_non_fqdn_helo_hostname`. Also records the current limit: verification emails for email authentication go through go-pkgz/auth's own sender, which has no equivalent setting, so the greeting there is unchanged. * Bump go-pkgz/auth to v2.2.0 and apply SMTP_HELO_HOST to verification email The verification email sender had no way to set the greeting, so a relay that refuses the HELO would accept notifications and still reject sign-in emails. EmailParams gains HELOHost in go-pkgz/auth v2.2.0, so the same SMTP_HELO_HOST now drives both paths. The example module is tidied alongside, as any change to backend/go.mod requires. --------- Co-authored-by: oli <someone@somewhere.tld> Co-authored-by: Dmitry Verkhoturov <paskal.07@gmail.com>
This commit is contained in:
co-authored by
oli
Dmitry Verkhoturov
parent
43fccf3bc9
commit
3f5b3cdd98
@@ -4,7 +4,7 @@ go 1.25.0
|
||||
|
||||
require (
|
||||
github.com/go-pkgz/jrpc v0.4.0
|
||||
github.com/go-pkgz/lgr v0.12.3
|
||||
github.com/go-pkgz/lgr v0.12.4
|
||||
github.com/jessevdk/go-flags v1.6.1
|
||||
github.com/stretchr/testify v1.12.0
|
||||
github.com/umputun/remark42/backend v1.1000.0
|
||||
@@ -26,9 +26,9 @@ require (
|
||||
github.com/rs/xid v1.6.0 // indirect
|
||||
github.com/russross/blackfriday/v2 v2.1.0 // indirect
|
||||
go.etcd.io/bbolt v1.5.0 // indirect
|
||||
golang.org/x/crypto v0.53.0 // indirect
|
||||
golang.org/x/crypto v0.55.0 // indirect
|
||||
golang.org/x/image v0.45.0 // indirect
|
||||
golang.org/x/net v0.56.0 // indirect
|
||||
golang.org/x/net v0.57.0 // indirect
|
||||
golang.org/x/sys v0.47.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
)
|
||||
|
||||
@@ -17,8 +17,8 @@ github.com/dlclark/regexp2/v2 v2.2.2 h1:MYWvNYw8okuqNhwTYO587EZMiDruVa2vhV6fsGpf
|
||||
github.com/dlclark/regexp2/v2 v2.2.2/go.mod h1:avUrQvPaLz2DrFNHJF0taWAFFX2C1GMSSoeiqFjcBmU=
|
||||
github.com/go-pkgz/jrpc v0.4.0 h1:oD7xiGrzDkndkuCjeHGugQXxbggLSV7O1QmHhoc5pYY=
|
||||
github.com/go-pkgz/jrpc v0.4.0/go.mod h1:JFoY3bRjRyx4M3CbEVDFQStMB1m2gmQ7OjqFK7q3kOo=
|
||||
github.com/go-pkgz/lgr v0.12.3 h1:QDug7kRkEsuQtruT9fNF5PVT2kZUqCDPc4GmsgS3fP8=
|
||||
github.com/go-pkgz/lgr v0.12.3/go.mod h1:lpCDgVvCIxBHZp8+sGCj9MPctIzKZyZ3QdE19ddqd54=
|
||||
github.com/go-pkgz/lgr v0.12.4 h1:lDeQ4BR28ldXrKau6BOjq7A8nHzcXz+MF4xUfV4l1Ok=
|
||||
github.com/go-pkgz/lgr v0.12.4/go.mod h1:Lw6DkNRnCPyX07mqkiUK/p+eA1opq4GKkWfWia64RA8=
|
||||
github.com/go-pkgz/rest v1.24.0 h1:GAUCgx7U8xCOC2OynLjhCRMhtnMQH4d1mTdKpQyX2yI=
|
||||
github.com/go-pkgz/rest v1.24.0/go.mod h1:dl3EWiuFB4hRTo2Sknj6UrQGFRAYvANK6/NyW8qQPxc=
|
||||
github.com/go-pkgz/routegroup v1.6.0 h1:44XHZgF6JIIldRlv+zjg6SygULASmjifnfIQjwCT0e4=
|
||||
@@ -45,12 +45,12 @@ github.com/stretchr/testify v1.12.0 h1:K6Mr6jO9JICuend/5xzTM03ydSV3vdNRYAdPSukj8
|
||||
github.com/stretchr/testify v1.12.0/go.mod h1:bOYBZb5qJ00vPzWfIqBUZPaxK8jWiXc6d3ErP4Ca9Gw=
|
||||
go.etcd.io/bbolt v1.5.0 h1:S7GAl7Fxv12yohbwFfIbQCGDWbQbtDGPET4P/bD4lxU=
|
||||
go.etcd.io/bbolt v1.5.0/go.mod h1:mkltfYE5aUHQxUct9N9V+Kp7aSjFqjgrhcXIS70Lrdk=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
golang.org/x/image v0.45.0 h1:FMb1nTbH5H9vF55SriQHgFw5GnNL9Jg6L25BwXKzhB0=
|
||||
golang.org/x/image v0.45.0/go.mod h1:n62x/7RqlwXDvGsSU4u6IUTUf6KghUZ9Bt7cG/T9Fx4=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
|
||||
@@ -252,6 +252,7 @@ type TelegramGroup struct {
|
||||
type SMTPGroup struct {
|
||||
Host string `long:"host" env:"HOST" description:"SMTP host"`
|
||||
Port int `long:"port" env:"PORT" description:"SMTP port"`
|
||||
HELOHost string `long:"helo_host" env:"HELO_HOST" description:"SMTP HELO/EHLO hostname"`
|
||||
Username string `long:"username" env:"USERNAME" description:"SMTP user name"`
|
||||
Password string `long:"password" env:"PASSWORD" description:"SMTP password"`
|
||||
TLS bool `long:"tls" env:"TLS" description:"enable TLS"`
|
||||
@@ -1166,6 +1167,7 @@ func (s *ServerCommand) addAuthProviders(authenticator *auth.Service) error {
|
||||
params := sender.EmailParams{
|
||||
Host: s.SMTP.Host,
|
||||
Port: s.SMTP.Port,
|
||||
HELOHost: s.SMTP.HELOHost,
|
||||
SMTPUserName: s.SMTP.Username,
|
||||
SMTPPassword: s.SMTP.Password,
|
||||
TimeOut: s.SMTP.TimeOut,
|
||||
@@ -1322,6 +1324,7 @@ func (s *ServerCommand) makeNotifyDestinations(authenticator *auth.Service) ([]n
|
||||
smtpParams := ntf.SMTPParams{
|
||||
Host: s.SMTP.Host,
|
||||
Port: s.SMTP.Port,
|
||||
HELOHost: s.SMTP.HELOHost,
|
||||
TLS: s.SMTP.TLS,
|
||||
StartTLS: s.SMTP.StartTLS,
|
||||
InsecureSkipVerify: s.SMTP.InsecureSkipVerify,
|
||||
|
||||
+8
-9
@@ -7,11 +7,11 @@ require (
|
||||
github.com/PuerkitoBio/goquery v1.12.0
|
||||
github.com/alecthomas/chroma/v2 v2.27.0
|
||||
github.com/didip/tollbooth/v8 v8.0.1
|
||||
github.com/go-pkgz/auth/v2 v2.1.5
|
||||
github.com/go-pkgz/auth/v2 v2.2.0
|
||||
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
|
||||
github.com/go-pkgz/notify v1.3.0
|
||||
github.com/go-pkgz/lgr v0.12.4
|
||||
github.com/go-pkgz/notify v1.4.0
|
||||
github.com/go-pkgz/repeater/v2 v2.2.0
|
||||
github.com/go-pkgz/rest v1.24.0
|
||||
github.com/go-pkgz/routegroup v1.6.0
|
||||
@@ -28,9 +28,9 @@ require (
|
||||
github.com/stretchr/testify v1.12.0
|
||||
go.etcd.io/bbolt v1.5.0
|
||||
go.uber.org/goleak v1.3.0
|
||||
golang.org/x/crypto v0.53.0
|
||||
golang.org/x/crypto v0.55.0
|
||||
golang.org/x/image v0.45.0
|
||||
golang.org/x/net v0.56.0
|
||||
golang.org/x/net v0.57.0
|
||||
golang.org/x/oauth2 v0.36.0
|
||||
)
|
||||
|
||||
@@ -42,17 +42,16 @@ require (
|
||||
github.com/dghubble/oauth1 v0.7.3 // indirect
|
||||
github.com/dlclark/regexp2/v2 v2.2.2 // indirect
|
||||
github.com/go-oauth2/oauth2/v4 v4.5.4 // indirect
|
||||
github.com/go-pkgz/email v0.6.0 // indirect
|
||||
github.com/go-pkgz/email v0.8.0 // indirect
|
||||
github.com/go-pkgz/expirable-cache/v3 v3.1.0 // indirect
|
||||
github.com/go-pkgz/repeater v1.2.0 // indirect
|
||||
github.com/golang/snappy v1.0.0 // indirect
|
||||
github.com/gorilla/css v1.0.1 // indirect
|
||||
github.com/gorilla/websocket v1.5.3 // indirect
|
||||
github.com/hashicorp/errwrap v1.1.0 // indirect
|
||||
github.com/hashicorp/go-multierror v1.1.1 // indirect
|
||||
github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect
|
||||
github.com/klauspost/compress v1.18.7 // indirect
|
||||
github.com/montanaflynn/stats v0.9.0 // indirect
|
||||
github.com/klauspost/compress v1.19.2 // indirect
|
||||
github.com/montanaflynn/stats v0.12.4 // indirect
|
||||
github.com/redis/go-redis/v9 v9.21.0 // indirect
|
||||
github.com/rrivera/identicon v0.0.0-20240116195454-d5ba35832c0d // indirect
|
||||
github.com/slack-go/slack v0.27.0 // indirect
|
||||
|
||||
+18
-20
@@ -42,22 +42,20 @@ github.com/gavv/httpexpect v2.0.0+incompatible h1:1X9kcRshkSKEjNJJxX9Y9mQ5BRfbxU
|
||||
github.com/gavv/httpexpect v2.0.0+incompatible/go.mod h1:x+9tiU1YnrOvnB725RkpoLv1M62hOWzwo5OXotisrKc=
|
||||
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.5 h1:CFL7XxRMNPga0S0YCnAnlvO61OHHEYvVEGrIZXuA98Y=
|
||||
github.com/go-pkgz/auth/v2 v2.1.5/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/auth/v2 v2.2.0 h1:vQO+GTFDjAaBNSdcLLLr3Xibka67GZPtkRRItVVS4ow=
|
||||
github.com/go-pkgz/auth/v2 v2.2.0/go.mod h1:iZx2JiGZ8Aef+wM0BPLMQY8aur4fLE0uyPhFRO9dYQ4=
|
||||
github.com/go-pkgz/email v0.8.0 h1:6+Tgjfj7zFccFCPmURV2spKDXDb7aX/iWXBY2hBa9ww=
|
||||
github.com/go-pkgz/email v0.8.0/go.mod h1:+wgi4x7S33IuCzfcCM5euN0GwQG6XvO/PBLxrNffYLI=
|
||||
github.com/go-pkgz/expirable-cache/v3 v3.1.0 h1:s05P851/O6QJ6Mc+7o2bh9aGtD3romB1SxDTXifdoqc=
|
||||
github.com/go-pkgz/expirable-cache/v3 v3.1.0/go.mod h1:6pVgNleydKPj0J2/mzrI02/RDo4ivKx5v2XlNmIjhjo=
|
||||
github.com/go-pkgz/jrpc v0.4.0 h1:oD7xiGrzDkndkuCjeHGugQXxbggLSV7O1QmHhoc5pYY=
|
||||
github.com/go-pkgz/jrpc v0.4.0/go.mod h1:JFoY3bRjRyx4M3CbEVDFQStMB1m2gmQ7OjqFK7q3kOo=
|
||||
github.com/go-pkgz/lcw/v2 v2.0.0 h1:gTwXpiJBhQeA1rXuqkRuLcV79uATFna8CckH8ZBBrH0=
|
||||
github.com/go-pkgz/lcw/v2 v2.0.0/go.mod h1:yxJHOn+IbQBQHxUqkCtMrbGjIfdYcsBAZcVCBaL1Va8=
|
||||
github.com/go-pkgz/lgr v0.12.3 h1:QDug7kRkEsuQtruT9fNF5PVT2kZUqCDPc4GmsgS3fP8=
|
||||
github.com/go-pkgz/lgr v0.12.3/go.mod h1:lpCDgVvCIxBHZp8+sGCj9MPctIzKZyZ3QdE19ddqd54=
|
||||
github.com/go-pkgz/notify v1.3.0 h1:YxF/ThEoCetdcoghWdyeqaBpCkZ8mvyve7HXbCAOzYU=
|
||||
github.com/go-pkgz/notify v1.3.0/go.mod h1:qdfi5OsViKlIFPryIOaINHTOtS9GFhOYXPqJmAMlaGU=
|
||||
github.com/go-pkgz/repeater v1.2.0 h1:oJFvjyKdTDd5RCzpzxlzYIZFFj6Zfl17rE1aUfu6UjQ=
|
||||
github.com/go-pkgz/repeater v1.2.0/go.mod h1:vypP6xamA53MFmafnGUucqOmALKk36xgKu2hSG73LHM=
|
||||
github.com/go-pkgz/lgr v0.12.4 h1:lDeQ4BR28ldXrKau6BOjq7A8nHzcXz+MF4xUfV4l1Ok=
|
||||
github.com/go-pkgz/lgr v0.12.4/go.mod h1:Lw6DkNRnCPyX07mqkiUK/p+eA1opq4GKkWfWia64RA8=
|
||||
github.com/go-pkgz/notify v1.4.0 h1:4pP7UGdYqFO7e7V3OsQStYF006CO0cCh1ahdawt6l18=
|
||||
github.com/go-pkgz/notify v1.4.0/go.mod h1:UFpL9ZvCYnLBEjeay++3afh8GceZ6qT8wj5lqfSR9U4=
|
||||
github.com/go-pkgz/repeater/v2 v2.2.0 h1:8nZR/NaknmLfx2YMHbr78u9OL4Xj+8+romm9dz4FpMg=
|
||||
github.com/go-pkgz/repeater/v2 v2.2.0/go.mod h1:RgX5vUbLKq7PV82QUDP5pFbQS1os4Z+U9XzKymK23A8=
|
||||
github.com/go-pkgz/rest v1.24.0 h1:GAUCgx7U8xCOC2OynLjhCRMhtnMQH4d1mTdKpQyX2yI=
|
||||
@@ -101,8 +99,8 @@ 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.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw=
|
||||
github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||
github.com/klauspost/compress v1.19.2 h1:hMRETovs/pu/dVWN7zIT1PGG8t509MwT6bO7XSi26R8=
|
||||
github.com/klauspost/compress v1.19.2/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=
|
||||
@@ -113,8 +111,8 @@ github.com/kyokomi/emoji/v2 v2.2.13 h1:GhTfQa67venUUvmleTNFnb+bi7S3aocF7ZCXU9fSO
|
||||
github.com/kyokomi/emoji/v2 v2.2.13/go.mod h1:JUcn42DTdsXJo1SWanHh4HKDEyPaR5CqkmoirZZP9qE=
|
||||
github.com/microcosm-cc/bluemonday v1.0.27 h1:MpEUotklkwCSLeH+Qdx1VJgNqLlpY2KXwXFM08ygZfk=
|
||||
github.com/microcosm-cc/bluemonday v1.0.27/go.mod h1:jFi9vgW+H7c3V0lb6nR74Ib/DIB5OBs92Dimizgw2cA=
|
||||
github.com/montanaflynn/stats v0.9.0 h1:tsBJ0RXwph9BmAuFoCmqGv6e8xa0MENQ8m0ptKq29mQ=
|
||||
github.com/montanaflynn/stats v0.9.0/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||
github.com/montanaflynn/stats v0.12.4 h1:amtNRsti20yIhcrkfUJGwoYqBR82jKQFE8SNNYVgGn0=
|
||||
github.com/montanaflynn/stats v0.12.4/go.mod h1:etXPPgVO6n31NxCd9KQUMvCM+ve0ruNzt6R8Bnaayow=
|
||||
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/redis/go-redis/v9 v9.21.0 h1:FPBE4hhbAke+TLmcY3WkpbDffJEomdqPn3HYiqAtL9E=
|
||||
@@ -143,8 +141,8 @@ github.com/tidwall/btree v1.8.1 h1:27ehoXvm5AG/g+1VxLS1SD3vRhp/H7LuEfwNvddEdmA=
|
||||
github.com/tidwall/btree v1.8.1/go.mod h1:jBbTdUWhSZClZWoDg54VnvV7/54modSOzDN7VXftj1A=
|
||||
github.com/tidwall/buntdb v1.3.2 h1:qd+IpdEGs0pZci37G4jF51+fSKlkuUTMXuHhXL1AkKg=
|
||||
github.com/tidwall/buntdb v1.3.2/go.mod h1:lZZrZUWzlyDJKlLQ6DKAy53LnG7m5kHyrEHvvcDmBpU=
|
||||
github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
|
||||
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
|
||||
github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU=
|
||||
github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc=
|
||||
github.com/tidwall/grect v0.1.4 h1:dA3oIgNgWdSspFzn1kS4S/RDpZFLrIxAZOdJKjYapOg=
|
||||
github.com/tidwall/grect v0.1.4/go.mod h1:9FBsaYRaR0Tcy4UwefBX/UDcDcDy9V5jUcxHzv2jd5Q=
|
||||
github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM=
|
||||
@@ -194,16 +192,16 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
golang.org/x/image v0.45.0 h1:FMb1nTbH5H9vF55SriQHgFw5GnNL9Jg6L25BwXKzhB0=
|
||||
golang.org/x/image v0.45.0/go.mod h1:n62x/7RqlwXDvGsSU4u6IUTUf6KghUZ9Bt7cG/T9Fx4=
|
||||
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
|
||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||
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=
|
||||
|
||||
+26
@@ -237,6 +237,10 @@ func (s *Service) Handlers() (authHandler, avatarHandler http.Handler) {
|
||||
p.Handler(w, r)
|
||||
}
|
||||
|
||||
if s.avatarProxy == nil { // no avatar store configured, avatar route has nothing to serve
|
||||
return withSecurityHeaders(http.HandlerFunc(ah)), withSecurityHeaders(http.NotFoundHandler())
|
||||
}
|
||||
|
||||
return withSecurityHeaders(http.HandlerFunc(ah)), withSecurityHeaders(http.HandlerFunc(s.avatarProxy.Handler))
|
||||
}
|
||||
|
||||
@@ -383,6 +387,28 @@ func (s *Service) AddMicrosoftProvider(cid, csecret, tenant string) {
|
||||
s.addProvider(provider.NewMicrosoft(p))
|
||||
}
|
||||
|
||||
// AddGithubProviderWithNumericID adds github provider deriving the user id from the immutable
|
||||
// numeric account id instead of the login. Logins are released on rename or account removal and
|
||||
// can be claimed by someone else, so an id derived from login may be inherited by the next holder
|
||||
// of the name. This changes the id of every existing github user, see README for the migration note.
|
||||
// If the response carries no usable numeric id the login-derived id is kept.
|
||||
// For advanced configuration (e.g., UserAttributes), construct provider.Params directly.
|
||||
func (s *Service) AddGithubProviderWithNumericID(cid, csecret string) {
|
||||
p := provider.Params{
|
||||
URL: s.opts.URL,
|
||||
JwtService: s.jwtService,
|
||||
Issuer: s.issuer,
|
||||
AvatarSaver: s.avatarProxy,
|
||||
Cid: cid,
|
||||
Csecret: csecret,
|
||||
L: s.logger,
|
||||
UserAttributes: map[string]string{},
|
||||
GithubNumericID: true,
|
||||
AllowedRedirectHosts: s.opts.AllowedRedirectHosts,
|
||||
}
|
||||
s.addProvider(provider.NewGithub(p))
|
||||
}
|
||||
|
||||
// AddDevProvider with a custom host and port
|
||||
func (s *Service) AddDevProvider(host string, port int) {
|
||||
p := provider.Params{
|
||||
|
||||
+74
-16
@@ -43,8 +43,67 @@ func (gf *GridFS) Put(userID string, reader io.Reader) (avatar string, err error
|
||||
}
|
||||
|
||||
avaHash := hash(buf.Bytes(), id)
|
||||
_, err = bucket.UploadFromStream(id+imgSfx, buf, &options.UploadOptions{Metadata: bson.M{"hash": avaHash}})
|
||||
return id + imgSfx, err
|
||||
fileID, err := bucket.UploadFromStream(id+imgSfx, buf, &options.UploadOptions{Metadata: bson.M{"hash": avaHash}})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// gridfs turns every upload with the same name into another revision, so drop what
|
||||
// earlier Put calls left behind. Cleanup runs after the upload, keeping the previous
|
||||
// avatar in place if the write failed, and is best-effort as the new avatar is stored
|
||||
// either way.
|
||||
_ = gf.removeOlderRevisions(bucket, id+imgSfx, fileID)
|
||||
|
||||
return id + imgSfx, nil
|
||||
}
|
||||
|
||||
// removeOlderRevisions deletes the revisions of fileName stored before keepID. Only older
|
||||
// revisions go, so two concurrent Put calls cannot delete each other's upload and leave the
|
||||
// avatar missing; the newer of the two survives.
|
||||
func (gf *GridFS) removeOlderRevisions(bucket *gridfs.Bucket, fileName string, keepID primitive.ObjectID) error {
|
||||
ids, err := gf.revisionIDs(bucket, fileName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
older := false
|
||||
for _, id := range ids { // newest first, everything past keepID was uploaded earlier
|
||||
if id == keepID {
|
||||
older = true
|
||||
continue
|
||||
}
|
||||
if !older {
|
||||
continue
|
||||
}
|
||||
if e := bucket.Delete(id); e != nil {
|
||||
err = e
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// revisionIDs returns ids of all gridfs files stored under the given name, newest first
|
||||
func (gf *GridFS) revisionIDs(bucket *gridfs.Bucket, fileName string) ([]primitive.ObjectID, error) {
|
||||
sortNewestFirst := options.GridFSFind().SetSort(bson.D{{Key: "uploadDate", Value: -1}, {Key: "_id", Value: -1}})
|
||||
cursor, err := bucket.Find(bson.M{"filename": fileName}, sortNewestFirst)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), gf.timeout)
|
||||
defer cancel()
|
||||
|
||||
var ids []primitive.ObjectID
|
||||
for cursor.Next(ctx) {
|
||||
r := struct {
|
||||
ID primitive.ObjectID `bson:"_id"`
|
||||
}{}
|
||||
if err = cursor.Decode(&r); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ids = append(ids, r.ID)
|
||||
}
|
||||
return ids, cursor.Err()
|
||||
}
|
||||
|
||||
// Get avatar reader for avatar id.image
|
||||
@@ -78,7 +137,8 @@ func (gf *GridFS) ID(avatar string) (id string) {
|
||||
if err != nil {
|
||||
return encodeID(avatar)
|
||||
}
|
||||
cursor, err := bucket.Find(bson.M{"filename": avatar})
|
||||
sortNewestFirst := options.GridFSFind().SetSort(bson.D{{Key: "uploadDate", Value: -1}, {Key: "_id", Value: -1}})
|
||||
cursor, err := bucket.Find(bson.M{"filename": avatar}, sortNewestFirst)
|
||||
if err != nil {
|
||||
return encodeID(avatar)
|
||||
}
|
||||
@@ -100,23 +160,21 @@ func (gf *GridFS) Remove(avatar string) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cursor, err := bucket.Find(bson.M{"filename": avatar})
|
||||
ids, err := gf.revisionIDs(bucket, avatar)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
r := struct {
|
||||
ID primitive.ObjectID `bson:"_id"`
|
||||
}{}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), gf.timeout)
|
||||
defer cancel()
|
||||
if found := cursor.Next(ctx); found {
|
||||
if err := cursor.Decode(&r); err != nil {
|
||||
return err
|
||||
}
|
||||
return bucket.Delete(r.ID)
|
||||
if len(ids) == 0 {
|
||||
return fmt.Errorf("avatar %s not found: %w", avatar, ErrNotFound)
|
||||
}
|
||||
return fmt.Errorf("avatar %s not found: %w", avatar, ErrNotFound)
|
||||
|
||||
// every revision has to go, deleting the newest one alone leaves the avatar readable
|
||||
for _, id := range ids {
|
||||
if e := bucket.Delete(id); e != nil {
|
||||
err = e
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// List all avatars (ids) on gfs
|
||||
|
||||
+1
-1
@@ -72,7 +72,7 @@ func (fs *LocalFS) Get(avatar string) (reader io.ReadCloser, size int, err error
|
||||
func (fs *LocalFS) ID(avatar string) (id string) {
|
||||
location := fs.location(strings.TrimSuffix(avatar, imgSfx))
|
||||
avFile := path.Join(location, avatar)
|
||||
fi, err := os.Stat(avFile)
|
||||
fi, err := os.Stat(avFile) //nolint:gosec // avatar id is store-generated and validated by the proxy handler
|
||||
if err != nil {
|
||||
return encodeID(avatar)
|
||||
}
|
||||
|
||||
+23
-26
@@ -92,8 +92,9 @@ type AppleHandler struct {
|
||||
// infoURL string not implemented at Apple side
|
||||
endpoint oauth2.Endpoint
|
||||
|
||||
mapUser func(jwt.MapClaims) token.User // map info from InfoURL to User
|
||||
conf AppleConfig // main config for Apple auth provider
|
||||
mapUser func(jwt.MapClaims) token.User // map info from InfoURL to User
|
||||
conf AppleConfig // main config for Apple auth provider
|
||||
jwkCache *appleJWKCache // shared cache of Apple public keys
|
||||
|
||||
PrivateKeyLoader PrivateKeyLoaderInterface // custom function interface for load private key
|
||||
|
||||
@@ -181,6 +182,8 @@ func NewApple(p Params, appleCfg AppleConfig, privateKeyLoader PrivateKeyLoaderI
|
||||
TokenURL: appleTokenURL,
|
||||
},
|
||||
|
||||
jwkCache: &appleJWKCache{},
|
||||
|
||||
mapUser: func(claims jwt.MapClaims) token.User {
|
||||
var usr token.User
|
||||
if uid, ok := claims["sub"]; ok {
|
||||
@@ -261,7 +264,7 @@ func (ah *AppleHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
State: state,
|
||||
From: r.URL.Query().Get("from"),
|
||||
},
|
||||
SessionOnly: r.URL.Query().Get("session") != "" && r.URL.Query().Get("session") != "0",
|
||||
SessionOnly: sessionOnlyFromRequest(r),
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ID: cid,
|
||||
Audience: jwt.ClaimStrings{r.URL.Query().Get("site")},
|
||||
@@ -288,11 +291,11 @@ func (ah *AppleHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
ah.Logf("[DEBUG] login url %s, claims=%+v", loginURL, claims)
|
||||
|
||||
http.Redirect(w, r, loginURL, http.StatusFound)
|
||||
http.Redirect(w, r, loginURL, http.StatusFound) //nolint:gosec // redirect goes to the fixed apple auth endpoint, request path only affects redirect_uri query param
|
||||
}
|
||||
|
||||
// AuthHandler fills user info and redirects to "from" url. This is callback url redirected locally by browser
|
||||
// GET /callback
|
||||
// AuthHandler fills user info and redirects to "from" url. This is callback url redirected locally by browser.
|
||||
// POST /callback with the default form_post response mode, GET /callback when response mode is overridden
|
||||
func (ah AppleHandler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
// read response form data
|
||||
@@ -338,19 +341,12 @@ func (ah AppleHandler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// trying to fetch Apple public key (JWK) for verify token signature, it need for verify IDToken received from Apple
|
||||
keySet, err := fetchAppleJWK(r.Context(), ah.conf.jwkURL)
|
||||
if err != nil {
|
||||
ah.Logf("[ERROR] failed to fetch JWK from Apple key service: " + err.Error())
|
||||
rest.SendErrorJSON(w, r, ah.L, http.StatusInternalServerError, nil, fmt.Sprintf("failed to fetch JWK from Apple key service: %s", resp.Error))
|
||||
return
|
||||
}
|
||||
|
||||
// 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
|
||||
// the signature is verified with Apple public keys (JWK), served from the handler cache,
|
||||
// while 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, ah.jwkKeyFunc(r.Context()),
|
||||
jwt.WithIssuer(appleIDTokenIssuer),
|
||||
jwt.WithAudience(ah.conf.ClientID))
|
||||
if err != nil {
|
||||
@@ -391,7 +387,7 @@ func (ah AppleHandler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
ID: cid,
|
||||
Audience: oauthClaims.Audience,
|
||||
},
|
||||
SessionOnly: false,
|
||||
SessionOnly: oauthClaims.SessionOnly,
|
||||
AuthProvider: &token.AuthProvider{
|
||||
Name: ah.name,
|
||||
},
|
||||
@@ -411,7 +407,9 @@ func (ah AppleHandler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
rest.RenderJSON(w, &u)
|
||||
return
|
||||
}
|
||||
http.Redirect(w, r, oauthClaims.Handshake.From, http.StatusTemporaryRedirect)
|
||||
// see-other makes the browser retrieve the target with GET, so apple's form_post
|
||||
// callback is not replayed as a POST onto the "from" page
|
||||
http.Redirect(w, r, oauthClaims.Handshake.From, http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
rest.RenderJSON(w, &u)
|
||||
@@ -461,18 +459,17 @@ func (ah *AppleHandler) exchange(ctx context.Context, code, redirectURI string,
|
||||
return err
|
||||
}
|
||||
|
||||
// trying to decode (unmarshal json) data of response
|
||||
err = json.NewDecoder(res.Body).Decode(result)
|
||||
if err != nil {
|
||||
return fmt.Errorf("unmarshalling data from apple service response failed: %w", err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
if err = res.Body.Close(); err != nil {
|
||||
ah.Logf("[ERROR] close request body failed when get access token: %v", err)
|
||||
if e := res.Body.Close(); e != nil {
|
||||
ah.Logf("[ERROR] close request body failed when get access token: %v", e)
|
||||
}
|
||||
}()
|
||||
|
||||
// trying to decode (unmarshal json) data of response
|
||||
if err = json.NewDecoder(res.Body).Decode(result); err != nil {
|
||||
return fmt.Errorf("unmarshalling data from apple service response failed: %w", err)
|
||||
}
|
||||
|
||||
// if above operation done successfully checking a response code and error descriptions, if one exist.
|
||||
// apple service will response either 200 (OK) or 400 (any error).
|
||||
if res.StatusCode != http.StatusOK || result.Error != "" {
|
||||
|
||||
+77
-1
@@ -14,6 +14,7 @@ import (
|
||||
"io"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
@@ -65,17 +66,25 @@ func fetchAppleJWK(ctx context.Context, keyURL string) (set appleKeySet, err err
|
||||
if err != nil {
|
||||
return set, fmt.Errorf("failed to fetch Apple public keys: %w", err)
|
||||
}
|
||||
defer func() { _ = res.Body.Close() }()
|
||||
|
||||
// an error body parses as a key set with no keys, caching it would break every login until it expires
|
||||
if res.StatusCode < http.StatusOK || res.StatusCode >= http.StatusMultipleChoices {
|
||||
return set, fmt.Errorf("failed to fetch Apple public keys, status %s", res.Status)
|
||||
}
|
||||
|
||||
data, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
return set, fmt.Errorf("failed read data after Apple public key fetched: %w", err)
|
||||
}
|
||||
defer func() { _ = res.Body.Close() }()
|
||||
|
||||
set, err = parseAppleJWK(data)
|
||||
if err != nil {
|
||||
return set, fmt.Errorf("get set of apple public key failed: %w", err)
|
||||
}
|
||||
if len(set.keys) == 0 {
|
||||
return appleKeySet{}, fmt.Errorf("no keys in Apple public key response")
|
||||
}
|
||||
|
||||
return set, nil
|
||||
}
|
||||
@@ -176,3 +185,70 @@ func (aks *appleKeySet) keyFunc(token *jwt.Token) (any, error) {
|
||||
|
||||
return key.publicKey, nil
|
||||
}
|
||||
|
||||
// appleJWKTTL is how long a successfully fetched key set is reused before a refresh
|
||||
const appleJWKTTL = time.Hour
|
||||
|
||||
// appleJWKStaleTTL bounds the use of a cached key set when the key service is unreachable
|
||||
const appleJWKStaleTTL = 12 * time.Hour
|
||||
|
||||
// appleJWKCache keeps the last fetched Apple key set. AppleHandler is copied by value on
|
||||
// every request, so the cache is held behind a pointer and shared by all copies.
|
||||
type appleJWKCache struct {
|
||||
lock sync.Mutex
|
||||
set appleKeySet
|
||||
fetchedAt time.Time
|
||||
}
|
||||
|
||||
// jwkSet returns a set of Apple public keys able to verify kid. The cached set is reused
|
||||
// until appleJWKTTL passes; an unknown kid forces a refresh to pick up Apple's key rotation.
|
||||
// If the refresh fails, a cached set holding the kid is used for up to appleJWKStaleTTL, so
|
||||
// logins survive a short outage of the key service.
|
||||
func (ah AppleHandler) jwkSet(ctx context.Context, kid string) (appleKeySet, error) {
|
||||
if ah.jwkCache == nil { // handler constructed without NewApple
|
||||
return fetchAppleJWK(ctx, ah.conf.jwkURL)
|
||||
}
|
||||
|
||||
ah.jwkCache.lock.Lock()
|
||||
defer ah.jwkCache.lock.Unlock()
|
||||
|
||||
cached := ah.jwkCache.set
|
||||
_, cachedErr := cached.get(kid)
|
||||
if cachedErr == nil && time.Since(ah.jwkCache.fetchedAt) < appleJWKTTL {
|
||||
return cached, nil
|
||||
}
|
||||
|
||||
set, err := fetchAppleJWK(ctx, ah.conf.jwkURL)
|
||||
if err != nil {
|
||||
if cachedErr == nil && time.Since(ah.jwkCache.fetchedAt) < appleJWKStaleTTL {
|
||||
ah.Logf("[WARN] failed to refresh Apple public keys, using cached set: %v", err)
|
||||
return cached, nil
|
||||
}
|
||||
return set, err
|
||||
}
|
||||
|
||||
ah.jwkCache.set, ah.jwkCache.fetchedAt = set, time.Now()
|
||||
return set, nil
|
||||
}
|
||||
|
||||
// jwkKeyFunc verifies an Apple id_token signature with the key named by the token's kid header
|
||||
func (ah AppleHandler) jwkKeyFunc(ctx context.Context) jwt.Keyfunc {
|
||||
return func(jwtToken *jwt.Token) (any, error) {
|
||||
kid, ok := jwtToken.Header["kid"].(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("get JWT kid header not found")
|
||||
}
|
||||
|
||||
set, err := ah.jwkSet(ctx, kid)
|
||||
if err != nil {
|
||||
ah.Logf("[ERROR] failed to fetch JWK from Apple key service: %v", err)
|
||||
return nil, fmt.Errorf("failed to fetch JWK from Apple key service: %w", err)
|
||||
}
|
||||
|
||||
key, err := set.get(kid)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return key.publicKey, nil
|
||||
}
|
||||
}
|
||||
|
||||
+2
-1
@@ -193,7 +193,8 @@ func (c *CustomServer) handleAvatar(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
if _, err = w.Write(b); err != nil {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
if _, err = w.Write(b); err != nil { //nolint:gosec // generated identicon png, not reflected markup
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
+2
-1
@@ -130,7 +130,8 @@ func (d *DevAuthServer) Run(ctx context.Context) { // nolint (gocyclo)
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
if _, err = w.Write(b); err != nil {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
if _, err = w.Write(b); err != nil { //nolint:gosec // generated identicon png, not reflected markup
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
+3
-3
@@ -60,9 +60,9 @@ func (p DirectHandler) Name() string { return p.ProviderName }
|
||||
|
||||
// LoginHandler checks "user" and "passwd" against data store and makes jwt if all passed.
|
||||
//
|
||||
// GET /something?user=name&passwd=xyz&aud=bar&sess=[0|1]
|
||||
// GET /something?user=name&passwd=xyz&aud=bar&session=[0|1]
|
||||
//
|
||||
// POST /something?sess[0|1]
|
||||
// POST /something?session=[0|1]
|
||||
// Accepts application/x-www-form-urlencoded or application/json encoded requests.
|
||||
//
|
||||
// application/x-www-form-urlencoded body example:
|
||||
@@ -82,7 +82,7 @@ func (p DirectHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
rest.SendErrorJSON(w, logReq, p.L, http.StatusBadRequest, err, "failed to parse credentials")
|
||||
return
|
||||
}
|
||||
sessOnly := r.URL.Query().Get("sess") == "1"
|
||||
sessOnly := sessionOnlyFromRequest(r)
|
||||
if p.CredChecker == nil {
|
||||
rest.SendErrorJSON(w, logReq, p.L, http.StatusInternalServerError,
|
||||
fmt.Errorf("no credential checker"), "no credential checker")
|
||||
|
||||
+11
-2
@@ -3,6 +3,7 @@ package provider
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -54,7 +55,7 @@ func (h Oauth1Handler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
State: requestSecret,
|
||||
From: r.URL.Query().Get("from"),
|
||||
},
|
||||
SessionOnly: r.URL.Query().Get("session") != "" && r.URL.Query().Get("session") != "0",
|
||||
SessionOnly: sessionOnlyFromRequest(r),
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ID: cid,
|
||||
Audience: []string{r.URL.Query().Get("site")},
|
||||
@@ -116,6 +117,13 @@ func (h Oauth1Handler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}()
|
||||
|
||||
// an error body carries no identity fields, mapping it would hash empty values into a shared id
|
||||
if uinfo.StatusCode < http.StatusOK || uinfo.StatusCode >= http.StatusMultipleChoices {
|
||||
rest.SendErrorJSON(w, r, h.L, http.StatusServiceUnavailable,
|
||||
fmt.Errorf("status %s", uinfo.Status), "failed to get user info")
|
||||
return
|
||||
}
|
||||
|
||||
data, err := io.ReadAll(uinfo.Body)
|
||||
if err != nil {
|
||||
rest.SendErrorJSON(w, r, h.L, http.StatusInternalServerError, err, "failed to read user info")
|
||||
@@ -167,7 +175,8 @@ func (h Oauth1Handler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
rest.RenderJSON(w, &u)
|
||||
return
|
||||
}
|
||||
http.Redirect(w, r, oauthClaims.Handshake.From, http.StatusTemporaryRedirect)
|
||||
// see-other keeps the redirect a GET whatever method the callback arrived with
|
||||
http.Redirect(w, r, oauthClaims.Handshake.From, http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
rest.RenderJSON(w, &u)
|
||||
|
||||
+17
-2
@@ -54,6 +54,13 @@ type Params struct {
|
||||
Host string // relevant for providers supporting host customization, for example dev oauth2
|
||||
|
||||
MicrosoftTenant string // tenant for microsoft provider, default "common"
|
||||
|
||||
// GithubNumericID makes github provider derive the user id from the immutable numeric account id
|
||||
// instead of the login. Logins are released on rename or account removal and can be claimed by
|
||||
// someone else, so an id derived from login may be inherited by the next holder of the name.
|
||||
// Enabling this changes the id of every existing github user, see README for the migration note.
|
||||
// Best-effort: if the response carries no usable numeric id the login-derived id is kept.
|
||||
GithubNumericID bool
|
||||
}
|
||||
|
||||
// UserData is type for user information returned from oauth2 providers /info API method
|
||||
@@ -120,7 +127,7 @@ func (p Oauth2Handler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
State: state,
|
||||
From: r.URL.Query().Get("from"),
|
||||
},
|
||||
SessionOnly: r.URL.Query().Get("session") != "" && r.URL.Query().Get("session") != "0",
|
||||
SessionOnly: sessionOnlyFromRequest(r),
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ID: cid,
|
||||
Audience: []string{aud},
|
||||
@@ -191,6 +198,13 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}()
|
||||
|
||||
// an error body carries no identity fields, mapping it would hash empty values into a shared id
|
||||
if uinfo.StatusCode < http.StatusOK || uinfo.StatusCode >= http.StatusMultipleChoices {
|
||||
rest.SendErrorJSON(w, r, p.L, http.StatusServiceUnavailable,
|
||||
fmt.Errorf("status %s", uinfo.Status), "failed to get user info")
|
||||
return
|
||||
}
|
||||
|
||||
data, err := io.ReadAll(uinfo.Body)
|
||||
if err != nil {
|
||||
rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to read user info")
|
||||
@@ -251,7 +265,8 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
rest.RenderJSON(w, &u)
|
||||
return
|
||||
}
|
||||
http.Redirect(w, r, oauthClaims.Handshake.From, http.StatusTemporaryRedirect)
|
||||
// see-other keeps the redirect a GET whatever method the callback arrived with
|
||||
http.Redirect(w, r, oauthClaims.Handshake.From, http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
rest.RenderJSON(w, &u)
|
||||
|
||||
+22
-4
@@ -5,10 +5,12 @@ import (
|
||||
"crypto/sha1" //nolint
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/dghubble/oauth1"
|
||||
"github.com/dghubble/oauth1/twitter"
|
||||
"github.com/go-pkgz/auth/v2/logger"
|
||||
"github.com/go-pkgz/auth/v2/token"
|
||||
"golang.org/x/oauth2"
|
||||
"golang.org/x/oauth2/facebook"
|
||||
@@ -48,17 +50,33 @@ func NewGoogle(p Params) Oauth2Handler {
|
||||
|
||||
// NewGithub makes github oauth2 provider
|
||||
func NewGithub(p Params) Oauth2Handler {
|
||||
if p.L == nil {
|
||||
p.L = logger.NoOp // mapUser below captures p, initOauth2Handler defaults its own copy only
|
||||
}
|
||||
return initOauth2Handler(p, Oauth2Handler{
|
||||
name: "github",
|
||||
endpoint: github.Endpoint,
|
||||
scopes: []string{},
|
||||
infoURL: "https://api.github.com/user",
|
||||
mapUser: func(data UserData, _ []byte) token.User {
|
||||
mapUser: func(data UserData, bdata []byte) token.User {
|
||||
userInfo := token.User{
|
||||
ID: "github_" + token.HashID(sha1.New(), data.Value("login")),
|
||||
Name: data.Value("name"),
|
||||
Picture: data.Value("avatar_url"),
|
||||
}
|
||||
if p.GithubNumericID {
|
||||
// data.Value is not usable here, json numbers decode to float64 and format as "1.345027e+06".
|
||||
// the "gid:" prefix keeps numeric ids out of the login hash space, logins may be all-digit
|
||||
var uinfoJSON struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := json.Unmarshal(bdata, &uinfoJSON); err == nil && uinfoJSON.ID != 0 {
|
||||
userInfo.ID = "github_" + token.HashID(sha1.New(), "gid:"+strconv.FormatInt(uinfoJSON.ID, 10))
|
||||
} else {
|
||||
// keep the login-based value, matching the default derivation and its recycling caveat
|
||||
p.Logf("[WARN] github numeric id not available, keeping login-based id")
|
||||
}
|
||||
}
|
||||
// github may have no user name, use login in this case
|
||||
if userInfo.Name == "" {
|
||||
userInfo.Name = data.Value("login")
|
||||
@@ -171,7 +189,7 @@ func NewTwitter(p Params) Oauth1Handler {
|
||||
func NewBattlenet(p Params) Oauth2Handler {
|
||||
return initOauth2Handler(p, Oauth2Handler{
|
||||
name: "battlenet",
|
||||
endpoint: oauth2.Endpoint{
|
||||
endpoint: oauth2.Endpoint{ //nolint:gosec // G101 false positive, oauth endpoint urls are not credentials
|
||||
AuthURL: "https://eu.battle.net/oauth/authorize",
|
||||
TokenURL: "https://eu.battle.net/oauth/token",
|
||||
AuthStyle: oauth2.AuthStyleInParams,
|
||||
@@ -241,7 +259,7 @@ func NewPatreon(p Params) Oauth2Handler {
|
||||
return initOauth2Handler(p, Oauth2Handler{
|
||||
name: "patreon",
|
||||
// see https://docs.patreon.com/?shell#oauth
|
||||
endpoint: oauth2.Endpoint{
|
||||
endpoint: oauth2.Endpoint{ //nolint:gosec // G101 false positive, oauth endpoint urls are not credentials
|
||||
AuthURL: "https://www.patreon.com/oauth2/authorize",
|
||||
TokenURL: "https://api.patreon.com/oauth2/token",
|
||||
AuthStyle: oauth2.AuthStyleInParams,
|
||||
@@ -276,7 +294,7 @@ func NewDiscord(p Params) Oauth2Handler {
|
||||
return initOauth2Handler(p, Oauth2Handler{
|
||||
name: "discord",
|
||||
// see https://discord.com/developers/docs/topics/oauth2
|
||||
endpoint: oauth2.Endpoint{
|
||||
endpoint: oauth2.Endpoint{ //nolint:gosec // G101 false positive, oauth endpoint urls are not credentials
|
||||
AuthURL: "https://discord.com/oauth2/authorize",
|
||||
TokenURL: "https://discord.com/api/oauth2/token",
|
||||
},
|
||||
|
||||
+20
-6
@@ -2,6 +2,7 @@
|
||||
package sender
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/go-pkgz/auth/v2/logger"
|
||||
@@ -20,6 +21,7 @@ type Email struct {
|
||||
type EmailParams struct {
|
||||
Host string // SMTP host
|
||||
Port int // SMTP port
|
||||
HELOHost string // SMTP HELO/EHLO hostname, defaults to localhost when unset
|
||||
From string // from email field
|
||||
Subject string // email subject
|
||||
ContentType string // content type
|
||||
@@ -58,6 +60,10 @@ func NewEmailClient(emailParams EmailParams, l logger.L) *Email {
|
||||
opts = append(opts, email.Port(emailParams.Port))
|
||||
}
|
||||
|
||||
if emailParams.HELOHost != "" {
|
||||
opts = append(opts, email.HELOHost(emailParams.HELOHost))
|
||||
}
|
||||
|
||||
if emailParams.TimeOut != 0 {
|
||||
opts = append(opts, email.TimeOut(emailParams.TimeOut))
|
||||
}
|
||||
@@ -79,14 +85,22 @@ func NewEmailClient(emailParams EmailParams, l logger.L) *Email {
|
||||
return &Email{EmailParams: emailParams, L: l, sender: sender}
|
||||
}
|
||||
|
||||
// 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.
|
||||
// Send email with given text, with no cancellation and with TimeOut applied to the connection setup only.
|
||||
// See SendContext for the details.
|
||||
func (e *Email) Send(to, text string) error {
|
||||
return e.SendContext(context.Background(), to, text)
|
||||
}
|
||||
|
||||
// SendContext sends email with given text and terminates the whole SMTP transaction as soon as ctx is done.
|
||||
// TimeOut covers the connection setup only, so without a context a server accepting the connection and
|
||||
// stalling afterwards blocks the caller for as long as it likes.
|
||||
//
|
||||
// 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) SendContext(ctx context.Context, to, text string) error {
|
||||
e.Logf("[DEBUG] send %d-byte message to %s", len(text), to)
|
||||
return e.sender.Send(text, email.Params{
|
||||
return e.sender.SendContext(ctx, text, email.Params{
|
||||
From: e.From,
|
||||
To: []string{to},
|
||||
Subject: e.Subject,
|
||||
|
||||
+12
@@ -118,6 +118,18 @@ func setAvatar(ava AvatarSaver, u token.User, client *http.Client) (token.User,
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// sessionOnlyFromRequest reads the session-only flag from the request. Both "session",
|
||||
// as documented in README, and "sess", accepted by the direct and verify providers since
|
||||
// the beginning, are honored.
|
||||
func sessionOnlyFromRequest(r *http.Request) bool {
|
||||
for _, name := range []string{"session", "sess"} {
|
||||
if v := r.URL.Query().Get(name); v != "" {
|
||||
return v != "0"
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func randToken() (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
|
||||
+25
-5
@@ -2,6 +2,7 @@ package provider
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
@@ -166,6 +167,12 @@ type Sender interface {
|
||||
Send(address, text string) error
|
||||
}
|
||||
|
||||
// ContextSender is an optional extension of Sender. A Sender implementing it gets the request context,
|
||||
// so the delivery is bound to the request instead of running until the transport times out on its own.
|
||||
type ContextSender interface {
|
||||
SendContext(ctx context.Context, address, text string) error
|
||||
}
|
||||
|
||||
// SenderFunc type is an adapter to allow the use of ordinary functions as Sender.
|
||||
type SenderFunc func(address, text string) error
|
||||
|
||||
@@ -174,6 +181,14 @@ func (f SenderFunc) Send(address, text string) error {
|
||||
return f(address, text)
|
||||
}
|
||||
|
||||
// send delivers the message, passing ctx to senders which accept it
|
||||
func (e VerifyHandler) send(ctx context.Context, address, text string) error {
|
||||
if s, ok := e.Sender.(ContextSender); ok {
|
||||
return s.SendContext(ctx, address, text)
|
||||
}
|
||||
return e.Sender.Send(address, text)
|
||||
}
|
||||
|
||||
// VerifTokenService defines interface accessing tokens
|
||||
type VerifTokenService interface {
|
||||
Token(claims token.Claims) (string, error)
|
||||
@@ -252,7 +267,7 @@ func (e VerifyHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
user, address := elems[0], elems[1]
|
||||
sessOnly := r.URL.Query().Get("sess") == "1"
|
||||
sessOnly := sessionOnlyFromRequest(r)
|
||||
|
||||
u := token.User{
|
||||
Name: user,
|
||||
@@ -299,7 +314,8 @@ func (e VerifyHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
rest.RenderJSON(w, claims.User)
|
||||
return
|
||||
}
|
||||
http.Redirect(w, r, confClaims.Handshake.From, http.StatusTemporaryRedirect)
|
||||
// see-other keeps the redirect a GET whatever method the callback arrived with
|
||||
http.Redirect(w, r, confClaims.Handshake.From, http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
rest.RenderJSON(w, claims.User)
|
||||
@@ -308,7 +324,11 @@ func (e VerifyHandler) LoginHandler(w http.ResponseWriter, r *http.Request) {
|
||||
// GET /login?site=site&user=name&address=someone@example.com
|
||||
func (e VerifyHandler) sendConfirmation(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
user, address, site := r.URL.Query().Get("user"), r.URL.Query().Get("address"), r.URL.Query().Get("site")
|
||||
user, address := r.URL.Query().Get("user"), r.URL.Query().Get("address")
|
||||
site := r.URL.Query().Get("site")
|
||||
if site == "" { // documented as aud, kept reading site for backward compatibility
|
||||
site = r.URL.Query().Get("aud")
|
||||
}
|
||||
|
||||
if user == "" || address == "" {
|
||||
rest.SendErrorJSON(w, r, e.L, http.StatusBadRequest, fmt.Errorf("wrong request"), "can't get user and address")
|
||||
@@ -327,7 +347,7 @@ func (e VerifyHandler) sendConfirmation(w http.ResponseWriter, r *http.Request)
|
||||
// honor from at all.
|
||||
From: r.URL.Query().Get("from"),
|
||||
},
|
||||
SessionOnly: r.URL.Query().Get("session") != "" && r.URL.Query().Get("session") != "0",
|
||||
SessionOnly: sessionOnlyFromRequest(r),
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Audience: []string{site},
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(30 * time.Minute)),
|
||||
@@ -372,7 +392,7 @@ func (e VerifyHandler) sendConfirmation(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
if err := e.Sender.Send(address, buf.String()); err != nil {
|
||||
if err := e.send(r.Context(), address, buf.String()); err != nil {
|
||||
rest.SendErrorJSON(w, r, e.L, http.StatusInternalServerError, err, "failed to send confirmation")
|
||||
return
|
||||
}
|
||||
|
||||
+4
-4
@@ -277,11 +277,11 @@ func (j *Service) Set(w http.ResponseWriter, claims Claims) (Claims, error) {
|
||||
cookieExpiration = int(j.CookieDuration.Seconds())
|
||||
}
|
||||
|
||||
jwtCookie := http.Cookie{Name: j.JWTCookieName, Value: tokenString, HttpOnly: true, Path: "/", Domain: j.JWTCookieDomain,
|
||||
jwtCookie := http.Cookie{Name: j.JWTCookieName, Value: tokenString, HttpOnly: true, Path: "/", Domain: j.JWTCookieDomain, //nolint:gosec // Secure and SameSite come from service config
|
||||
MaxAge: cookieExpiration, Secure: j.SecureCookies, SameSite: j.SameSite}
|
||||
http.SetCookie(w, &jwtCookie)
|
||||
|
||||
xsrfCookie := http.Cookie{Name: j.XSRFCookieName, Value: claims.ID, HttpOnly: false, Path: "/", Domain: j.JWTCookieDomain,
|
||||
xsrfCookie := http.Cookie{Name: j.XSRFCookieName, Value: claims.ID, HttpOnly: false, Path: "/", Domain: j.JWTCookieDomain, //nolint:gosec // HttpOnly false by design, JS reads it for the X-XSRF-Token header
|
||||
MaxAge: cookieExpiration, Secure: j.SecureCookies, SameSite: j.SameSite}
|
||||
http.SetCookie(w, &xsrfCookie)
|
||||
|
||||
@@ -357,11 +357,11 @@ func (j *Service) IsExpired(claims Claims) bool {
|
||||
|
||||
// Reset token's cookies
|
||||
func (j *Service) Reset(w http.ResponseWriter) {
|
||||
jwtCookie := http.Cookie{Name: j.JWTCookieName, Value: "", HttpOnly: false, Path: "/", Domain: j.JWTCookieDomain,
|
||||
jwtCookie := http.Cookie{Name: j.JWTCookieName, Value: "", HttpOnly: false, Path: "/", Domain: j.JWTCookieDomain, //nolint:gosec // expired removal cookie, carries no value
|
||||
MaxAge: -1, Expires: time.Unix(0, 0), Secure: j.SecureCookies, SameSite: j.SameSite}
|
||||
http.SetCookie(w, &jwtCookie)
|
||||
|
||||
xsrfCookie := http.Cookie{Name: j.XSRFCookieName, Value: "", HttpOnly: false, Path: "/", Domain: j.JWTCookieDomain,
|
||||
xsrfCookie := http.Cookie{Name: j.XSRFCookieName, Value: "", HttpOnly: false, Path: "/", Domain: j.JWTCookieDomain, //nolint:gosec // expired removal cookie, carries no value
|
||||
MaxAge: -1, Expires: time.Unix(0, 0), Secure: j.SecureCookies, SameSite: j.SameSite}
|
||||
http.SetCookie(w, &xsrfCookie)
|
||||
|
||||
|
||||
+1
@@ -20,6 +20,7 @@ linters:
|
||||
- unconvert
|
||||
- unparam
|
||||
- unused
|
||||
- testifylint
|
||||
settings:
|
||||
goconst:
|
||||
min-len: 2
|
||||
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
# CLAUDE.md
|
||||
|
||||
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
||||
|
||||
## What this is
|
||||
|
||||
`github.com/go-pkgz/email` is a single-package library wrapping the stdlib `net/smtp` to simplify sending
|
||||
messages (alerts, notifications, password-reset mails). It has no runtime dependencies — only `testify` for
|
||||
tests. Not designed for high-volume/low-latency bulk sending. Go 1.19.
|
||||
|
||||
## Commands
|
||||
|
||||
- Test: `go test -race ./...` — CI runs with `TZ=America/Chicago`; the Date header uses RFC1123Z with an
|
||||
injectable clock, so timezone can affect assertions. Match CI: `TZ=America/Chicago go test -race ./...`.
|
||||
- Single test: `go test -run TestEmail_Send`
|
||||
- Lint: `golangci-lint run` from repo root (config is golangci-lint **v2** format; CI pins v2.6)
|
||||
- Regenerate mocks: `go generate ./...` (requires `moq`; do not hand-edit files under `mocks/`)
|
||||
|
||||
## Architecture
|
||||
|
||||
Three source files, one package:
|
||||
|
||||
- `email.go` — `Sender` type, `Send`, and MIME message construction (`buildMessage`).
|
||||
- `options.go` — functional options (`Option = func(*Sender)`), all applied in `NewSender`.
|
||||
- `auth.go` — custom LOGIN SASL auth mechanism.
|
||||
|
||||
Key design points that span files:
|
||||
|
||||
- **Functional options.** `Sender` fields are unexported and set only through `Option` funcs in `options.go`,
|
||||
applied in `NewSender` after defaults. Add a new config knob = new `Option` func + field, nothing else.
|
||||
|
||||
- **`SMTPClient` interface is the seam.** It's a consumer-side subset of `net/smtp.Client` (Mail/Auth/Rcpt/
|
||||
Data/Quit/Close). If the caller injects one via the `SMTP()` option it's reused; otherwise `Send` builds a
|
||||
fresh `net/smtp` client per call via `em.client()`. This interface is what the moq mock implements, so all
|
||||
`Send` tests run without a real SMTP server.
|
||||
|
||||
- **`em.client()` handles the three transport modes:** plain (dial + optional STARTTLS), and implicit `TLS`
|
||||
(dial over TLS, port 465). `STARTTLS` upgrades a plain connection (port 587). `InsecureSkipVerify` feeds the
|
||||
`tls.Config`.
|
||||
|
||||
- **`buildMessage` is the intricate part.** It assembles headers manually, then bodies: `multipart/mixed` for
|
||||
attachments, `multipart/related` for inline images, quoted-printable for the text body. Inline images get an
|
||||
auto `Content-ID` equal to their filename. Attachment/inline type is sniffed from the first 512 bytes via
|
||||
`http.DetectContentType`, then the file is re-seeked to 0 and base64-encoded. Changes here are easy to break
|
||||
silently — the tests assert on the exact serialized message string.
|
||||
|
||||
- **LOGIN auth (`auth.go`) exists because stdlib only ships PLAIN.** Needed for Office 365 / Outlook.com.
|
||||
Enabled with the `LoginAuth()` option; it refuses to send credentials over an unencrypted, non-localhost
|
||||
connection. `Sender.auth()` picks PLAIN vs LOGIN and returns nil (no auth) when username/password are empty.
|
||||
|
||||
- **Envelope vs. headers.** `extractEmailAddress` (via `net/mail`) strips a display name so `MAIL FROM` /
|
||||
`RCPT TO` get a bare address, while the `From`/`To` *headers* keep the full `"Name" <addr>` form. Falls back
|
||||
to the raw string if parsing fails.
|
||||
|
||||
- **`timeNow` field** is an injectable clock (`func() time.Time`, defaults to `time.Now`) so tests can pin the
|
||||
Date header deterministically — set `s.timeNow` directly in tests.
|
||||
|
||||
## Testing conventions
|
||||
|
||||
- Mocks are moq-generated into `mocks/` from `//go:generate` directives at the top of `email.go` for the
|
||||
`SMTPClient` and `Logger` interfaces. Use the mock's `*Calls()` accessors to assert interactions.
|
||||
- `testdata/` holds attachment fixtures (`1.txt`, `2.txt`, `image.jpg`, `nullfile` for the empty-file path).
|
||||
- Coverage in CI strips `mocks`/`_mock.go` lines before submitting to coveralls.
|
||||
+15
@@ -24,6 +24,8 @@ err := client.Send("<html>some content, foo bar</html>",
|
||||
- `TLS`: Use TLS SMTP (default: false)
|
||||
- `STARTTLS`: Use STARTTLS (default: false)
|
||||
- `InsecureSkipVerify`: skip certificate verification (default: false)
|
||||
- `HELOHost`: SMTP HELO/EHLO hostname (default: empty, greets as `localhost`). Some servers reject `localhost`,
|
||||
e.g. Postfix with `reject_non_fqdn_helo_hostname`. Not applied to a custom client set with `SMTP`.
|
||||
- `Auth(user, password)`: Username and password for SMTP authentication (default: empty, no authentication)
|
||||
- `LoginAuth`: Use [LOGIN mechanism](https://www.ietf.org/archive/id/draft-murchison-sasl-login-00.txt) instead of PLAIN mechanism for SMTP authentication, e.g. this is relevant for Office 365 and Outlook.com
|
||||
- `ContentType`: Content type for the email (default: "text/plain")
|
||||
@@ -56,6 +58,19 @@ To send email user need to create a sender first and then use `Send` method. The
|
||||
|
||||
See [go docs](https://pkg.go.dev/github.com/go-pkgz/email#Sender.Send) for `Send` function.
|
||||
|
||||
`SendContext` takes the same parameters with a context added and is the way to bound the time spent on sending.
|
||||
The `TimeOut` option covers the connection setup only, while the context covers the whole SMTP transaction,
|
||||
so a server accepting the connection and stalling afterwards terminates the send instead of blocking the caller:
|
||||
|
||||
```go
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Minute)
|
||||
defer cancel()
|
||||
err := client.SendContext(ctx, "some content", email.Params{From: "me@example.com", To: []string{"to@example.com"}})
|
||||
```
|
||||
|
||||
A custom smtp client set with the `SMTP` option owns its connection, and such a transaction can't be
|
||||
terminated in the middle; the context is checked before it starts in that case.
|
||||
|
||||
## technical details
|
||||
|
||||
- Content-Transfer-Encoding set to `quoted-printable`
|
||||
|
||||
+286
-123
@@ -3,6 +3,7 @@ package email
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
@@ -20,6 +21,7 @@ import (
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -31,6 +33,7 @@ type Sender struct {
|
||||
smtpClient SMTPClient
|
||||
logger Logger
|
||||
host string // SMTP host
|
||||
heloHost string // SMTP HELO/EHLO host
|
||||
port int // SMTP port
|
||||
contentType string // content type, optional. Will trigger MIME and Content-Type headers
|
||||
tls bool // TLS auth
|
||||
@@ -91,53 +94,73 @@ func NewSender(smtpHost string, options ...Option) *Sender {
|
||||
opt(&res)
|
||||
}
|
||||
|
||||
res.logger.Logf("[INFO] new email sender created with host: %s:%d, tls: %v, insecureSkipVerify: %v, username: %q, timeout: %v, "+
|
||||
"content type: %q, charset: %q", smtpHost,
|
||||
res.port, res.tls, res.insecureSkipVerify, res.smtpUserName, res.timeOut, res.contentType, res.contentCharset)
|
||||
res.logger.Logf("[INFO] new email sender created with host: %s:%d, helo: %q, tls: %v, insecureSkipVerify: %v, username: %q, timeout: %v, "+
|
||||
"content type: %q, charset: %q", smtpHost, res.port, res.effectiveHELOHost(),
|
||||
res.tls, res.insecureSkipVerify, res.smtpUserName, res.timeOut, res.contentType, res.contentCharset)
|
||||
return &res
|
||||
}
|
||||
|
||||
// Send email with given text
|
||||
// If SMTPClient defined in Email struct it will be used, if not - new smtp.Client on each send.
|
||||
// Always closes client on completion or failure.
|
||||
// Send email with given text, with no cancellation and with TimeOut applied to the connection setup only.
|
||||
// See SendContext for the details.
|
||||
func (em *Sender) Send(text string, params Params) error {
|
||||
return em.SendContext(context.Background(), text, params)
|
||||
}
|
||||
|
||||
// SendContext sends email with given text and terminates the whole SMTP transaction as soon as ctx is done,
|
||||
// including the greeting, the authentication and the message body transfer.
|
||||
// If SMTPClient set with the SMTP option it will be used, if not - new smtp.Client on each send.
|
||||
// Note that a client set that way owns its connection, so such a transaction can't be terminated in the middle.
|
||||
// Always closes client on completion or failure.
|
||||
func (em *Sender) SendContext(ctx context.Context, text string, params Params) error {
|
||||
em.logger.Logf("[DEBUG] send %q to %v", text, params.To)
|
||||
|
||||
client := em.smtpClient
|
||||
if client == nil { // if client not set make new net/smtp
|
||||
c, err := em.client()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to make smtp client: %w", err)
|
||||
}
|
||||
client = c
|
||||
}
|
||||
client := em.smtpClient // set by the SMTP option, nil when SendContext makes its own client below
|
||||
|
||||
var quit bool
|
||||
defer func() {
|
||||
if quit || client == nil { // quit set if Quit() call passed because it's closing connection as well.
|
||||
return
|
||||
}
|
||||
if err := client.Close(); err != nil {
|
||||
em.logger.Logf("[WARN] can't close smtp connection, %v", err)
|
||||
if e := client.Close(); e != nil {
|
||||
em.logger.Logf("[WARN] can't close smtp connection, %v", e)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := ctx.Err(); err != nil { // nothing started yet, a client set with the SMTP option is closed by the defer
|
||||
return err
|
||||
}
|
||||
|
||||
if len(params.To) == 0 {
|
||||
return errors.New("no recipients")
|
||||
}
|
||||
|
||||
// message is built before the connection is made, this way a bad message doesn't reach the server at all
|
||||
msg, err := em.buildMessage(text, params)
|
||||
if err != nil {
|
||||
return fmt.Errorf("can't make email message: %w", err)
|
||||
}
|
||||
|
||||
if client == nil { // if client not set make new net/smtp
|
||||
c, stop, e := em.client(ctx)
|
||||
if e != nil {
|
||||
return fmt.Errorf("failed to make smtp client: %w", e)
|
||||
}
|
||||
defer stop() // runs before the deferred close above, releasing the ctx watcher first
|
||||
client = c
|
||||
}
|
||||
|
||||
if auth := em.auth(); auth != nil {
|
||||
if err := client.Auth(auth); err != nil {
|
||||
if err = client.Auth(auth); err != nil {
|
||||
return fmt.Errorf("failed to auth to smtp %s:%d, %w", em.host, em.port, err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := client.Mail(extractEmailAddress(params.From)); err != nil {
|
||||
if err = client.Mail(extractEmailAddress(params.From)); err != nil {
|
||||
return fmt.Errorf("bad from address %q: %w", params.From, err)
|
||||
}
|
||||
|
||||
for _, rcpt := range params.To {
|
||||
if err := client.Rcpt(extractEmailAddress(rcpt)); err != nil {
|
||||
if err = client.Rcpt(extractEmailAddress(rcpt)); err != nil {
|
||||
return fmt.Errorf("bad to address %q: %w", params.To, err)
|
||||
}
|
||||
}
|
||||
@@ -147,16 +170,12 @@ func (em *Sender) Send(text string, params Params) error {
|
||||
return fmt.Errorf("can't make email writer: %w", err)
|
||||
}
|
||||
|
||||
msg, err := em.buildMessage(text, params)
|
||||
if err != nil {
|
||||
return fmt.Errorf("can't make email message: %w", err)
|
||||
}
|
||||
buf := bytes.NewBufferString(msg)
|
||||
if _, err = buf.WriteTo(writer); err != nil {
|
||||
if _, err = msg.WriteTo(writer); err != nil {
|
||||
return fmt.Errorf("failed to send email body to %q: %w", params.To, err)
|
||||
}
|
||||
// closing the writer reports the final response to the DATA command, i.e. the actual delivery result
|
||||
if err = writer.Close(); err != nil {
|
||||
em.logger.Logf("[WARN] can't close smtp body writer, %v", err)
|
||||
return fmt.Errorf("failed to send email to %q: %w", params.To, err)
|
||||
}
|
||||
|
||||
if err = client.Quit(); err != nil {
|
||||
@@ -179,11 +198,25 @@ func extractEmailAddress(from string) string {
|
||||
}
|
||||
|
||||
func (em *Sender) String() string {
|
||||
return fmt.Sprintf("smtp://%s:%d, auth:%v, tls:%v, starttls:%v, insecureSkipVerify:%v, timeout:%v, content-type:%q, charset:%q",
|
||||
em.host, em.port, em.smtpUserName != "", em.tls, em.starttls, em.insecureSkipVerify, em.timeOut, em.contentType, em.contentCharset)
|
||||
return fmt.Sprintf("smtp://%s:%d, helo:%q, auth:%v, tls:%v, starttls:%v, insecureSkipVerify:%v, timeout:%v, content-type:%q, charset:%q",
|
||||
em.host, em.port, em.effectiveHELOHost(), em.smtpUserName != "", em.tls, em.starttls, em.insecureSkipVerify,
|
||||
em.timeOut, em.contentType, em.contentCharset)
|
||||
}
|
||||
|
||||
func (em *Sender) client() (c *smtp.Client, err error) {
|
||||
func (em *Sender) effectiveHELOHost() string {
|
||||
if em.smtpClient != nil {
|
||||
return "client-managed"
|
||||
}
|
||||
if em.heloHost == "" {
|
||||
return "localhost"
|
||||
}
|
||||
return em.heloHost
|
||||
}
|
||||
|
||||
// client makes smtp client with the connection bound to ctx: it is closed as soon as ctx is done,
|
||||
// which is the only way to interrupt net/smtp calls as they take no context.
|
||||
// Returned stop function releases that binding and has to be called when the client is not needed anymore.
|
||||
func (em *Sender) client(ctx context.Context) (c *smtp.Client, stop func(), err error) {
|
||||
srvAddress := net.JoinHostPort(em.host, strconv.Itoa(em.port))
|
||||
// #nosec G402
|
||||
tlsConf := &tls.Config{
|
||||
@@ -192,34 +225,74 @@ func (em *Sender) client() (c *smtp.Client, err error) {
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
dialer := &net.Dialer{Timeout: em.timeOut}
|
||||
|
||||
var conn net.Conn
|
||||
if em.tls {
|
||||
conn, e := tls.DialWithDialer(&net.Dialer{Timeout: em.timeOut}, "tcp", srvAddress, tlsConf)
|
||||
if e != nil {
|
||||
return nil, fmt.Errorf("failed to dial smtp tls to %s: %w", srvAddress, e)
|
||||
if conn, err = (&tls.Dialer{NetDialer: dialer, Config: tlsConf}).DialContext(ctx, "tcp", srvAddress); err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to dial smtp tls to %s: %w", srvAddress, err)
|
||||
}
|
||||
if c, err = smtp.NewClient(conn, em.host); err != nil {
|
||||
return nil, fmt.Errorf("failed to make smtp client for %s: %w", srvAddress, err)
|
||||
} else {
|
||||
if conn, err = dialer.DialContext(ctx, "tcp", srvAddress); err != nil {
|
||||
return nil, nil, fmt.Errorf("timeout connecting to %s: %w", srvAddress, err)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
conn, err := net.DialTimeout("tcp", srvAddress, em.timeOut)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("timeout connecting to %s: %w", srvAddress, err)
|
||||
// closing the connection is the only way to interrupt net/smtp calls, as they take no context
|
||||
watchDone := make(chan struct{})
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = conn.Close()
|
||||
case <-watchDone:
|
||||
}
|
||||
}()
|
||||
var stopOnce sync.Once
|
||||
stop = func() { stopOnce.Do(func() { close(watchDone) }) }
|
||||
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
if e := conn.SetDeadline(deadline); e != nil {
|
||||
em.logger.Logf("[WARN] can't set deadline on smtp connection to %s, %v", srvAddress, e)
|
||||
}
|
||||
}
|
||||
|
||||
c, err = smtp.NewClient(conn, em.host)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to dial: %w", err)
|
||||
if c, err = smtp.NewClient(conn, em.host); err != nil {
|
||||
stop()
|
||||
_ = conn.Close()
|
||||
if ctxErr := ctx.Err(); ctxErr != nil {
|
||||
return nil, nil, fmt.Errorf("failed to make smtp client for %s: %w", srvAddress, ctxErr)
|
||||
}
|
||||
if em.tls {
|
||||
return nil, nil, fmt.Errorf("failed to make smtp client for %s: %w", srvAddress, err)
|
||||
}
|
||||
return nil, nil, fmt.Errorf("failed to dial: %w", err)
|
||||
}
|
||||
|
||||
if em.starttls {
|
||||
if err = em.hello(c); err != nil {
|
||||
stop()
|
||||
_ = c.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
if !em.tls && em.starttls {
|
||||
if err = c.StartTLS(tlsConf); err != nil {
|
||||
return nil, fmt.Errorf("failed to start tls: %w", err)
|
||||
stop()
|
||||
_ = c.Close()
|
||||
return nil, nil, fmt.Errorf("failed to start tls: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return c, nil
|
||||
return c, stop, nil
|
||||
}
|
||||
|
||||
func (em *Sender) hello(client *smtp.Client) error {
|
||||
if em.heloHost == "" {
|
||||
return nil
|
||||
}
|
||||
if err := client.Hello(em.heloHost); err != nil {
|
||||
return fmt.Errorf("failed to send SMTP greeting: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// auth returns an smtp.Auth that implements SMTP authentication mechanism
|
||||
@@ -235,78 +308,112 @@ func (em *Sender) auth() smtp.Auth {
|
||||
return smtp.PlainAuth("", em.smtpUserName, em.smtpPassword, em.host)
|
||||
}
|
||||
|
||||
func (em *Sender) buildMessage(text string, params Params) (message string, err error) {
|
||||
addHeader := func(msg, h, v string) string {
|
||||
msg += fmt.Sprintf("%s: %s\n", h, v)
|
||||
return msg
|
||||
}
|
||||
message = addHeader(message, "From", params.From)
|
||||
message = addHeader(message, "To", strings.Join(params.To, ","))
|
||||
message = addHeader(message, "Subject", mime.BEncoding.Encode("utf-8", params.Subject))
|
||||
|
||||
if params.UnsubscribeLink != "" {
|
||||
message = addHeader(message, "List-Unsubscribe-Post", "List-Unsubscribe=One-Click")
|
||||
message = addHeader(message, "List-Unsubscribe", "<"+params.UnsubscribeLink+">")
|
||||
// validateHeaders rejects user-provided values which would break out of the header they are put into.
|
||||
// CR and LF allow injecting arbitrary headers and message body, i.e. sending a different email than the caller intended.
|
||||
func (params Params) validateHeaders() error {
|
||||
check := func(name, value string) error {
|
||||
if strings.ContainsAny(value, "\r\n") {
|
||||
return fmt.Errorf("invalid %s header value %q: contains CR or LF", name, value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if params.InReplyTo != "" {
|
||||
message = addHeader(message, "In-reply-to", "<"+params.InReplyTo+">")
|
||||
for _, h := range [][2]string{
|
||||
{"From", params.From},
|
||||
{"Subject", params.Subject},
|
||||
{"List-Unsubscribe", params.UnsubscribeLink},
|
||||
{"In-reply-to", params.InReplyTo},
|
||||
} {
|
||||
if err := check(h[0], h[1]); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
withAttachments := len(params.Attachments) > 0
|
||||
withInlineImg := len(params.InlineImages) > 0
|
||||
|
||||
if em.contentType != "" || withAttachments || withInlineImg {
|
||||
message = addHeader(message, "MIME-version", "1.0")
|
||||
for _, to := range params.To {
|
||||
if err := check("To", to); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
message = addHeader(message, "Date", em.timeNow().Format(time.RFC1123Z))
|
||||
// buildMessage makes the complete message, headers and body, in a single buffer the caller sends as is
|
||||
func (em *Sender) buildMessage(text string, params Params) (*bytes.Buffer, error) {
|
||||
if err := params.validateHeaders(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
buff := &bytes.Buffer{}
|
||||
addHeader := func(h, v string) {
|
||||
fmt.Fprintf(buff, "%s: %s\n", h, v)
|
||||
}
|
||||
|
||||
// body writers are made upfront because the boundaries they pick are needed in the headers,
|
||||
// they write nothing until used, i.e. after all the headers are in the buffer
|
||||
qp := quotedprintable.NewWriter(buff)
|
||||
mpMixed := multipart.NewWriter(buff)
|
||||
boundaryMixed := mpMixed.Boundary()
|
||||
mpRelated := multipart.NewWriter(buff)
|
||||
boundaryRelated := mpRelated.Boundary()
|
||||
|
||||
addHeader("From", params.From)
|
||||
addHeader("To", strings.Join(params.To, ","))
|
||||
addHeader("Subject", mime.BEncoding.Encode("utf-8", params.Subject))
|
||||
|
||||
if params.UnsubscribeLink != "" {
|
||||
addHeader("List-Unsubscribe-Post", "List-Unsubscribe=One-Click")
|
||||
addHeader("List-Unsubscribe", "<"+params.UnsubscribeLink+">")
|
||||
}
|
||||
|
||||
if params.InReplyTo != "" {
|
||||
addHeader("In-reply-to", "<"+params.InReplyTo+">")
|
||||
}
|
||||
|
||||
withAttachments := len(params.Attachments) > 0
|
||||
withInlineImg := len(params.InlineImages) > 0
|
||||
|
||||
if em.contentType != "" || withAttachments || withInlineImg {
|
||||
addHeader("MIME-version", "1.0")
|
||||
}
|
||||
|
||||
addHeader("Date", em.timeNow().Format(time.RFC1123Z))
|
||||
|
||||
if withAttachments {
|
||||
message = addHeader(message, "Content-Type", fmt.Sprintf("multipart/mixed; boundary=%q\r\n\r\n%s\r",
|
||||
addHeader("Content-Type", fmt.Sprintf("multipart/mixed; boundary=%q\r\n\r\n%s\r",
|
||||
boundaryMixed, "--"+boundaryMixed))
|
||||
}
|
||||
|
||||
if withInlineImg {
|
||||
message = addHeader(message, "Content-Type", fmt.Sprintf("multipart/related; boundary=%q\r\n\r\n%s\r",
|
||||
addHeader("Content-Type", fmt.Sprintf("multipart/related; boundary=%q\r\n\r\n%s\r",
|
||||
boundaryRelated, "--"+boundaryRelated))
|
||||
}
|
||||
|
||||
if em.contentType != "" {
|
||||
message = addHeader(message, "Content-Transfer-Encoding", "quoted-printable")
|
||||
message = addHeader(message, "Content-Type", fmt.Sprintf("%s; charset=%q", em.contentType, em.contentCharset))
|
||||
|
||||
addHeader("Content-Transfer-Encoding", "quoted-printable")
|
||||
addHeader("Content-Type", fmt.Sprintf("%s; charset=%q", em.contentType, em.contentCharset))
|
||||
}
|
||||
|
||||
buff.WriteString("\n") // empty line between the headers and the body
|
||||
|
||||
if err := em.writeBody(qp, text); err != nil {
|
||||
return "", fmt.Errorf("failed to write body: %w", err)
|
||||
return nil, fmt.Errorf("failed to write body: %w", err)
|
||||
}
|
||||
|
||||
if withInlineImg {
|
||||
buff.WriteString("\r\n\r\n")
|
||||
if err := em.writeFiles(mpRelated, params.InlineImages, "inline"); err != nil {
|
||||
return "", fmt.Errorf("failed to write inline images: %w", err)
|
||||
return nil, fmt.Errorf("failed to write inline images: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if withAttachments {
|
||||
buff.WriteString("\r\n\r\n")
|
||||
if err := em.writeFiles(mpMixed, params.Attachments, "attachment"); err != nil {
|
||||
return "", fmt.Errorf("failed to write attachments: %w", err)
|
||||
return nil, fmt.Errorf("failed to write attachments: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
m := buff.String()
|
||||
message += "\n" + m
|
||||
// returns base part of the file location
|
||||
return message, nil
|
||||
return buff, nil
|
||||
}
|
||||
|
||||
func (em *Sender) writeBody(wc io.WriteCloser, text string) error {
|
||||
@@ -321,53 +428,7 @@ func (em *Sender) writeBody(wc io.WriteCloser, text string) error {
|
||||
|
||||
func (em *Sender) writeFiles(mp *multipart.Writer, files []string, disposition string) error {
|
||||
for _, attachment := range files {
|
||||
file, err := os.Open(filepath.Clean(attachment))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// we need first 512 bytes to detect file type
|
||||
fTypeBuff := make([]byte, 512)
|
||||
_, err = file.Read(fTypeBuff)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read file type %q: %w", attachment, err)
|
||||
}
|
||||
|
||||
// remove null bytes in case file less than 512 bytes
|
||||
fTypeBuff = bytes.Trim(fTypeBuff, "\x00")
|
||||
fName := filepath.Base(attachment)
|
||||
header := textproto.MIMEHeader{}
|
||||
header.Set("Content-Type", http.DetectContentType(fTypeBuff)+"; name=\""+fName+"\"")
|
||||
header.Set("Content-Transfer-Encoding", "base64")
|
||||
|
||||
switch disposition {
|
||||
case "attachment":
|
||||
header.Set("Content-Disposition", fmt.Sprintf("attachment; filename=%q", fName))
|
||||
case "inline":
|
||||
header.Set("Content-Disposition", fmt.Sprintf("inline; filename=%q", fName))
|
||||
header.Set("Content-ID", fmt.Sprintf("<%s>", fName))
|
||||
}
|
||||
|
||||
writer, err := mp.CreatePart(header)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// set reader offset at the beginning of the file because we read first 512 bytes
|
||||
_, err = file.Seek(0, io.SeekStart)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
encoder := base64.NewEncoder(base64.StdEncoding, writer)
|
||||
if _, err := io.Copy(encoder, file); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := encoder.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := file.Close(); err != nil {
|
||||
if err := em.writeFile(mp, attachment, disposition); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
@@ -377,6 +438,108 @@ func (em *Sender) writeFiles(mp *multipart.Writer, files []string, disposition s
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeFile adds a single file as a mime part, the file is closed on every return path
|
||||
func (em *Sender) writeFile(mp *multipart.Writer, attachment, disposition string) (err error) {
|
||||
file, err := os.Open(filepath.Clean(attachment))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
if e := file.Close(); e != nil && err == nil {
|
||||
err = e
|
||||
}
|
||||
}()
|
||||
|
||||
// we need first 512 bytes to detect file type, an empty file is fine and detected as plain text
|
||||
fTypeBuff := make([]byte, 512)
|
||||
n, err := file.Read(fTypeBuff)
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return fmt.Errorf("failed to read file type %q: %w", attachment, err)
|
||||
}
|
||||
fTypeBuff = fTypeBuff[:n] // file can be shorter than the buffer
|
||||
|
||||
fName := filepath.Base(attachment)
|
||||
// CR and LF are legal in file names but would terminate the header the name goes into
|
||||
if strings.ContainsAny(fName, "\r\n") {
|
||||
return fmt.Errorf("invalid file name %q: contains CR or LF", attachment)
|
||||
}
|
||||
|
||||
// the detected type is always parseable, it comes from a fixed set of sniffed types
|
||||
contentType, ctParams, _ := mime.ParseMediaType(http.DetectContentType(fTypeBuff))
|
||||
params := map[string]string{"name": fName}
|
||||
for k, v := range ctParams { // carries the charset of the text types over
|
||||
params[k] = v
|
||||
}
|
||||
|
||||
// mime formatting quotes and encodes the file name, plain interpolation would let it break out of the header
|
||||
header := textproto.MIMEHeader{}
|
||||
header.Set("Content-Type", mime.FormatMediaType(contentType, params))
|
||||
header.Set("Content-Transfer-Encoding", "base64")
|
||||
|
||||
switch disposition {
|
||||
case "attachment", "inline":
|
||||
header.Set("Content-Disposition", mime.FormatMediaType(disposition, map[string]string{"filename": fName}))
|
||||
}
|
||||
if disposition == "inline" {
|
||||
header.Set("Content-ID", fmt.Sprintf("<%s>", fName))
|
||||
}
|
||||
|
||||
writer, err := mp.CreatePart(header)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// set reader offset at the beginning of the file because we read first 512 bytes
|
||||
if _, err = file.Seek(0, io.SeekStart); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
encoder := base64.NewEncoder(base64.StdEncoding, &lineWrapper{w: writer, limit: base64LineLimit})
|
||||
if _, err = io.Copy(encoder, file); err != nil {
|
||||
return err
|
||||
}
|
||||
return encoder.Close()
|
||||
}
|
||||
|
||||
// base64LineLimit is the maximum line length for base64 encoded mime parts, set by RFC 2045
|
||||
const base64LineLimit = 76
|
||||
|
||||
// crlf is shared by the line wrapper to keep it from allocating a separator for every line
|
||||
var crlf = []byte("\r\n")
|
||||
|
||||
// lineWrapper breaks the stream written to it into lines of limit characters, separated by CRLF.
|
||||
// Base64 encoders produce a single unbroken line, which mime doesn't allow and strict relays reject.
|
||||
type lineWrapper struct {
|
||||
w io.Writer
|
||||
limit int
|
||||
n int // characters already written to the current line
|
||||
}
|
||||
|
||||
func (lw *lineWrapper) Write(p []byte) (int, error) {
|
||||
written := 0
|
||||
for len(p) > 0 {
|
||||
if lw.n == lw.limit {
|
||||
if _, err := lw.w.Write(crlf); err != nil {
|
||||
return written, err
|
||||
}
|
||||
lw.n = 0
|
||||
}
|
||||
|
||||
size := lw.limit - lw.n
|
||||
if size > len(p) {
|
||||
size = len(p)
|
||||
}
|
||||
n, err := lw.w.Write(p[:size])
|
||||
written += n
|
||||
lw.n += n
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
p = p[size:]
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
type nopLogger struct{}
|
||||
|
||||
func (nopLogger) Logf(_ string, _ ...interface{}) {}
|
||||
|
||||
+11
-1
@@ -5,7 +5,7 @@ import "time"
|
||||
// Option func type
|
||||
type Option func(s *Sender)
|
||||
|
||||
// SMTP sets SMTP client
|
||||
// SMTP sets SMTP client. Such a client greets the server itself, so HELOHost doesn't apply to it.
|
||||
func SMTP(smtp SMTPClient) Option {
|
||||
return func(s *Sender) {
|
||||
s.smtpClient = smtp
|
||||
@@ -61,6 +61,16 @@ func InsecureSkipVerify(enabled bool) Option {
|
||||
}
|
||||
}
|
||||
|
||||
// HELOHost sets the SMTP HELO/EHLO hostname for connections created by Sender.
|
||||
// Unset, the greeting stays net/smtp's "localhost". The value is passed to the server as-is,
|
||||
// so an address literal like "[192.0.2.10]" works, and it has no effect on a client
|
||||
// supplied with SMTP, which greets on its own.
|
||||
func HELOHost(host string) Option {
|
||||
return func(s *Sender) {
|
||||
s.heloHost = host
|
||||
}
|
||||
}
|
||||
|
||||
// Auth sets smtp username and password
|
||||
func Auth(smtpUserName, smtpPasswd string) Option {
|
||||
return func(s *Sender) {
|
||||
|
||||
+6
@@ -164,6 +164,12 @@ func (l *Logger) logf(format string, args ...interface{}) {
|
||||
return
|
||||
}
|
||||
|
||||
// hide secrets in the message itself, covers slog handler and panic dump paths
|
||||
// as well as the formatted line redacted later in this function
|
||||
if len(l.secrets) > 0 {
|
||||
msg = string(l.hideSecrets([]byte(msg)))
|
||||
}
|
||||
|
||||
// if slog handler is set, use it
|
||||
if l.slogHandler != nil {
|
||||
// get the caller's PC so slog handlers can resolve source info when AddSource is enabled
|
||||
|
||||
+6
-1
@@ -40,15 +40,19 @@ func main() {
|
||||
}
|
||||
err = notify.Send(context.Background(), notifiers, "https://example.com/webhook", "Hello, world!")
|
||||
if err != nil {
|
||||
fmt.Printf("Sent message error: %s", err))
|
||||
fmt.Printf("Sent message error: %s", err)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
A runnable version of the same flow, sending to a local test server instead of a real webhook, is in [`example_test.go`](example_test.go).
|
||||
|
||||
### Email
|
||||
|
||||
`mailto:` [scheme](https://datatracker.ietf.org/doc/html/rfc6068) is supported. Only `subject` and `from` query params are used.
|
||||
|
||||
**Note:** Query parameter values must be URL-encoded. In particular, email addresses containing `+` (e.g. `noreply+tag@example.com`) must use `%2B`, otherwise `+` is interpreted as a space. Use `url.QueryEscape` for all parameter values.
|
||||
|
||||
Examples:
|
||||
|
||||
- `mailto:"John Wayne"<john@example.org>?subject=test-subj&from="Notifier"<notify@example.org>`
|
||||
@@ -70,6 +74,7 @@ func main() {
|
||||
Host: "localhost", // the only required field, others are optional
|
||||
Port: 25,
|
||||
TLS: false, // TLS, but not STARTTLS
|
||||
HELOHost: "mail.example.org", // hostname sent in the SMTP greeting, "localhost" by default
|
||||
ContentType: "text/html",
|
||||
Charset: "UTF-8",
|
||||
Username: "username",
|
||||
|
||||
+20
-9
@@ -2,6 +2,7 @@ package notify
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/mail"
|
||||
"net/url"
|
||||
@@ -16,14 +17,15 @@ type SMTPParams struct {
|
||||
Host string // SMTP host
|
||||
Port int // SMTP port
|
||||
TLS bool // TLS auth
|
||||
StartTLS bool // StartTLS auth
|
||||
HELOHost string // SMTP HELO/EHLO hostname
|
||||
StartTLS bool // startTLS auth
|
||||
InsecureSkipVerify bool // skip certificate verification
|
||||
ContentType string // Content type
|
||||
Charset string // Character set
|
||||
ContentType string // content type
|
||||
Charset string // character set
|
||||
LoginAuth bool // LOGIN auth method instead of default PLAIN, needed for Office 365 and outlook.com
|
||||
Username string // username
|
||||
Password string // password
|
||||
TimeOut time.Duration // TCP connection timeout
|
||||
TimeOut time.Duration // TCP connection timeout, the rest of the transaction is bound by the context of Send
|
||||
}
|
||||
|
||||
// Email notifications client
|
||||
@@ -56,6 +58,10 @@ func NewEmail(smtpParams SMTPParams) *Email {
|
||||
opts = append(opts, email.Port(smtpParams.Port))
|
||||
}
|
||||
|
||||
if smtpParams.HELOHost != "" {
|
||||
opts = append(opts, email.HELOHost(smtpParams.HELOHost))
|
||||
}
|
||||
|
||||
if smtpParams.TimeOut != 0 {
|
||||
opts = append(opts, email.TimeOut(smtpParams.TimeOut))
|
||||
}
|
||||
@@ -81,6 +87,10 @@ func NewEmail(smtpParams SMTPParams) *Email {
|
||||
// with "mailto:" schema.
|
||||
// "unsubscribeLink" passed as a header, https://support.google.com/mail/answer/81126 -> "Use one-click unsubscribe"
|
||||
//
|
||||
// Note: query parameter values in the mailto URL must be properly URL-encoded. In particular, email addresses
|
||||
// containing "+" (e.g. "noreply+tag@example.com") must use "%2B" instead, otherwise "+" is interpreted as a space
|
||||
// per standard URL query string parsing. Use url.QueryEscape for all parameter values.
|
||||
//
|
||||
// Example:
|
||||
//
|
||||
// - mailto:"John Wayne"<john@example.org>?subject=test-subj&from="Notifier"<notify@example.org>
|
||||
@@ -91,12 +101,13 @@ func (e *Email) Send(ctx context.Context, destination, text string) error {
|
||||
return fmt.Errorf("problem parsing destination: %w", err)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
return e.sender.Send(text, emailParams)
|
||||
// SendContext terminates the transaction when ctx is done, including the parts after the connection is made
|
||||
err = e.sender.SendContext(ctx, text, emailParams)
|
||||
if err != nil && ctx.Err() != nil && !errors.Is(err, ctx.Err()) {
|
||||
// transaction was interrupted, report why on top of the error it failed with
|
||||
return fmt.Errorf("%w: %w", ctx.Err(), err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Schema returns schema prefix supported by this client
|
||||
|
||||
+2
-1
@@ -35,7 +35,8 @@ func (s *Slack) Send(ctx context.Context, destination, text string) error {
|
||||
return fmt.Errorf("problem parsing destination: %w", err)
|
||||
}
|
||||
options := []slack.MsgOption{slack.MsgOptionText(text, false)}
|
||||
if attachment.Title != "" {
|
||||
// titleLink alone carries nothing, slack renders it as a link on the title and drops it without one
|
||||
if attachment.Title != "" || attachment.Text != "" {
|
||||
options = append(options, slack.MsgOptionAttachments(attachment))
|
||||
}
|
||||
|
||||
|
||||
+61
-34
@@ -12,11 +12,10 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
log "github.com/go-pkgz/lgr"
|
||||
"github.com/go-pkgz/repeater"
|
||||
"github.com/go-pkgz/repeater/v2"
|
||||
"github.com/microcosm-cc/bluemonday"
|
||||
"golang.org/x/net/html"
|
||||
)
|
||||
@@ -34,16 +33,19 @@ type TelegramParams struct {
|
||||
type Telegram struct {
|
||||
TelegramParams
|
||||
|
||||
// Identifier of the first update to be requested.
|
||||
// Should be equal to LastSeenUpdateID + 1
|
||||
// identifier of the first update to be requested.
|
||||
// should be equal to LastSeenUpdateID + 1
|
||||
// See https://core.telegram.org/bots/api#getupdates
|
||||
updateOffset int
|
||||
apiPollInterval time.Duration // interval to check updates from Telegram API and answer to users
|
||||
expiredCleanupInterval time.Duration // interval to check and clean up expired notification requests
|
||||
username string // bot username
|
||||
run int32 // non-zero if Run goroutine has started
|
||||
requests struct {
|
||||
updates struct {
|
||||
sync.RWMutex
|
||||
running bool // set while the Run goroutine is active, ProcessUpdate is not allowed then
|
||||
}
|
||||
requests struct {
|
||||
sync.Mutex
|
||||
data map[string]tgAuthRequest
|
||||
}
|
||||
}
|
||||
@@ -224,18 +226,17 @@ func (t *Telegram) AddToken(token, user, site string, expires time.Time) {
|
||||
|
||||
// CheckToken verifies incoming token, returns the user address if it's confirmed and empty string otherwise
|
||||
func (t *Telegram) CheckToken(token, user string) (telegram, site string, err error) {
|
||||
t.requests.RLock()
|
||||
authRequest, ok := t.requests.data[token]
|
||||
t.requests.RUnlock()
|
||||
// lookup and consumption are a single transaction, so the one-time token can't be used twice
|
||||
t.requests.Lock()
|
||||
defer t.requests.Unlock()
|
||||
|
||||
authRequest, ok := t.requests.data[token]
|
||||
if !ok {
|
||||
return "", "", errors.New("request is not found")
|
||||
}
|
||||
|
||||
if time.Now().After(authRequest.expires) {
|
||||
t.requests.Lock()
|
||||
delete(t.requests.data, token)
|
||||
t.requests.Unlock()
|
||||
return "", "", errors.New("request expired")
|
||||
}
|
||||
|
||||
@@ -247,10 +248,7 @@ func (t *Telegram) CheckToken(token, user string) (telegram, site string, err er
|
||||
return "", "", errors.New("user does not match original requester")
|
||||
}
|
||||
|
||||
// Delete request
|
||||
t.requests.Lock()
|
||||
delete(t.requests.data, token)
|
||||
t.requests.Unlock()
|
||||
|
||||
return authRequest.telegramID, authRequest.site, nil
|
||||
}
|
||||
@@ -258,7 +256,21 @@ func (t *Telegram) CheckToken(token, user string) (telegram, site string, err er
|
||||
// Run starts processing login requests sent in Telegram, required for user notifications to work
|
||||
// Blocks caller
|
||||
func (t *Telegram) Run(ctx context.Context) {
|
||||
atomic.AddInt32(&t.run, 1)
|
||||
t.updates.Lock()
|
||||
if t.updates.running {
|
||||
t.updates.Unlock()
|
||||
log.Print("[WARN] telegram updates processing is already running, ignoring the call")
|
||||
return
|
||||
}
|
||||
t.updates.running = true
|
||||
t.updates.Unlock()
|
||||
|
||||
defer func() {
|
||||
t.updates.Lock()
|
||||
t.updates.running = false
|
||||
t.updates.Unlock()
|
||||
}()
|
||||
|
||||
processUpdatedTicker := time.NewTicker(t.apiPollInterval)
|
||||
cleanupTicker := time.NewTicker(t.expiredCleanupInterval)
|
||||
|
||||
@@ -267,7 +279,6 @@ func (t *Telegram) Run(ctx context.Context) {
|
||||
case <-ctx.Done():
|
||||
processUpdatedTicker.Stop()
|
||||
cleanupTicker.Stop()
|
||||
atomic.AddInt32(&t.run, -1)
|
||||
return
|
||||
case <-processUpdatedTicker.C:
|
||||
updates, err := t.getUpdates(ctx)
|
||||
@@ -292,7 +303,12 @@ func (t *Telegram) Run(ctx context.Context) {
|
||||
// ProcessUpdate is alternative to Run, it processes provided plain text update from Telegram
|
||||
// so that caller could get updates and send it not only there but to multiple sources
|
||||
func (t *Telegram) ProcessUpdate(ctx context.Context, textUpdate string) error {
|
||||
if atomic.LoadInt32(&t.run) != 0 {
|
||||
// read lock is held for the whole call, so Run can't start in the middle of it,
|
||||
// while parallel ProcessUpdate calls are still allowed
|
||||
t.updates.RLock()
|
||||
defer t.updates.RUnlock()
|
||||
|
||||
if t.updates.running {
|
||||
return errors.New("the Run goroutine should not be used with ProcessUpdate")
|
||||
}
|
||||
defer func() {
|
||||
@@ -385,10 +401,18 @@ func (t *Telegram) processUpdates(ctx context.Context, updates *TelegramUpdate)
|
||||
|
||||
token := strings.TrimPrefix(update.Message.Text, "/start ")
|
||||
|
||||
t.requests.RLock()
|
||||
// confirmation is a single transaction, otherwise a request consumed by CheckToken
|
||||
// in the middle of it would be restored from the stale copy
|
||||
t.requests.Lock()
|
||||
authRequest, ok := t.requests.data[token]
|
||||
if !ok { // No such token
|
||||
t.requests.RUnlock()
|
||||
if ok {
|
||||
authRequest.confirmed = true
|
||||
authRequest.telegramID = strconv.Itoa(update.Message.Chat.ID)
|
||||
t.requests.data[token] = authRequest
|
||||
}
|
||||
t.requests.Unlock()
|
||||
|
||||
if !ok { // no such token
|
||||
if t.ErrorMsg != "" {
|
||||
if err := t.sendText(ctx, update.Message.Chat.ID, t.ErrorMsg); err != nil {
|
||||
log.Printf("[WARN] failed to notify telegram peer: %v", err)
|
||||
@@ -396,14 +420,6 @@ func (t *Telegram) processUpdates(ctx context.Context, updates *TelegramUpdate)
|
||||
}
|
||||
continue
|
||||
}
|
||||
t.requests.RUnlock()
|
||||
|
||||
authRequest.confirmed = true
|
||||
authRequest.telegramID = strconv.Itoa(update.Message.Chat.ID)
|
||||
|
||||
t.requests.Lock()
|
||||
t.requests.data[token] = authRequest
|
||||
t.requests.Unlock()
|
||||
|
||||
if err := t.sendText(ctx, update.Message.Chat.ID, t.SuccessMsg); err != nil {
|
||||
log.Printf("[ERROR] failed to notify telegram peer: %v", err)
|
||||
@@ -413,7 +429,7 @@ func (t *Telegram) processUpdates(ctx context.Context, updates *TelegramUpdate)
|
||||
|
||||
// sendText sends a plain text message to telegram peer
|
||||
func (t *Telegram) sendText(ctx context.Context, recipientID int, msg string) error {
|
||||
url := fmt.Sprintf("sendMessage?chat_id=%d&text=%s", recipientID, neturl.PathEscape(msg))
|
||||
url := fmt.Sprintf("sendMessage?chat_id=%d&text=%s", recipientID, neturl.QueryEscape(msg))
|
||||
return t.Request(ctx, url, nil, &struct{}{})
|
||||
}
|
||||
|
||||
@@ -436,7 +452,7 @@ func (t *Telegram) botInfo(ctx context.Context) (*TelegramBotInfo, error) {
|
||||
|
||||
// Request makes a request to the Telegram API and return the result
|
||||
func (t *Telegram) Request(ctx context.Context, method string, b []byte, data any) error {
|
||||
return repeater.NewDefault(3, time.Millisecond*250).Do(ctx, func() error {
|
||||
return repeater.NewFixed(3, time.Millisecond*250).Do(ctx, func() error {
|
||||
url := fmt.Sprintf("%s%s/%s", t.apiPrefix, t.Token, method)
|
||||
|
||||
var req *http.Request
|
||||
@@ -445,17 +461,18 @@ func (t *Telegram) Request(ctx context.Context, method string, b []byte, data an
|
||||
req, err = http.NewRequestWithContext(ctx, "GET", url, http.NoBody)
|
||||
} else {
|
||||
req, err = http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(b))
|
||||
req.Header.Set("Content-Type", "application/json; charset=utf-8")
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create request: %w", err)
|
||||
return fmt.Errorf("failed to create request: %w", t.redactToken(err))
|
||||
}
|
||||
if b != nil {
|
||||
req.Header.Set("Content-Type", "application/json; charset=utf-8")
|
||||
}
|
||||
|
||||
client := http.Client{Timeout: t.Timeout}
|
||||
defer client.CloseIdleConnections()
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to send request: %w", err)
|
||||
return fmt.Errorf("failed to send request: %w", t.redactToken(err))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
@@ -471,6 +488,16 @@ func (t *Telegram) Request(ctx context.Context, method string, b []byte, data an
|
||||
})
|
||||
}
|
||||
|
||||
// redactToken hides the bot token in the URL of *url.Error returned by the http client,
|
||||
// as the token is a part of every API URL and otherwise leaks into the logs of the caller printing the error
|
||||
func (t *Telegram) redactToken(err error) error {
|
||||
var urlErr *neturl.Error
|
||||
if t.Token == "" || !errors.As(err, &urlErr) || !strings.Contains(urlErr.URL, t.Token) {
|
||||
return err
|
||||
}
|
||||
return &neturl.Error{Op: urlErr.Op, URL: strings.ReplaceAll(urlErr.URL, t.Token, "<redacted>"), Err: urlErr.Err}
|
||||
}
|
||||
|
||||
func (t *Telegram) parseError(r io.Reader, statusCode int) error {
|
||||
tgErr := struct {
|
||||
Description string `json:"description"`
|
||||
|
||||
+20
-5
@@ -13,6 +13,13 @@ import (
|
||||
|
||||
const webhookTimeOut = 5000 * time.Millisecond
|
||||
|
||||
// limits for the response body of the webhook request: how much of it is quoted in the error
|
||||
// for non-OK responses, and how much is drained to make the connection reusable
|
||||
const (
|
||||
webhookErrBodyLimit = 16 * 1024
|
||||
webhookDrainBodyLimit = 64 * 1024
|
||||
)
|
||||
|
||||
// WebhookParams contain settings for webhook notifications
|
||||
type WebhookParams struct {
|
||||
Timeout time.Duration
|
||||
@@ -56,7 +63,7 @@ func (wh *Webhook) Send(ctx context.Context, destination, text string) error {
|
||||
}
|
||||
|
||||
for _, h := range wh.Headers {
|
||||
elems := strings.Split(h, ":")
|
||||
elems := strings.SplitN(h, ":", 2)
|
||||
if len(elems) != 2 {
|
||||
continue
|
||||
}
|
||||
@@ -67,14 +74,21 @@ func (wh *Webhook) Send(ctx context.Context, destination, text string) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("webhook request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() {
|
||||
// drain the remaining body to let the underlying connection be reused
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, webhookDrainBodyLimit))
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
|
||||
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
|
||||
errMsg := fmt.Sprintf("webhook request failed with non-OK status code: %d", resp.StatusCode)
|
||||
respBody, e := io.ReadAll(resp.Body)
|
||||
respBody, e := io.ReadAll(io.LimitReader(resp.Body, webhookErrBodyLimit+1))
|
||||
if e != nil {
|
||||
return errors.New(errMsg)
|
||||
}
|
||||
if len(respBody) > webhookErrBodyLimit {
|
||||
return fmt.Errorf("%s, body: %s... (truncated)", errMsg, respBody[:webhookErrBodyLimit])
|
||||
}
|
||||
return fmt.Errorf("%s, body: %s", errMsg, respBody)
|
||||
}
|
||||
|
||||
@@ -89,8 +103,9 @@ func (wh *Webhook) Schema() string {
|
||||
// String describes the webhook instance
|
||||
func (wh *Webhook) String() string {
|
||||
str := fmt.Sprintf("webhook notification with timeout %s", wh.Timeout)
|
||||
if wh.Headers != nil {
|
||||
str += fmt.Sprintf(" and headers %v", wh.Headers)
|
||||
if len(wh.Headers) != 0 {
|
||||
// header values are not printed as they might contain secrets, like authorization tokens
|
||||
str += fmt.Sprintf(" and %d headers", len(wh.Headers))
|
||||
}
|
||||
return str
|
||||
}
|
||||
|
||||
@@ -1,12 +0,0 @@
|
||||
# Binaries for programs and plugins
|
||||
*.exe
|
||||
*.exe~
|
||||
*.dll
|
||||
*.so
|
||||
*.dylib
|
||||
|
||||
# Test binary, build with `go test -c`
|
||||
*.test
|
||||
|
||||
# Output of the go coverage tool, specifically when used with LiteIDE
|
||||
*.out
|
||||
-83
@@ -1,83 +0,0 @@
|
||||
linters-settings:
|
||||
govet:
|
||||
check-shadowing: true
|
||||
golint:
|
||||
min-confidence: 0.6
|
||||
gocyclo:
|
||||
min-complexity: 15
|
||||
maligned:
|
||||
suggest-new: true
|
||||
dupl:
|
||||
threshold: 100
|
||||
goconst:
|
||||
min-len: 2
|
||||
min-occurrences: 2
|
||||
misspell:
|
||||
locale: US
|
||||
lll:
|
||||
line-length: 140
|
||||
gocritic:
|
||||
enabled-tags:
|
||||
- performance
|
||||
- style
|
||||
- experimental
|
||||
disabled-checks:
|
||||
- wrapperFunc
|
||||
- hugeParam
|
||||
- rangeValCopy
|
||||
|
||||
linters:
|
||||
disable-all: true
|
||||
enable:
|
||||
- megacheck
|
||||
- revive
|
||||
- govet
|
||||
- unconvert
|
||||
- gas
|
||||
- misspell
|
||||
- unused
|
||||
- typecheck
|
||||
- ineffassign
|
||||
- stylecheck
|
||||
- gochecknoinits
|
||||
- exportloopref
|
||||
- nakedret
|
||||
- gosimple
|
||||
- prealloc
|
||||
|
||||
fast: false
|
||||
|
||||
|
||||
run:
|
||||
# modules-download-mode: vendor
|
||||
skip-dirs:
|
||||
- vendor
|
||||
concurrency: 4
|
||||
|
||||
issues:
|
||||
exclude-rules:
|
||||
- text: "should have a package comment, unless it's in another file for this package"
|
||||
linters:
|
||||
- golint
|
||||
- text: "exitAfterDefer:"
|
||||
linters:
|
||||
- gocritic
|
||||
- text: "whyNoLint: include an explanation for nolint directive"
|
||||
linters:
|
||||
- gocritic
|
||||
- text: "go.mongodb.org/mongo-driver/bson/primitive.E"
|
||||
linters:
|
||||
- govet
|
||||
- text: "weak cryptographic primitive"
|
||||
linters:
|
||||
- gosec
|
||||
- text: "at least one file in a package should have a package comment"
|
||||
linters:
|
||||
- stylecheck
|
||||
- text: "should have a package comment"
|
||||
linters:
|
||||
- revive
|
||||
- text: 'Deferring unsafe method "Close" on type "io.ReadCloser"'
|
||||
linters:
|
||||
- gosec
|
||||
exclude-use-default: false
|
||||
-21
@@ -1,21 +0,0 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2023 Umputun
|
||||
|
||||
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.
|
||||
-33
@@ -1,33 +0,0 @@
|
||||
# Repeater [](https://github.com/go-pkgz/repeater/actions) [](https://goreportcard.com/report/github.com/go-pkgz/repeater) [](https://coveralls.io/github/go-pkgz/repeater?branch=master)
|
||||
|
||||
Repeater calls a function until it returns no error, up to some number of iterations and delays defined by strategy. It terminates immediately on err from the provided (optional) list of critical errors.
|
||||
|
||||
## Install and update
|
||||
|
||||
`go get -u github.com/go-pkgz/repeater`
|
||||
|
||||
## How to use
|
||||
|
||||
New Repeater created by `New(strtg strategy.Interface)` or shortcut for default - `NewDefault(repeats int, delay time.Duration) *Repeater`.
|
||||
|
||||
To activate invoke `Do` method. `Do` repeats func until no error returned. Predefined (optional) errors terminate the loop immediately.
|
||||
|
||||
`func (r Repeater) Do(ctx context.Context, fun func() error, errors ...error) (err error)`
|
||||
|
||||
### Repeating strategy
|
||||
|
||||
User can provide his own strategy implementing the interface:
|
||||
|
||||
```go
|
||||
type Interface interface {
|
||||
Start(ctx context.Context) chan struct{}
|
||||
}
|
||||
```
|
||||
|
||||
Returned channels used as "ticks," i.e., for each repeat or initial operation one read from this channel needed. Closing the channel indicates "done with retries." It is pretty much the same idea as `time.Timer` or `time.Tick` implements. Note - the first (technically not-repeated-yet) call won't happen **until something sent to the channel**. For this reason, the typical strategy sends the first "tick" before the first wait/sleep.
|
||||
|
||||
Three strategies provided byt the package:
|
||||
|
||||
1. **Fixed delay**, up to max number of attempts. It is the default strategy used by `repeater.NewDefault` constructor.
|
||||
2. **BackOff** with jitter provides an exponential backoff. It starts from `Duration` interval and goes in steps with `last * math.Pow(factor, attempt)`. Optional jitter randomizes intervals a little. _Factor = 1 effectively makes this strategy fixed with `Duration` delay._
|
||||
3. **Once** strategy does not do any repeats and mainly used for tests/mocks`.
|
||||
-69
@@ -1,69 +0,0 @@
|
||||
// Package repeater call fun till it returns no error, up to repeat some number of iterations and delays defined by strategy.
|
||||
// Repeats number and delays defined by strategy.Interface. Terminates immediately on err from
|
||||
// provided, optional list of critical errors
|
||||
package repeater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/go-pkgz/repeater/strategy"
|
||||
)
|
||||
|
||||
// Repeater is the main object, should be made by New or NewDefault, embeds strategy
|
||||
type Repeater struct {
|
||||
Strategy
|
||||
}
|
||||
|
||||
// Strategy interface for repeater strategy
|
||||
type Strategy interface {
|
||||
Start(ctx context.Context) <-chan struct{} // returns channel with repeater ticks
|
||||
}
|
||||
|
||||
// New repeater with a given strategy. If strategy=nil initializes with FixedDelay 5sec, 10 times.
|
||||
func New(strtg strategy.Interface) *Repeater {
|
||||
if strtg == nil {
|
||||
strtg = &strategy.FixedDelay{Repeats: 10, Delay: time.Second * 5}
|
||||
}
|
||||
result := Repeater{Strategy: strtg}
|
||||
return &result
|
||||
}
|
||||
|
||||
// NewDefault makes repeater with FixedDelay strategy
|
||||
func NewDefault(repeats int, delay time.Duration) *Repeater {
|
||||
return New(&strategy.FixedDelay{Repeats: repeats, Delay: delay})
|
||||
}
|
||||
|
||||
// Do repeats fun till no error. Predefined (optional) errors terminate immediately
|
||||
func (r Repeater) Do(ctx context.Context, fun func() error, errs ...error) (err error) {
|
||||
ctx, cancelFunc := context.WithCancel(ctx)
|
||||
defer cancelFunc() // ensure strategy's channel termination
|
||||
|
||||
inErrors := func(err error) bool {
|
||||
for _, e := range errs {
|
||||
if errors.Is(err, e) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
ch := r.Start(ctx) // channel of ticks-like events provided by strategy
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case _, ok := <-ch:
|
||||
if !ok { // closed channel indicates completion or early termination, set by strategy
|
||||
return err
|
||||
}
|
||||
if err = fun(); err == nil {
|
||||
return nil
|
||||
}
|
||||
if err != nil && inErrors(err) { // terminate on critical error from provided list
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
-59
@@ -1,59 +0,0 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"math"
|
||||
"math/rand"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Backoff implements strategy.Interface for exponential-backoff
|
||||
// it starts from 100ms (by default, if no Duration set) and goes in steps with last * math.Pow(factor, attempt)
|
||||
// optional jitter randomize intervals a little bit.
|
||||
type Backoff struct {
|
||||
Duration time.Duration
|
||||
Repeats int
|
||||
Factor float64
|
||||
Jitter bool
|
||||
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
// Start returns channel, similar to time.Timer
|
||||
// then publishing signals to channel ch for retries attempt. Closed ch indicates "done" event
|
||||
// consumer (repeater) should stop it explicitly after completion
|
||||
func (b *Backoff) Start(ctx context.Context) <-chan struct{} {
|
||||
|
||||
b.once.Do(func() {
|
||||
if b.Duration == 0 {
|
||||
b.Duration = 100 * time.Millisecond
|
||||
}
|
||||
if b.Repeats == 0 {
|
||||
b.Repeats = 1
|
||||
}
|
||||
if b.Factor <= 0 {
|
||||
b.Factor = 1
|
||||
}
|
||||
})
|
||||
|
||||
ch := make(chan struct{})
|
||||
go func() {
|
||||
defer close(ch)
|
||||
rnd := rand.New(rand.NewSource(int64(time.Now().Nanosecond()))) //nolint:gosec
|
||||
for i := 0; i < b.Repeats; i++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- struct{}{}:
|
||||
}
|
||||
|
||||
delay := float64(b.Duration) * math.Pow(b.Factor, float64(i))
|
||||
if b.Jitter {
|
||||
delay = rnd.Float64()*(float64(2*b.Duration)) + (delay - float64(b.Duration))
|
||||
}
|
||||
sleep(ctx, time.Duration(delay))
|
||||
}
|
||||
}()
|
||||
return ch
|
||||
}
|
||||
-36
@@ -1,36 +0,0 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// FixedDelay implements strategy.Interface for fixed intervals up to max repeats
|
||||
type FixedDelay struct {
|
||||
Repeats int
|
||||
Delay time.Duration
|
||||
}
|
||||
|
||||
// Start returns channel, similar to time.Timer
|
||||
// then publishing signals to channel ch for retries attempt.
|
||||
// can be terminated (canceled) via context.
|
||||
func (s *FixedDelay) Start(ctx context.Context) <-chan struct{} {
|
||||
if s.Repeats == 0 {
|
||||
s.Repeats = 1
|
||||
}
|
||||
ch := make(chan struct{})
|
||||
go func() {
|
||||
defer func() {
|
||||
close(ch)
|
||||
}()
|
||||
for i := 0; i < s.Repeats; i++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- struct{}{}:
|
||||
}
|
||||
sleep(ctx, s.Delay)
|
||||
}
|
||||
}()
|
||||
return ch
|
||||
}
|
||||
-35
@@ -1,35 +0,0 @@
|
||||
// Package strategy defines repeater's strategy and implements some.
|
||||
// Strategy result is a channel acting like time.Timer ot time.Tick
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Interface for repeater strategy. Returns channel with ticks
|
||||
type Interface interface {
|
||||
Start(ctx context.Context) <-chan struct{}
|
||||
}
|
||||
|
||||
// Once strategy eliminate repeats and makes a single try only
|
||||
type Once struct{}
|
||||
|
||||
// Start returns closed channel with a single element to prevent any repeats
|
||||
func (s *Once) Start(_ context.Context) <-chan struct{} {
|
||||
ch := make(chan struct{})
|
||||
go func() {
|
||||
ch <- struct{}{}
|
||||
close(ch)
|
||||
}()
|
||||
return ch
|
||||
}
|
||||
|
||||
func sleep(ctx context.Context, duration time.Duration) {
|
||||
select {
|
||||
case <-time.After(duration):
|
||||
return
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
+12
@@ -27,6 +27,18 @@ Use the links above for more information on each.
|
||||
|
||||
# changelog
|
||||
|
||||
* Jul 1st, 2026 [1.19.0](https://github.com/klauspost/compress/releases/tag/v1.19.0)
|
||||
* zstd: Add true concurrent stream encodingin https://github.com/klauspost/compress/pull/1136
|
||||
* zstd: arm64 decoder asm by @lizthegrey in https://github.com/klauspost/compress/pull/1160
|
||||
* flate: Add inflate checkpoints in https://github.com/klauspost/compress/pull/1154
|
||||
* zstd: avoid unused BuildDict encoder allocation by @snissn in https://github.com/klauspost/compress/pull/1147
|
||||
* snappy/s2: Limit length of varint in `decodedLen` by @eustas in https://github.com/klauspost/compress/pull/1148
|
||||
* gzhttp: match qvalue parameter case-insensitively (RFC 7231) by @z9z in https://github.com/klauspost/compress/pull/1149
|
||||
* zip: add NameDecoder callback for legacy encoding rewrite by @SAY-5 in https://github.com/klauspost/compress/pull/1150
|
||||
* huff0: Allow building tables from histogram in https://github.com/klauspost/compress/pull/1155
|
||||
* huff0: Allow building table from oversized histogram in https://github.com/klauspost/compress/pull/1156
|
||||
* s2sx: Clean symlink targets in https://github.com/klauspost/compress/pull/1163
|
||||
|
||||
* Feb 9th, 2026 [1.18.4](https://github.com/klauspost/compress/releases/tag/v1.18.4)
|
||||
* gzhttp: Add zstandard to server handler wrapper https://github.com/klauspost/compress/pull/1121
|
||||
* zstd: Add ResetWithOptions to encoder/decoder https://github.com/klauspost/compress/pull/1122
|
||||
|
||||
+168
@@ -0,0 +1,168 @@
|
||||
package huff0
|
||||
|
||||
import "errors"
|
||||
|
||||
// BuildCTable builds a Huffman compression table from a precomputed symbol
|
||||
// histogram and installs it as the previous (reuse) table on s.
|
||||
//
|
||||
// After this call:
|
||||
// - EstimateSize/CanUseTable can probe the table against other histograms.
|
||||
// - Compress1X/Compress4X with Reuse = ReusePolicyMust will encode without
|
||||
// emitting a new table header.
|
||||
// - TransferCTable can hand the table to a sibling Scratch.
|
||||
//
|
||||
// count[i] is the number of occurrences of symbol i. The histogram must have
|
||||
// at least 2 distinct non-zero symbols; ErrUseRLE is returned for a single
|
||||
// symbol and an error is returned for an empty histogram.
|
||||
func (s *Scratch) BuildCTable(count *[256]uint32) error {
|
||||
if s == nil {
|
||||
return errors.New("huff0: BuildCTable on nil Scratch")
|
||||
}
|
||||
if count == nil {
|
||||
return errors.New("huff0: nil count passed to BuildCTable")
|
||||
}
|
||||
var err error
|
||||
s, err = s.prepare(nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.count = *count
|
||||
var total, maxCount int
|
||||
var symLen uint16
|
||||
for i, v := range s.count {
|
||||
total += int(v)
|
||||
if int(v) > maxCount {
|
||||
maxCount = int(v)
|
||||
}
|
||||
if v != 0 {
|
||||
symLen = uint16(i) + 1
|
||||
}
|
||||
}
|
||||
if total == 0 {
|
||||
return errors.New("huff0: empty histogram")
|
||||
}
|
||||
if symLen < 2 || maxCount == total {
|
||||
return ErrUseRLE
|
||||
}
|
||||
// huff0's internal rank table assumes total ≤ BlockSizeMax (it uses
|
||||
// highBit32(count+1) + 1 as a rank index into a fixed-size array).
|
||||
// Histograms summed across multiple blocks can exceed that; scale the
|
||||
// counts down preserving the distribution. Non-zero entries round up so
|
||||
// rare symbols stay representable.
|
||||
if total > BlockSizeMax {
|
||||
shift := uint(0)
|
||||
for total>>shift > BlockSizeMax {
|
||||
shift++
|
||||
}
|
||||
round := uint32(1<<shift) - 1
|
||||
var newTotal, newMax int
|
||||
for i, v := range s.count {
|
||||
if v == 0 {
|
||||
continue
|
||||
}
|
||||
scaled := (v + round) >> shift
|
||||
if scaled == 0 {
|
||||
scaled = 1
|
||||
}
|
||||
s.count[i] = scaled
|
||||
newTotal += int(scaled)
|
||||
if int(scaled) > newMax {
|
||||
newMax = int(scaled)
|
||||
}
|
||||
}
|
||||
total = newTotal
|
||||
maxCount = newMax
|
||||
if maxCount == total {
|
||||
return ErrUseRLE
|
||||
}
|
||||
}
|
||||
s.symbolLen = symLen
|
||||
s.maxCount = maxCount
|
||||
s.srcLen = total
|
||||
if err := s.buildCTable(); err != nil {
|
||||
return err
|
||||
}
|
||||
if cap(s.prevTable) < len(s.cTable) {
|
||||
s.prevTable = make(cTable, 0, maxSymbolValue+1)
|
||||
}
|
||||
s.prevTable = s.prevTable[:len(s.cTable)]
|
||||
copy(s.prevTable, s.cTable)
|
||||
s.prevTableLog = s.actualTableLog
|
||||
// Force the next Compress* to recount from real input.
|
||||
s.clearCount = true
|
||||
s.maxCount = 0
|
||||
return nil
|
||||
}
|
||||
|
||||
// EstimateSize returns an estimated compressed payload size in bytes for the
|
||||
// supplied histogram using the table currently stored in prevTable. It returns
|
||||
// -1 when the table cannot encode every non-zero symbol of hist (i.e. when
|
||||
// CanUseTable would return false). The estimate excludes the table header.
|
||||
func (s *Scratch) EstimateSize(hist *[256]uint32) int {
|
||||
if s == nil || hist == nil || len(s.prevTable) == 0 {
|
||||
return -1
|
||||
}
|
||||
pt := s.prevTable
|
||||
nbBits := uint32(7)
|
||||
for i, v := range hist {
|
||||
if v == 0 {
|
||||
continue
|
||||
}
|
||||
if i >= len(pt) || pt[i].nBits == 0 {
|
||||
return -1
|
||||
}
|
||||
nbBits += uint32(pt[i].nBits) * v
|
||||
}
|
||||
return int(nbBits >> 3)
|
||||
}
|
||||
|
||||
// CanUseTable reports whether the table in prevTable can encode every
|
||||
// non-zero symbol present in hist.
|
||||
func (s *Scratch) CanUseTable(hist *[256]uint32) bool {
|
||||
if s == nil || hist == nil || len(s.prevTable) == 0 {
|
||||
return false
|
||||
}
|
||||
pt := s.prevTable
|
||||
for i, v := range hist {
|
||||
if v == 0 {
|
||||
continue
|
||||
}
|
||||
if i >= len(pt) || pt[i].nBits == 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// AppendTable serializes the table currently stored in prevTable (e.g. as
|
||||
// installed by BuildCTable or carried over from a previous Compress call)
|
||||
// into a self-delimiting zstd-style header and appends it to dst. The
|
||||
// returned slice can be parsed back by ReadTable.
|
||||
func (s *Scratch) AppendTable(dst []byte) ([]byte, error) {
|
||||
if s == nil || len(s.prevTable) == 0 {
|
||||
return dst, errors.New("huff0: AppendTable with empty table")
|
||||
}
|
||||
// cTable.write reads s.actualTableLog, s.symbolLen, s.huffWeight, s.fse
|
||||
// and writes into s.Out. Save/restore Out so we don't disturb in-flight
|
||||
// compression buffers.
|
||||
saveOut := s.Out
|
||||
saveTL := s.actualTableLog
|
||||
saveSL := s.symbolLen
|
||||
if s.fse == nil {
|
||||
// Lazily init in case AppendTable is called on a fresh Scratch.
|
||||
if _, err := s.prepare(nil); err != nil {
|
||||
return dst, err
|
||||
}
|
||||
saveOut = s.Out
|
||||
}
|
||||
s.Out = s.Out[:0]
|
||||
s.actualTableLog = s.prevTableLog
|
||||
s.symbolLen = uint16(len(s.prevTable))
|
||||
if err := s.prevTable.write(s); err != nil {
|
||||
s.Out, s.actualTableLog, s.symbolLen = saveOut, saveTL, saveSL
|
||||
return dst, err
|
||||
}
|
||||
dst = append(dst, s.Out...)
|
||||
s.Out, s.actualTableLog, s.symbolLen = saveOut, saveTL, saveSL
|
||||
return dst, nil
|
||||
}
|
||||
+15
-188
@@ -1,222 +1,49 @@
|
||||
//go:build amd64 && !appengine && !noasm && gc
|
||||
|
||||
// This file contains the specialisation of Decoder.Decompress4X
|
||||
// and Decoder.Decompress1X that use an asm implementation of thir main loops.
|
||||
// amd64 stubs and dispatch for the asm loops used by decompress_asm.go.
|
||||
package huff0
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/klauspost/compress/internal/cpuinfo"
|
||||
)
|
||||
|
||||
// decompress4x_main_loop_x86 is an x86 assembler implementation
|
||||
// decompress4x_main_loop_amd64 is an x86 assembler implementation
|
||||
// of Decompress4X when tablelog > 8.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress4x_main_loop_amd64(ctx *decompress4xContext)
|
||||
|
||||
// decompress4x_8b_loop_x86 is an x86 assembler implementation
|
||||
// decompress4x_8b_main_loop_amd64 is an x86 assembler implementation
|
||||
// of Decompress4X when tablelog <= 8 which decodes 4 entries
|
||||
// per loop.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress4x_8b_main_loop_amd64(ctx *decompress4xContext)
|
||||
|
||||
// fallback8BitSize is the size where using Go version is faster.
|
||||
const fallback8BitSize = 800
|
||||
|
||||
type decompress4xContext struct {
|
||||
pbr *[4]bitReaderShifted
|
||||
peekBits uint8
|
||||
out *byte
|
||||
dstEvery int
|
||||
tbl *dEntrySingle
|
||||
decoded int
|
||||
limit *byte
|
||||
}
|
||||
|
||||
// Decompress4X will decompress a 4X encoded stream.
|
||||
// The length of the supplied input must match the end of a block exactly.
|
||||
// The *capacity* of the dst slice must match the destination size of
|
||||
// the uncompressed data exactly.
|
||||
func (d *Decoder) Decompress4X(dst, src []byte) ([]byte, error) {
|
||||
if len(d.dt.single) == 0 {
|
||||
return nil, errors.New("no table loaded")
|
||||
}
|
||||
if len(src) < 6+(4*1) {
|
||||
return nil, errors.New("input too small")
|
||||
}
|
||||
|
||||
use8BitTables := d.actualTableLog <= 8
|
||||
if cap(dst) < fallback8BitSize && use8BitTables {
|
||||
return d.decompress4X8bit(dst, src)
|
||||
}
|
||||
|
||||
var br [4]bitReaderShifted
|
||||
// Decode "jump table"
|
||||
start := 6
|
||||
for i := range 3 {
|
||||
length := int(src[i*2]) | (int(src[i*2+1]) << 8)
|
||||
if start+length >= len(src) {
|
||||
return nil, errors.New("truncated input (or invalid offset)")
|
||||
}
|
||||
err := br[i].init(src[start : start+length])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
start += length
|
||||
}
|
||||
err := br[3].init(src[start:])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// destination, offset to match first output
|
||||
dstSize := cap(dst)
|
||||
dst = dst[:dstSize]
|
||||
out := dst
|
||||
dstEvery := (dstSize + 3) / 4
|
||||
|
||||
const tlSize = 1 << tableLogMax
|
||||
const tlMask = tlSize - 1
|
||||
single := d.dt.single[:tlSize]
|
||||
|
||||
var decoded int
|
||||
|
||||
if len(out) > 4*4 && !(br[0].off < 4 || br[1].off < 4 || br[2].off < 4 || br[3].off < 4) {
|
||||
ctx := decompress4xContext{
|
||||
pbr: &br,
|
||||
peekBits: uint8((64 - d.actualTableLog) & 63), // see: bitReaderShifted.peekBitsFast()
|
||||
out: &out[0],
|
||||
dstEvery: dstEvery,
|
||||
tbl: &single[0],
|
||||
limit: &out[dstEvery-4], // Always stop decoding when first buffer gets here to avoid writing OOB on last.
|
||||
}
|
||||
if use8BitTables {
|
||||
decompress4x_8b_main_loop_amd64(&ctx)
|
||||
} else {
|
||||
decompress4x_main_loop_amd64(&ctx)
|
||||
}
|
||||
|
||||
decoded = ctx.decoded
|
||||
out = out[decoded/4:]
|
||||
}
|
||||
|
||||
// Decode remaining.
|
||||
remainBytes := dstEvery - (decoded / 4)
|
||||
for i := range br {
|
||||
offset := dstEvery * i
|
||||
endsAt := min(offset+remainBytes, len(out))
|
||||
br := &br[i]
|
||||
bitsLeft := br.remaining()
|
||||
for bitsLeft > 0 {
|
||||
br.fill()
|
||||
if offset >= endsAt {
|
||||
return nil, errors.New("corruption detected: stream overrun 4")
|
||||
}
|
||||
|
||||
// Read value and increment offset.
|
||||
val := br.peekBitsFast(d.actualTableLog)
|
||||
v := single[val&tlMask].entry
|
||||
nBits := uint8(v)
|
||||
br.advance(nBits)
|
||||
bitsLeft -= uint(nBits)
|
||||
out[offset] = uint8(v >> 8)
|
||||
offset++
|
||||
}
|
||||
if offset != endsAt {
|
||||
return nil, fmt.Errorf("corruption detected: short output block %d, end %d != %d", i, offset, endsAt)
|
||||
}
|
||||
decoded += offset - dstEvery*i
|
||||
err = br.close()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if dstSize != decoded {
|
||||
return nil, errors.New("corruption detected: short output block")
|
||||
}
|
||||
return dst, nil
|
||||
}
|
||||
|
||||
// decompress4x_main_loop_x86 is an x86 assembler implementation
|
||||
// decompress1x_main_loop_amd64 is an x86 assembler implementation
|
||||
// of Decompress1X when tablelog > 8.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress1x_main_loop_amd64(ctx *decompress1xContext)
|
||||
|
||||
// decompress4x_main_loop_x86 is an x86 with BMI2 assembler implementation
|
||||
// decompress1x_main_loop_bmi2 is an x86 with BMI2 assembler implementation
|
||||
// of Decompress1X when tablelog > 8.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress1x_main_loop_bmi2(ctx *decompress1xContext)
|
||||
|
||||
type decompress1xContext struct {
|
||||
pbr *bitReaderShifted
|
||||
peekBits uint8
|
||||
out *byte
|
||||
outCap int
|
||||
tbl *dEntrySingle
|
||||
decoded int
|
||||
func decompress4x_main_loop_asm(ctx *decompress4xContext) {
|
||||
decompress4x_main_loop_amd64(ctx)
|
||||
}
|
||||
|
||||
// Error reported by asm implementations
|
||||
const error_max_decoded_size_exeeded = -1
|
||||
func decompress4x_8b_main_loop_asm(ctx *decompress4xContext) {
|
||||
decompress4x_8b_main_loop_amd64(ctx)
|
||||
}
|
||||
|
||||
// Decompress1X will decompress a 1X encoded stream.
|
||||
// The cap of the output buffer will be the maximum decompressed size.
|
||||
// The length of the supplied input must match the end of a block exactly.
|
||||
func (d *Decoder) Decompress1X(dst, src []byte) ([]byte, error) {
|
||||
if len(d.dt.single) == 0 {
|
||||
return nil, errors.New("no table loaded")
|
||||
func decompress1x_main_loop_asm(ctx *decompress1xContext) {
|
||||
if cpuinfo.HasBMI2() {
|
||||
decompress1x_main_loop_bmi2(ctx)
|
||||
} else {
|
||||
decompress1x_main_loop_amd64(ctx)
|
||||
}
|
||||
var br bitReaderShifted
|
||||
err := br.init(src)
|
||||
if err != nil {
|
||||
return dst, err
|
||||
}
|
||||
maxDecodedSize := cap(dst)
|
||||
dst = dst[:maxDecodedSize]
|
||||
|
||||
const tlSize = 1 << tableLogMax
|
||||
const tlMask = tlSize - 1
|
||||
|
||||
if maxDecodedSize >= 4 {
|
||||
ctx := decompress1xContext{
|
||||
pbr: &br,
|
||||
out: &dst[0],
|
||||
outCap: maxDecodedSize,
|
||||
peekBits: uint8((64 - d.actualTableLog) & 63), // see: bitReaderShifted.peekBitsFast()
|
||||
tbl: &d.dt.single[0],
|
||||
}
|
||||
|
||||
if cpuinfo.HasBMI2() {
|
||||
decompress1x_main_loop_bmi2(&ctx)
|
||||
} else {
|
||||
decompress1x_main_loop_amd64(&ctx)
|
||||
}
|
||||
if ctx.decoded == error_max_decoded_size_exeeded {
|
||||
return nil, ErrMaxDecodedSizeExceeded
|
||||
}
|
||||
|
||||
dst = dst[:ctx.decoded]
|
||||
}
|
||||
|
||||
// br < 8, so uint8 is fine
|
||||
bitsLeft := uint8(br.off)*8 + 64 - br.bitsRead
|
||||
for bitsLeft > 0 {
|
||||
br.fill()
|
||||
if len(dst) >= maxDecodedSize {
|
||||
br.close()
|
||||
return nil, ErrMaxDecodedSizeExceeded
|
||||
}
|
||||
v := d.dt.single[br.peekBitsFast(d.actualTableLog)&tlMask]
|
||||
nBits := uint8(v.entry)
|
||||
br.advance(nBits)
|
||||
bitsLeft -= nBits
|
||||
dst = append(dst, uint8(v.entry>>8))
|
||||
}
|
||||
return dst, br.close()
|
||||
}
|
||||
|
||||
+2
-2
@@ -1,6 +1,6 @@
|
||||
// Code generated by command: go run gen.go -out ../decompress_amd64.s -pkg=huff0. DO NOT EDIT.
|
||||
// Code generated by command: go run gen.go -out ../decompress.s -arch amd64,arm64 -pkg=huff0. DO NOT EDIT.
|
||||
|
||||
//go:build amd64 && !appengine && !noasm && gc
|
||||
//go:build !appengine && !noasm && gc
|
||||
|
||||
// func decompress4x_main_loop_amd64(ctx *decompress4xContext)
|
||||
TEXT ·decompress4x_main_loop_amd64(SB), $0-8
|
||||
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
//go:build arm64 && !appengine && !noasm && gc
|
||||
|
||||
// arm64 stubs and dispatch for the asm loops used by decompress_asm.go.
|
||||
// The asm (decompress_arm64.s) is generated by the avo arm64 lowering
|
||||
// printer from the same source as the amd64 asm; see _generate/gen.go.
|
||||
package huff0
|
||||
|
||||
// decompress4x_main_loop_arm64 is an arm64 assembler implementation
|
||||
// of Decompress4X when tablelog > 8.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress4x_main_loop_arm64(ctx *decompress4xContext)
|
||||
|
||||
// decompress4x_8b_main_loop_arm64 is an arm64 assembler implementation
|
||||
// of Decompress4X when tablelog <= 8 which decodes 4 entries
|
||||
// per loop.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress4x_8b_main_loop_arm64(ctx *decompress4xContext)
|
||||
|
||||
// decompress1x_main_loop_arm64 is an arm64 assembler implementation
|
||||
// of Decompress1X when tablelog > 8.
|
||||
//
|
||||
//go:noescape
|
||||
func decompress1x_main_loop_arm64(ctx *decompress1xContext)
|
||||
|
||||
func decompress4x_main_loop_asm(ctx *decompress4xContext) {
|
||||
decompress4x_main_loop_arm64(ctx)
|
||||
}
|
||||
|
||||
func decompress4x_8b_main_loop_asm(ctx *decompress4xContext) {
|
||||
decompress4x_8b_main_loop_arm64(ctx)
|
||||
}
|
||||
|
||||
func decompress1x_main_loop_asm(ctx *decompress1xContext) {
|
||||
decompress1x_main_loop_arm64(ctx)
|
||||
}
|
||||
+851
@@ -0,0 +1,851 @@
|
||||
// Code generated by command: go run gen.go -out ../decompress.s -arch amd64,arm64 -pkg=huff0. DO NOT EDIT.
|
||||
// EXPERIMENTAL arm64 output lowered from an amd64 avo program.
|
||||
|
||||
//go:build arm64 && !appengine && !noasm && gc
|
||||
|
||||
// func decompress4x_main_loop_amd64(ctx *decompress4xContext)
|
||||
TEXT ·decompress4x_main_loop_arm64(SB), $0-8
|
||||
// Preload values
|
||||
MOVD ctx+0(FP), R0
|
||||
MOVBU 8(R0), R6
|
||||
MOVD 16(R0), R3
|
||||
MOVD 48(R0), R5
|
||||
MOVD 24(R0), R7
|
||||
MOVD 32(R0), R8
|
||||
MOVD (R0), R9
|
||||
|
||||
// Main loop
|
||||
main_loop:
|
||||
MOVD $0, R2
|
||||
CMP R5, R3
|
||||
CSET GE, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
// br0.fillFast32()
|
||||
MOVD 32(R9), R10
|
||||
MOVBU 40(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill0
|
||||
MOVD 24(R9), R0
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R0, R0
|
||||
MOVD (R9), R12
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R12, R0, R15
|
||||
MOVWU (R15), R12
|
||||
MOVD R11, R1
|
||||
LSL R1, R12, R12
|
||||
MOVD R0, 24(R9)
|
||||
ORR R12, R10, R10
|
||||
|
||||
// exhausted += (br0.off < 4)
|
||||
CMP $0x04, R0
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill0:
|
||||
// val0 := br0.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br0.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br0.peekTopBits(peekBits)
|
||||
MOVD R6, R1
|
||||
MOVD R10, R12
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val1&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br0.advance(uint8(v1.entry))
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// these two writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
MOVH R0, (R3)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 32(R9)
|
||||
MOVB R11, 40(R9)
|
||||
|
||||
// br1.fillFast32()
|
||||
MOVD 80(R9), R10
|
||||
MOVBU 88(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill1
|
||||
MOVD 72(R9), R0
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R0, R0
|
||||
MOVD 48(R9), R12
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R12, R0, R15
|
||||
MOVWU (R15), R12
|
||||
MOVD R11, R1
|
||||
LSL R1, R12, R12
|
||||
MOVD R0, 72(R9)
|
||||
ORR R12, R10, R10
|
||||
|
||||
// exhausted += (br1.off < 4)
|
||||
CMP $0x04, R0
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill1:
|
||||
// val0 := br1.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br1.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br1.peekTopBits(peekBits)
|
||||
MOVD R6, R1
|
||||
MOVD R10, R12
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val1&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br1.advance(uint8(v1.entry))
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// these two writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
ADD R7, R3, R15
|
||||
MOVH R0, (R15)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 80(R9)
|
||||
MOVB R11, 88(R9)
|
||||
|
||||
// br2.fillFast32()
|
||||
MOVD 128(R9), R10
|
||||
MOVBU 136(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill2
|
||||
MOVD 120(R9), R0
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R0, R0
|
||||
MOVD 96(R9), R12
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R12, R0, R15
|
||||
MOVWU (R15), R12
|
||||
MOVD R11, R1
|
||||
LSL R1, R12, R12
|
||||
MOVD R0, 120(R9)
|
||||
ORR R12, R10, R10
|
||||
|
||||
// exhausted += (br2.off < 4)
|
||||
CMP $0x04, R0
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill2:
|
||||
// val0 := br2.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br2.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br2.peekTopBits(peekBits)
|
||||
MOVD R6, R1
|
||||
MOVD R10, R12
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val1&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br2.advance(uint8(v1.entry))
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// these two writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
ADD R7<<1, R3, R15
|
||||
MOVH R0, (R15)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 128(R9)
|
||||
MOVB R11, 136(R9)
|
||||
|
||||
// br3.fillFast32()
|
||||
MOVD 176(R9), R10
|
||||
MOVBU 184(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill3
|
||||
MOVD 168(R9), R0
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R0, R0
|
||||
MOVD 144(R9), R12
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R12, R0, R15
|
||||
MOVWU (R15), R12
|
||||
MOVD R11, R1
|
||||
LSL R1, R12, R12
|
||||
MOVD R0, 168(R9)
|
||||
ORR R12, R10, R10
|
||||
|
||||
// exhausted += (br3.off < 4)
|
||||
CMP $0x04, R0
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill3:
|
||||
// val0 := br3.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br3.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br3.peekTopBits(peekBits)
|
||||
MOVD R6, R1
|
||||
MOVD R10, R12
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val1&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br3.advance(uint8(v1.entry))
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// these two writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
ADD R7<<1, R7, R1
|
||||
ADD R1, R3, R15
|
||||
MOVH R0, (R15)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 176(R9)
|
||||
MOVB R11, 184(R9)
|
||||
ADD $0x02, R3, R3
|
||||
AND $0xff, R2, R15
|
||||
AND $0xff, R2, R16
|
||||
TST R16, R15
|
||||
BEQ main_loop
|
||||
MOVD ctx+0(FP), R0
|
||||
MOVD 16(R0), R16
|
||||
SUB R16, R3, R3
|
||||
LSL $0x02, R3, R3
|
||||
MOVD R3, 40(R0)
|
||||
RET
|
||||
|
||||
// func decompress4x_8b_main_loop_amd64(ctx *decompress4xContext)
|
||||
TEXT ·decompress4x_8b_main_loop_arm64(SB), $0-8
|
||||
// Preload values
|
||||
MOVD ctx+0(FP), R1
|
||||
MOVBU 8(R1), R6
|
||||
MOVD 16(R1), R3
|
||||
MOVD 48(R1), R5
|
||||
MOVD 24(R1), R7
|
||||
MOVD 32(R1), R8
|
||||
MOVD (R1), R9
|
||||
|
||||
// Main loop
|
||||
main_loop:
|
||||
MOVD $0, R2
|
||||
CMP R5, R3
|
||||
CSET GE, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
// br0.fillFast32()
|
||||
MOVD 32(R9), R10
|
||||
MOVBU 40(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill0
|
||||
MOVD 24(R9), R12
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R12, R12
|
||||
MOVD (R9), R13
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R13, R12, R15
|
||||
MOVWU (R15), R13
|
||||
MOVD R11, R1
|
||||
LSL R1, R13, R13
|
||||
MOVD R12, 24(R9)
|
||||
ORR R13, R10, R10
|
||||
|
||||
// exhausted += (br0.off < 4)
|
||||
CMP $0x04, R12
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill0:
|
||||
// val0 := br0.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br0.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br0.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br0.advance(uint8(v1.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// val2 := br0.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v2 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br0.advance(uint8(v2.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val3 := br0.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v3 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br0.advance(uint8(v3.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// these four writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
// out[id * dstEvery + 3] = uint8(v2.entry >> 8)
|
||||
// out[id * dstEvery + 4] = uint8(v3.entry >> 8)
|
||||
MOVW R0, (R3)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 32(R9)
|
||||
MOVB R11, 40(R9)
|
||||
|
||||
// br1.fillFast32()
|
||||
MOVD 80(R9), R10
|
||||
MOVBU 88(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill1
|
||||
MOVD 72(R9), R12
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R12, R12
|
||||
MOVD 48(R9), R13
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R13, R12, R15
|
||||
MOVWU (R15), R13
|
||||
MOVD R11, R1
|
||||
LSL R1, R13, R13
|
||||
MOVD R12, 72(R9)
|
||||
ORR R13, R10, R10
|
||||
|
||||
// exhausted += (br1.off < 4)
|
||||
CMP $0x04, R12
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill1:
|
||||
// val0 := br1.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br1.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br1.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br1.advance(uint8(v1.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// val2 := br1.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v2 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br1.advance(uint8(v2.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val3 := br1.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v3 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br1.advance(uint8(v3.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// these four writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
// out[id * dstEvery + 3] = uint8(v2.entry >> 8)
|
||||
// out[id * dstEvery + 4] = uint8(v3.entry >> 8)
|
||||
ADD R7, R3, R15
|
||||
MOVW R0, (R15)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 80(R9)
|
||||
MOVB R11, 88(R9)
|
||||
|
||||
// br2.fillFast32()
|
||||
MOVD 128(R9), R10
|
||||
MOVBU 136(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill2
|
||||
MOVD 120(R9), R12
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R12, R12
|
||||
MOVD 96(R9), R13
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R13, R12, R15
|
||||
MOVWU (R15), R13
|
||||
MOVD R11, R1
|
||||
LSL R1, R13, R13
|
||||
MOVD R12, 120(R9)
|
||||
ORR R13, R10, R10
|
||||
|
||||
// exhausted += (br2.off < 4)
|
||||
CMP $0x04, R12
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill2:
|
||||
// val0 := br2.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br2.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br2.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br2.advance(uint8(v1.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// val2 := br2.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v2 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br2.advance(uint8(v2.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val3 := br2.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v3 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br2.advance(uint8(v3.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// these four writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
// out[id * dstEvery + 3] = uint8(v2.entry >> 8)
|
||||
// out[id * dstEvery + 4] = uint8(v3.entry >> 8)
|
||||
ADD R7<<1, R3, R15
|
||||
MOVW R0, (R15)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 128(R9)
|
||||
MOVB R11, 136(R9)
|
||||
|
||||
// br3.fillFast32()
|
||||
MOVD 176(R9), R10
|
||||
MOVBU 184(R9), R11
|
||||
CMP $0x20, R11
|
||||
BLS skip_fill3
|
||||
MOVD 168(R9), R12
|
||||
SUB $0x20, R11, R11
|
||||
SUB $0x04, R12, R12
|
||||
MOVD 144(R9), R13
|
||||
|
||||
// b.value |= uint64(low) << (b.bitsRead & 63)
|
||||
ADD R13, R12, R15
|
||||
MOVWU (R15), R13
|
||||
MOVD R11, R1
|
||||
LSL R1, R13, R13
|
||||
MOVD R12, 168(R9)
|
||||
ORR R13, R10, R10
|
||||
|
||||
// exhausted += (br3.off < 4)
|
||||
CMP $0x04, R12
|
||||
CSINC HS, R2, R2, R16
|
||||
BFI $0, R16, $8, R2
|
||||
|
||||
skip_fill3:
|
||||
// val0 := br3.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v0 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br3.advance(uint8(v0.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val1 := br3.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v1 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br3.advance(uint8(v1.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// val2 := br3.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v2 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br3.advance(uint8(v2.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
|
||||
// val3 := br3.peekTopBits(peekBits)
|
||||
MOVD R10, R12
|
||||
MOVD R6, R1
|
||||
LSR R1, R12, R12
|
||||
|
||||
// v3 := table[val0&mask]
|
||||
ADD R12<<1, R8, R15
|
||||
MOVHU (R15), R1
|
||||
|
||||
// br3.advance(uint8(v3.entry)
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
LSL R1, R10, R10
|
||||
ADD R1, R11, R15
|
||||
BFI $0, R15, $8, R11
|
||||
REVW R0, R0
|
||||
|
||||
// these four writes get coalesced
|
||||
// out[id * dstEvery + 0] = uint8(v0.entry >> 8)
|
||||
// out[id * dstEvery + 1] = uint8(v1.entry >> 8)
|
||||
// out[id * dstEvery + 3] = uint8(v2.entry >> 8)
|
||||
// out[id * dstEvery + 4] = uint8(v3.entry >> 8)
|
||||
ADD R7<<1, R7, R1
|
||||
ADD R1, R3, R15
|
||||
MOVW R0, (R15)
|
||||
|
||||
// update the bitreader structure
|
||||
MOVD R10, 176(R9)
|
||||
MOVB R11, 184(R9)
|
||||
ADD $0x04, R3, R3
|
||||
AND $0xff, R2, R15
|
||||
AND $0xff, R2, R16
|
||||
TST R16, R15
|
||||
BEQ main_loop
|
||||
MOVD ctx+0(FP), R0
|
||||
MOVD 16(R0), R16
|
||||
SUB R16, R3, R3
|
||||
LSL $0x02, R3, R3
|
||||
MOVD R3, 40(R0)
|
||||
RET
|
||||
|
||||
// func decompress1x_main_loop_amd64(ctx *decompress1xContext)
|
||||
TEXT ·decompress1x_main_loop_arm64(SB), $0-8
|
||||
MOVD ctx+0(FP), R1
|
||||
MOVD 16(R1), R2
|
||||
MOVD 24(R1), R3
|
||||
CMP $0x04, R3
|
||||
BLO error_max_decoded_size_exceeded
|
||||
ADD R3, R2, R3
|
||||
MOVD (R1), R5
|
||||
MOVD (R5), R7
|
||||
MOVD 24(R5), R8
|
||||
MOVD 32(R5), R9
|
||||
MOVBU 40(R5), R10
|
||||
MOVD 32(R1), R5
|
||||
MOVBU 8(R1), R6
|
||||
JMP loop_condition
|
||||
|
||||
main_loop:
|
||||
// Check if we have room for 4 bytes in the output buffer
|
||||
ADD $4, R2, R1
|
||||
CMP R3, R1
|
||||
BGE error_max_decoded_size_exceeded
|
||||
|
||||
// Decode 4 values
|
||||
CMP $0x20, R10
|
||||
BLT bitReader_fillFast_1_end
|
||||
SUB $0x20, R10, R10
|
||||
SUB $0x04, R8, R8
|
||||
ADD R8, R7, R15
|
||||
MOVWU (R15), R11
|
||||
MOVD R10, R1
|
||||
LSL R1, R11, R11
|
||||
ORR R11, R9, R9
|
||||
|
||||
bitReader_fillFast_1_end:
|
||||
MOVD R6, R1
|
||||
MOVD R9, R11
|
||||
LSR R1, R11, R11
|
||||
ADD R11<<1, R5, R15
|
||||
MOVHU (R15), R1
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
MOVBU R1, R1
|
||||
ADD R1, R10, R10
|
||||
LSL R1, R9, R9
|
||||
MOVD R6, R1
|
||||
MOVD R9, R11
|
||||
LSR R1, R11, R11
|
||||
ADD R11<<1, R5, R15
|
||||
MOVHU (R15), R1
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
MOVBU R1, R1
|
||||
ADD R1, R10, R10
|
||||
LSL R1, R9, R9
|
||||
REVW R0, R0
|
||||
CMP $0x20, R10
|
||||
BLT bitReader_fillFast_2_end
|
||||
SUB $0x20, R10, R10
|
||||
SUB $0x04, R8, R8
|
||||
ADD R8, R7, R15
|
||||
MOVWU (R15), R11
|
||||
MOVD R10, R1
|
||||
LSL R1, R11, R11
|
||||
ORR R11, R9, R9
|
||||
|
||||
bitReader_fillFast_2_end:
|
||||
MOVD R6, R1
|
||||
MOVD R9, R11
|
||||
LSR R1, R11, R11
|
||||
ADD R11<<1, R5, R15
|
||||
MOVHU (R15), R1
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $8, R16, $8, R0
|
||||
MOVBU R1, R1
|
||||
ADD R1, R10, R10
|
||||
LSL R1, R9, R9
|
||||
MOVD R6, R1
|
||||
MOVD R9, R11
|
||||
LSR R1, R11, R11
|
||||
ADD R11<<1, R5, R15
|
||||
MOVHU (R15), R1
|
||||
UBFX $8, R1, $8, R16
|
||||
BFI $0, R16, $8, R0
|
||||
MOVBU R1, R1
|
||||
ADD R1, R10, R10
|
||||
LSL R1, R9, R9
|
||||
REVW R0, R0
|
||||
|
||||
// Store the decoded values
|
||||
MOVW R0, (R2)
|
||||
ADD $0x04, R2, R2
|
||||
|
||||
loop_condition:
|
||||
CMP $0x08, R8
|
||||
BGE main_loop
|
||||
|
||||
// Update ctx structure
|
||||
MOVD ctx+0(FP), R0
|
||||
MOVD 16(R0), R16
|
||||
SUB R16, R2, R2
|
||||
MOVD R2, 40(R0)
|
||||
MOVD (R0), R0
|
||||
MOVD R8, 24(R0)
|
||||
MOVD R9, 32(R0)
|
||||
MOVB R10, 40(R0)
|
||||
RET
|
||||
|
||||
// Report error
|
||||
error_max_decoded_size_exceeded:
|
||||
MOVD ctx+0(FP), R0
|
||||
MOVD $-1, R1
|
||||
MOVD R1, 40(R0)
|
||||
RET
|
||||
|
||||
// skipped decompress1x_main_loop_bmi2 (generic twin preferred on arm64)
|
||||
+193
@@ -0,0 +1,193 @@
|
||||
//go:build (amd64 || arm64) && !appengine && !noasm && gc
|
||||
|
||||
// This file contains the specialisation of Decoder.Decompress4X
|
||||
// and Decoder.Decompress1X that use an asm implementation of their main loops.
|
||||
// The asm function stubs and any per-arch dispatch live in decompress_amd64.go
|
||||
// and decompress_arm64.go.
|
||||
package huff0
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// fallback8BitSize is the size where using Go version is faster.
|
||||
const fallback8BitSize = 800
|
||||
|
||||
type decompress4xContext struct {
|
||||
pbr *[4]bitReaderShifted
|
||||
peekBits uint8
|
||||
out *byte
|
||||
dstEvery int
|
||||
tbl *dEntrySingle
|
||||
decoded int
|
||||
limit *byte
|
||||
}
|
||||
|
||||
// Decompress4X will decompress a 4X encoded stream.
|
||||
// The length of the supplied input must match the end of a block exactly.
|
||||
// The *capacity* of the dst slice must match the destination size of
|
||||
// the uncompressed data exactly.
|
||||
func (d *Decoder) Decompress4X(dst, src []byte) ([]byte, error) {
|
||||
if len(d.dt.single) == 0 {
|
||||
return nil, errors.New("no table loaded")
|
||||
}
|
||||
if len(src) < 6+(4*1) {
|
||||
return nil, errors.New("input too small")
|
||||
}
|
||||
|
||||
use8BitTables := d.actualTableLog <= 8
|
||||
if cap(dst) < fallback8BitSize && use8BitTables {
|
||||
return d.decompress4X8bit(dst, src)
|
||||
}
|
||||
|
||||
var br [4]bitReaderShifted
|
||||
// Decode "jump table"
|
||||
start := 6
|
||||
for i := range 3 {
|
||||
length := int(src[i*2]) | (int(src[i*2+1]) << 8)
|
||||
if start+length >= len(src) {
|
||||
return nil, errors.New("truncated input (or invalid offset)")
|
||||
}
|
||||
err := br[i].init(src[start : start+length])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
start += length
|
||||
}
|
||||
err := br[3].init(src[start:])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// destination, offset to match first output
|
||||
dstSize := cap(dst)
|
||||
dst = dst[:dstSize]
|
||||
out := dst
|
||||
dstEvery := (dstSize + 3) / 4
|
||||
|
||||
const tlSize = 1 << tableLogMax
|
||||
const tlMask = tlSize - 1
|
||||
single := d.dt.single[:tlSize]
|
||||
|
||||
var decoded int
|
||||
|
||||
if len(out) > 4*4 && !(br[0].off < 4 || br[1].off < 4 || br[2].off < 4 || br[3].off < 4) {
|
||||
ctx := decompress4xContext{
|
||||
pbr: &br,
|
||||
peekBits: uint8((64 - d.actualTableLog) & 63), // see: bitReaderShifted.peekBitsFast()
|
||||
out: &out[0],
|
||||
dstEvery: dstEvery,
|
||||
tbl: &single[0],
|
||||
limit: &out[dstEvery-4], // Always stop decoding when first buffer gets here to avoid writing OOB on last.
|
||||
}
|
||||
if use8BitTables {
|
||||
decompress4x_8b_main_loop_asm(&ctx)
|
||||
} else {
|
||||
decompress4x_main_loop_asm(&ctx)
|
||||
}
|
||||
|
||||
decoded = ctx.decoded
|
||||
out = out[decoded/4:]
|
||||
}
|
||||
|
||||
// Decode remaining.
|
||||
remainBytes := dstEvery - (decoded / 4)
|
||||
for i := range br {
|
||||
offset := dstEvery * i
|
||||
endsAt := min(offset+remainBytes, len(out))
|
||||
br := &br[i]
|
||||
bitsLeft := br.remaining()
|
||||
for bitsLeft > 0 {
|
||||
br.fill()
|
||||
if offset >= endsAt {
|
||||
return nil, errors.New("corruption detected: stream overrun 4")
|
||||
}
|
||||
|
||||
// Read value and increment offset.
|
||||
val := br.peekBitsFast(d.actualTableLog)
|
||||
v := single[val&tlMask].entry
|
||||
nBits := uint8(v)
|
||||
br.advance(nBits)
|
||||
bitsLeft -= uint(nBits)
|
||||
out[offset] = uint8(v >> 8)
|
||||
offset++
|
||||
}
|
||||
if offset != endsAt {
|
||||
return nil, fmt.Errorf("corruption detected: short output block %d, end %d != %d", i, offset, endsAt)
|
||||
}
|
||||
decoded += offset - dstEvery*i
|
||||
err = br.close()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if dstSize != decoded {
|
||||
return nil, errors.New("corruption detected: short output block")
|
||||
}
|
||||
return dst, nil
|
||||
}
|
||||
|
||||
type decompress1xContext struct {
|
||||
pbr *bitReaderShifted
|
||||
peekBits uint8
|
||||
out *byte
|
||||
outCap int
|
||||
tbl *dEntrySingle
|
||||
decoded int
|
||||
}
|
||||
|
||||
// Error reported by asm implementations
|
||||
const error_max_decoded_size_exeeded = -1
|
||||
|
||||
// Decompress1X will decompress a 1X encoded stream.
|
||||
// The cap of the output buffer will be the maximum decompressed size.
|
||||
// The length of the supplied input must match the end of a block exactly.
|
||||
func (d *Decoder) Decompress1X(dst, src []byte) ([]byte, error) {
|
||||
if len(d.dt.single) == 0 {
|
||||
return nil, errors.New("no table loaded")
|
||||
}
|
||||
var br bitReaderShifted
|
||||
err := br.init(src)
|
||||
if err != nil {
|
||||
return dst, err
|
||||
}
|
||||
maxDecodedSize := cap(dst)
|
||||
dst = dst[:maxDecodedSize]
|
||||
|
||||
const tlSize = 1 << tableLogMax
|
||||
const tlMask = tlSize - 1
|
||||
|
||||
if maxDecodedSize >= 4 {
|
||||
ctx := decompress1xContext{
|
||||
pbr: &br,
|
||||
out: &dst[0],
|
||||
outCap: maxDecodedSize,
|
||||
peekBits: uint8((64 - d.actualTableLog) & 63), // see: bitReaderShifted.peekBitsFast()
|
||||
tbl: &d.dt.single[0],
|
||||
}
|
||||
|
||||
decompress1x_main_loop_asm(&ctx)
|
||||
if ctx.decoded == error_max_decoded_size_exeeded {
|
||||
return nil, ErrMaxDecodedSizeExceeded
|
||||
}
|
||||
|
||||
dst = dst[:ctx.decoded]
|
||||
}
|
||||
|
||||
// br < 8, so uint8 is fine
|
||||
bitsLeft := uint8(br.off)*8 + 64 - br.bitsRead
|
||||
for bitsLeft > 0 {
|
||||
br.fill()
|
||||
if len(dst) >= maxDecodedSize {
|
||||
br.close()
|
||||
return nil, ErrMaxDecodedSizeExceeded
|
||||
}
|
||||
v := d.dt.single[br.peekBitsFast(d.actualTableLog)&tlMask]
|
||||
nBits := uint8(v.entry)
|
||||
br.advance(nBits)
|
||||
bitsLeft -= nBits
|
||||
dst = append(dst, uint8(v.entry>>8))
|
||||
}
|
||||
return dst, br.close()
|
||||
}
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
//go:build !amd64 || appengine || !gc || noasm
|
||||
//go:build (!amd64 && !arm64) || appengine || !gc || noasm
|
||||
|
||||
// This file contains a generic implementation of Decoder.Decompress4X.
|
||||
package huff0
|
||||
|
||||
+1
-1
@@ -31,7 +31,7 @@ func DecodedLen(src []byte) (int, error) {
|
||||
// that the length header occupied.
|
||||
func decodedLen(src []byte) (blockLen, headerLen int, err error) {
|
||||
v, n := binary.Uvarint(src)
|
||||
if n <= 0 || v > 0xffffffff {
|
||||
if n <= 0 || n > 5 || v > 0xffffffff {
|
||||
return 0, 0, ErrCorrupt
|
||||
}
|
||||
|
||||
|
||||
+35
-2
@@ -75,14 +75,47 @@ The above is fine for big encodes. However, whenever possible try to *reuse* the
|
||||
To reuse the encoder, you can use the `Reset(io.Writer)` function to change to another output.
|
||||
This will allow the encoder to reuse all resources and avoid wasteful allocations.
|
||||
|
||||
Currently stream encoding has 'light' concurrency, meaning up to 2 goroutines can be working on part
|
||||
of a stream. This is independent of the `WithEncoderConcurrency(n)`, but that is likely to change
|
||||
By default, stream encoding has 'light' concurrency, meaning up to 2 goroutines can be working on part
|
||||
of a stream. This is independent of the `WithEncoderConcurrency(n)`, but that is likely to change
|
||||
in the future. So if you want to limit concurrency for future updates, specify the concurrency
|
||||
you would like.
|
||||
|
||||
If you would like stream encoding to be done without spawning async goroutines, use `WithEncoderConcurrency(1)`
|
||||
which will compress input as each block is completed, blocking on writes until each has completed.
|
||||
|
||||
#### Parallel Stream Compression
|
||||
|
||||
For maximum throughput on large streams, use `WithConcurrentBlocks(true)` together with
|
||||
`WithEncoderConcurrency(n)` where n is the number of CPU cores you want to use.
|
||||
This splits the input into large sections (jobs) that are compressed simultaneously by multiple goroutines,
|
||||
similar to how the C zstd library does multithreaded compression.
|
||||
|
||||
```Go
|
||||
enc, err := zstd.NewWriter(out,
|
||||
zstd.WithEncoderLevel(zstd.SpeedDefault),
|
||||
zstd.WithEncoderConcurrency(runtime.GOMAXPROCS(0)),
|
||||
zstd.WithConcurrentBlocks(true),
|
||||
)
|
||||
```
|
||||
|
||||
Each non-first job receives an overlap prefix from the previous job for match context,
|
||||
so compression ratio is only marginally affected. Output is flushed in order,
|
||||
producing a valid single-frame zstd stream.
|
||||
|
||||
Benchmark on 1.8GB GOB stream (AMD Ryzen 9 9950X):
|
||||
|
||||
| Level | 1 thread | 4 threads | 16 threads | 1T ratio | 16T ratio |
|
||||
|---------|:----------:|:------------------:|:-------------------:|:--------:|:---------:|
|
||||
| fastest | 783 MB/s | 2950 MB/s (3.8×) | 6939 MB/s (8.9×) | 12.24% | 12.26% |
|
||||
| default | 728 MB/s | 2533 MB/s (3.5×) | 5340 MB/s (7.3×) | 10.67% | 10.68% |
|
||||
| better | 434 MB/s | 1105 MB/s (2.5×) | 2206 MB/s (5.1×) | 9.14% | 9.21% |
|
||||
| best | 129 MB/s | 367 MB/s (2.8×) | 884 MB/s (6.8×) | 8.48% | 8.63% |
|
||||
|
||||
Notes:
|
||||
* Not compatible with dictionary encoding.
|
||||
* `Flush()` dispatches the current partial job, so latency-sensitive callers can force output.
|
||||
* `EncodeAll` is unaffected — it uses its own concurrency via the encoder pool.
|
||||
|
||||
You can specify your desired compression level using `WithEncoderLevel()` option. Currently only pre-defined
|
||||
compression settings can be specified.
|
||||
|
||||
|
||||
+3
-6
@@ -400,8 +400,9 @@ func (b *blockDec) decodeLiterals(in []byte, hist *history) (remain []byte, err
|
||||
}
|
||||
}
|
||||
var err error
|
||||
// Use our out buffer.
|
||||
huff.MaxDecodedSize = litRegenSize
|
||||
// Decoder.Decompress* uses cap(dst) for the size limit. Do not write
|
||||
// MaxDecodedSize on hist.huffTree: with a trained dictionary that
|
||||
// pointer aliases the shared dict.litEnc and concurrent DecodeAll races.
|
||||
if fourStreams {
|
||||
literals, err = huff.Decoder().Decompress4X(b.literalBuf[:0:litRegenSize], literals)
|
||||
} else {
|
||||
@@ -673,10 +674,6 @@ func (b *blockDec) executeSequences(hist *history) error {
|
||||
hbytes := hist.b
|
||||
if len(hbytes) > hist.windowSize {
|
||||
hbytes = hbytes[len(hbytes)-hist.windowSize:]
|
||||
// We do not need history anymore.
|
||||
if hist.dict != nil {
|
||||
hist.dict.content = nil
|
||||
}
|
||||
}
|
||||
hist.decoders.windowSize = hist.windowSize
|
||||
hist.decoders.out = b.dst[:0]
|
||||
|
||||
+66
-22
@@ -230,7 +230,7 @@ func BuildDict(o BuildDictOptions) ([]byte, error) {
|
||||
}
|
||||
block := blockEnc{lowMem: false}
|
||||
block.init()
|
||||
enc := encoder(&bestFastEncoder{fastBase: fastBase{maxMatchOff: int32(maxMatchLen), bufferReset: math.MaxInt32 - int32(maxMatchLen*2), lowMem: false}})
|
||||
var enc encoder
|
||||
if o.Level != 0 {
|
||||
eOpts := encoderOptions{
|
||||
level: o.Level,
|
||||
@@ -242,6 +242,7 @@ func BuildDict(o BuildDictOptions) ([]byte, error) {
|
||||
enc = eOpts.encoder()
|
||||
} else {
|
||||
o.Level = SpeedBestCompression
|
||||
enc = encoder(&bestFastEncoder{fastBase: fastBase{maxMatchOff: int32(maxMatchLen), bufferReset: math.MaxInt32 - int32(maxMatchLen*2), lowMem: false}})
|
||||
}
|
||||
var (
|
||||
remain [256]int
|
||||
@@ -295,40 +296,81 @@ func BuildDict(o BuildDictOptions) ([]byte, error) {
|
||||
if offset > 3 {
|
||||
newOffsets[offset-3]++
|
||||
} else {
|
||||
newOffsets[uint32(o.Offsets[offset-1])]++
|
||||
// Repeat codes reference the training Offsets. Skip unset
|
||||
// (zero) entries so they are not ranked as real offsets.
|
||||
prev := o.Offsets[offset-1]
|
||||
if prev > 0 {
|
||||
newOffsets[uint32(prev)]++
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Find most used offsets.
|
||||
var sortedOffsets []uint32
|
||||
for k := range newOffsets {
|
||||
if k == 0 {
|
||||
continue
|
||||
}
|
||||
sortedOffsets = append(sortedOffsets, k)
|
||||
}
|
||||
sort.Slice(sortedOffsets, func(i, j int) bool {
|
||||
a, b := sortedOffsets[i], sortedOffsets[j]
|
||||
if a == b {
|
||||
ca, cb := newOffsets[a], newOffsets[b]
|
||||
if ca == cb {
|
||||
// Prefer the longer offset
|
||||
return sortedOffsets[i] > sortedOffsets[j]
|
||||
return a > b
|
||||
}
|
||||
return newOffsets[sortedOffsets[i]] > newOffsets[sortedOffsets[j]]
|
||||
return ca > cb
|
||||
})
|
||||
if len(sortedOffsets) > 3 {
|
||||
if debug {
|
||||
print("Offsets:")
|
||||
for i, v := range sortedOffsets {
|
||||
if i > 20 {
|
||||
break
|
||||
}
|
||||
printf("[%d: %d],", v, newOffsets[v])
|
||||
if debug {
|
||||
print("Offsets:")
|
||||
for i, v := range sortedOffsets {
|
||||
if i > 20 {
|
||||
break
|
||||
}
|
||||
println("")
|
||||
printf("[%d: %d],", v, newOffsets[v])
|
||||
}
|
||||
|
||||
sortedOffsets = sortedOffsets[:3]
|
||||
println("")
|
||||
}
|
||||
for i, v := range sortedOffsets {
|
||||
o.Offsets[i] = int(v)
|
||||
// Dictionary recent-offsets must be three positive values within the
|
||||
// history. Ranked matches may be fewer (or empty when only unset
|
||||
// repeat codes were seen), so fill remaining slots with defaults.
|
||||
used := make(map[int]bool, 3)
|
||||
var finalOffsets [3]int
|
||||
nOff := 0
|
||||
for _, v := range sortedOffsets {
|
||||
iv := int(v)
|
||||
if iv <= 0 || iv > len(hist) || used[iv] {
|
||||
continue
|
||||
}
|
||||
finalOffsets[nOff] = iv
|
||||
used[iv] = true
|
||||
nOff++
|
||||
if nOff == 3 {
|
||||
break
|
||||
}
|
||||
}
|
||||
for _, def := range []int{1, 4, 8} {
|
||||
if nOff == 3 {
|
||||
break
|
||||
}
|
||||
if def <= len(hist) && !used[def] {
|
||||
finalOffsets[nOff] = def
|
||||
used[def] = true
|
||||
nOff++
|
||||
}
|
||||
}
|
||||
for def := 1; nOff < 3 && def <= len(hist); def++ {
|
||||
if !used[def] {
|
||||
finalOffsets[nOff] = def
|
||||
used[def] = true
|
||||
nOff++
|
||||
}
|
||||
}
|
||||
if nOff < 3 {
|
||||
return nil, fmt.Errorf("could not determine 3 valid dictionary offsets (history size %d)", len(hist))
|
||||
}
|
||||
o.Offsets = finalOffsets
|
||||
if debug {
|
||||
println("New repeat offsets", o.Offsets)
|
||||
}
|
||||
@@ -336,6 +378,9 @@ func BuildDict(o BuildDictOptions) ([]byte, error) {
|
||||
if nUsed == 0 || seqs == 0 {
|
||||
return nil, fmt.Errorf("%d blocks, %d sequences found", nUsed, seqs)
|
||||
}
|
||||
if litTotal == 0 {
|
||||
return nil, errors.New("0 literals found")
|
||||
}
|
||||
if debug {
|
||||
println("Sequences:", seqs, "Blocks:", nUsed, "Literals:", litTotal)
|
||||
}
|
||||
@@ -516,11 +561,10 @@ func BuildDict(o BuildDictOptions) ([]byte, error) {
|
||||
out.Write(binary.LittleEndian.AppendUint32(nil, uint32(o.Offsets[1])))
|
||||
out.Write(binary.LittleEndian.AppendUint32(nil, uint32(o.Offsets[2])))
|
||||
out.Write(hist)
|
||||
if _, err := loadDict(out.Bytes()); err != nil {
|
||||
return nil, fmt.Errorf("built dictionary failed validation: %w", err)
|
||||
}
|
||||
if debug {
|
||||
_, err := loadDict(out.Bytes())
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
i, err := InspectDictionary(out.Bytes())
|
||||
if err != nil {
|
||||
panic(err)
|
||||
|
||||
+28
@@ -128,6 +128,34 @@ func (e *fastBase) matchlen(s, t int32, src []byte) int32 {
|
||||
return int32(matchLen(src[s:], src[t:]))
|
||||
}
|
||||
|
||||
// resetBasePrefix resets the encoder state and loads prefix as initial history.
|
||||
// This is used for parallel job encoding where non-first jobs need overlap context.
|
||||
// Rep offsets are set to defaults [1,4,8] (invalidated, matching C behavior).
|
||||
func (e *fastBase) resetBasePrefix(prefix []byte) {
|
||||
if e.blk == nil {
|
||||
e.blk = &blockEnc{lowMem: e.lowMem}
|
||||
e.blk.init()
|
||||
} else {
|
||||
e.blk.reset(nil)
|
||||
}
|
||||
e.blk.initNewEncode()
|
||||
if e.crc == nil {
|
||||
e.crc = xxhash.New()
|
||||
} else {
|
||||
e.crc.Reset()
|
||||
}
|
||||
e.blk.dictLitEnc = nil
|
||||
e.ensureHist(len(prefix) + maxCompressedBlockSize)
|
||||
// Bump cur so old table entries fall outside the window.
|
||||
// When cur >= bufferReset, leave it; the first Encode call
|
||||
// will shift/clear tables, preserving valid prefix entries.
|
||||
if e.cur < e.bufferReset {
|
||||
e.cur += e.maxMatchOff + int32(len(e.hist))
|
||||
}
|
||||
e.hist = e.hist[:0]
|
||||
e.hist = append(e.hist, prefix...)
|
||||
}
|
||||
|
||||
// Reset the encoding table.
|
||||
func (e *fastBase) resetBase(d *dict, singleBlock bool) {
|
||||
if e.blk == nil {
|
||||
|
||||
+15
@@ -551,3 +551,18 @@ func (e *bestFastEncoder) Reset(d *dict, singleBlock bool) {
|
||||
// Reset table to initial state
|
||||
copy(e.table[:], e.dictTable)
|
||||
}
|
||||
|
||||
func (e *bestFastEncoder) ResetPrefix(prefix []byte) {
|
||||
e.resetBasePrefix(prefix)
|
||||
if len(prefix) < 8 {
|
||||
return
|
||||
}
|
||||
end := e.cur + int32(len(prefix)) - 8
|
||||
for i := e.cur; i < end; i++ {
|
||||
cv := load6432(prefix, i-e.cur)
|
||||
h := hashLen(cv, bestLongTableBits, bestLongLen)
|
||||
e.longTable[h] = prevEntry{offset: i, prev: e.longTable[h].offset}
|
||||
h0 := hashLen(cv, bestShortTableBits, bestShortLen)
|
||||
e.table[h0] = prevEntry{offset: i, prev: e.table[h0].offset}
|
||||
}
|
||||
}
|
||||
|
||||
+18
@@ -1096,6 +1096,20 @@ func (e *betterFastEncoder) Reset(d *dict, singleBlock bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func (e *betterFastEncoder) ResetPrefix(prefix []byte) {
|
||||
e.resetBasePrefix(prefix)
|
||||
if len(prefix) < 8 {
|
||||
return
|
||||
}
|
||||
end := e.cur + int32(len(prefix)) - 8
|
||||
for i := e.cur; i < end; i += 2 {
|
||||
cv := load6432(prefix, i-e.cur)
|
||||
h := hashLen(cv, betterLongTableBits, betterLongLen)
|
||||
e.longTable[h] = prevEntry{offset: i, prev: e.longTable[h].offset}
|
||||
e.table[hashLen(cv>>8, betterShortTableBits, betterShortLen)] = tableEntry{val: uint32(cv >> 8), offset: i + 1}
|
||||
}
|
||||
}
|
||||
|
||||
// ResetDict will reset and set a dictionary if not nil
|
||||
func (e *betterFastEncoderDict) Reset(d *dict, singleBlock bool) {
|
||||
e.resetBase(d, singleBlock)
|
||||
@@ -1229,6 +1243,10 @@ func (e *betterFastEncoderDict) Reset(d *dict, singleBlock bool) {
|
||||
e.allDirty = false
|
||||
}
|
||||
|
||||
func (e *betterFastEncoderDict) ResetPrefix([]byte) {
|
||||
panic("ResetPrefix not supported for dict encoders")
|
||||
}
|
||||
|
||||
func (e *betterFastEncoderDict) markLongShardDirty(entryNum uint32) {
|
||||
e.longTableShardDirty[entryNum/betterLongTableShardSize] = true
|
||||
}
|
||||
|
||||
+16
@@ -1037,6 +1037,18 @@ func (e *doubleFastEncoder) Reset(d *dict, singleBlock bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func (e *doubleFastEncoder) ResetPrefix(prefix []byte) {
|
||||
e.fastEncoder.ResetPrefix(prefix)
|
||||
if len(prefix) < 8 {
|
||||
return
|
||||
}
|
||||
end := e.cur + int32(len(prefix)) - 8
|
||||
for i := e.cur + 1; i < end; i += 2 {
|
||||
cv := load6432(prefix, i-e.cur)
|
||||
e.longTable[hashLen(cv, dFastLongTableBits, dFastLongLen)] = tableEntry{val: uint32(cv), offset: i}
|
||||
}
|
||||
}
|
||||
|
||||
// ResetDict will reset and set a dictionary if not nil
|
||||
func (e *doubleFastEncoderDict) Reset(d *dict, singleBlock bool) {
|
||||
allDirty := e.allDirty
|
||||
@@ -1102,6 +1114,10 @@ func (e *doubleFastEncoderDict) Reset(d *dict, singleBlock bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func (e *doubleFastEncoderDict) ResetPrefix([]byte) {
|
||||
panic("ResetPrefix not supported for dict encoders")
|
||||
}
|
||||
|
||||
func (e *doubleFastEncoderDict) markLongShardDirty(entryNum uint32) {
|
||||
e.longTableShardDirty[entryNum/dLongTableShardSize] = true
|
||||
}
|
||||
|
||||
+17
@@ -797,6 +797,19 @@ func (e *fastEncoder) Reset(d *dict, singleBlock bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func (e *fastEncoder) ResetPrefix(prefix []byte) {
|
||||
e.resetBasePrefix(prefix)
|
||||
if len(prefix) < 8 {
|
||||
return
|
||||
}
|
||||
end := e.cur + int32(len(prefix)) - 8
|
||||
// Index every 4th
|
||||
for i := e.cur + 1; i < end; i += 4 {
|
||||
cv := load6432(prefix, i-e.cur)
|
||||
e.table[hashLen(cv, tableBits, tableFastHashLen)] = tableEntry{val: uint32(cv), offset: i}
|
||||
}
|
||||
}
|
||||
|
||||
// ResetDict will reset and set a dictionary if not nil
|
||||
func (e *fastEncoderDict) Reset(d *dict, singleBlock bool) {
|
||||
e.resetBase(d, singleBlock)
|
||||
@@ -866,6 +879,10 @@ func (e *fastEncoderDict) Reset(d *dict, singleBlock bool) {
|
||||
e.allDirty = false
|
||||
}
|
||||
|
||||
func (e *fastEncoderDict) ResetPrefix([]byte) {
|
||||
panic("ResetPrefix not supported for dict encoders")
|
||||
}
|
||||
|
||||
func (e *fastEncoderDict) markAllShardsDirty() {
|
||||
e.allDirty = true
|
||||
}
|
||||
|
||||
+352
@@ -0,0 +1,352 @@
|
||||
// Copyright 2019+ Klaus Post. All rights reserved.
|
||||
// License information can be found in the LICENSE file.
|
||||
// Based on work by Yann Collet, released under BSD License.
|
||||
|
||||
package zstd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
rdebug "runtime/debug"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type encJob struct {
|
||||
prefix []byte // overlap from previous job (nil for first)
|
||||
input []byte // job's own input data (swapped from filling)
|
||||
last bool // last block of last job gets last=true
|
||||
output []byte // compressed blocks (filled by worker)
|
||||
err error // encoding error
|
||||
done chan struct{} // closed when complete
|
||||
}
|
||||
|
||||
type jobState struct {
|
||||
jobSize int
|
||||
overlapSize int
|
||||
filling []byte // accumulates input up to jobSize
|
||||
nextPrefix []byte // overlap prefix prepared for the next dispatched job
|
||||
|
||||
jobSeq int // next job sequence number
|
||||
|
||||
jobCh chan *encJob // dispatch to workers
|
||||
resultCh chan *encJob // ordered results to flusher
|
||||
|
||||
workerWg sync.WaitGroup
|
||||
flusherWg sync.WaitGroup
|
||||
|
||||
mu sync.Mutex
|
||||
flushedSeq int // last flushed sequence number
|
||||
cond *sync.Cond
|
||||
|
||||
flusherErr error
|
||||
started bool
|
||||
|
||||
inputPool sync.Pool // *[]byte buffers of jobSize cap
|
||||
outputPool sync.Pool // *[]byte buffers for compressed output
|
||||
overlapPool sync.Pool // *[]byte buffers for overlap prefixes
|
||||
}
|
||||
|
||||
func (e *Encoder) startJobWorkers() {
|
||||
js := &e.state.jobs
|
||||
n := e.o.concurrent
|
||||
js.jobCh = make(chan *encJob, n)
|
||||
js.resultCh = make(chan *encJob, n)
|
||||
js.flushedSeq = 0
|
||||
js.cond = sync.NewCond(&js.mu)
|
||||
|
||||
// Workers borrow encoders from the shared e.encoders pool per-job.
|
||||
// Ensure the pool is initialized before any worker tries to borrow.
|
||||
e.init.Do(e.initialize)
|
||||
|
||||
for range n {
|
||||
js.workerWg.Add(1)
|
||||
go e.jobWorker()
|
||||
}
|
||||
js.flusherWg.Add(1)
|
||||
go e.jobFlusher()
|
||||
js.started = true
|
||||
}
|
||||
|
||||
func (e *Encoder) jobWorker() {
|
||||
js := &e.state.jobs
|
||||
defer js.workerWg.Done()
|
||||
for job := range js.jobCh {
|
||||
enc := <-e.encoders
|
||||
e.compressJob(enc, job)
|
||||
e.encoders <- enc
|
||||
close(job.done)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Encoder) compressJob(enc encoder, job *encJob) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
job.err = fmt.Errorf("panic in parallel job: %v", r)
|
||||
rdebug.PrintStack()
|
||||
}
|
||||
}()
|
||||
|
||||
if len(job.prefix) > 0 {
|
||||
enc.ResetPrefix(job.prefix)
|
||||
} else {
|
||||
enc.Reset(nil, false)
|
||||
}
|
||||
|
||||
data := job.input
|
||||
if len(data) == 0 && job.last {
|
||||
blk := enc.Block()
|
||||
blk.reset(nil)
|
||||
blk.last = true
|
||||
blk.encodeRaw(nil)
|
||||
job.output = append(job.output, blk.output...)
|
||||
return
|
||||
}
|
||||
|
||||
blk := enc.Block()
|
||||
for len(data) > 0 {
|
||||
todo := data
|
||||
if len(todo) > e.o.blockSize {
|
||||
todo = todo[:e.o.blockSize]
|
||||
}
|
||||
data = data[len(todo):]
|
||||
|
||||
blk.pushOffsets()
|
||||
enc.Encode(blk, todo)
|
||||
blk.last = len(data) == 0 && job.last
|
||||
|
||||
err := blk.encode(todo, e.o.noEntropy, !e.o.allLitEntropy)
|
||||
if err != nil {
|
||||
job.err = err
|
||||
return
|
||||
}
|
||||
job.output = append(job.output, blk.output...)
|
||||
blk.reset(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func (js *jobState) getInputBuf(size int) []byte {
|
||||
if v := js.inputPool.Get(); v != nil {
|
||||
bp := v.(*[]byte)
|
||||
b := *bp
|
||||
if cap(b) >= size {
|
||||
return b[:0]
|
||||
}
|
||||
}
|
||||
return make([]byte, 0, size)
|
||||
}
|
||||
|
||||
func (js *jobState) putInputBuf(b []byte) {
|
||||
if cap(b) > 0 {
|
||||
b = b[:0]
|
||||
js.inputPool.Put(&b)
|
||||
}
|
||||
}
|
||||
|
||||
func (js *jobState) getOutputBuf(size int) []byte {
|
||||
if v := js.outputPool.Get(); v != nil {
|
||||
bp := v.(*[]byte)
|
||||
b := *bp
|
||||
if cap(b) >= size {
|
||||
return b[:0]
|
||||
}
|
||||
}
|
||||
return make([]byte, 0, size)
|
||||
}
|
||||
|
||||
func (js *jobState) putOutputBuf(b []byte) {
|
||||
if cap(b) > 0 {
|
||||
b = b[:0]
|
||||
js.outputPool.Put(&b)
|
||||
}
|
||||
}
|
||||
|
||||
func (js *jobState) getOverlapBuf(size int) []byte {
|
||||
if v := js.overlapPool.Get(); v != nil {
|
||||
bp := v.(*[]byte)
|
||||
b := *bp
|
||||
if cap(b) >= size {
|
||||
return b[:size]
|
||||
}
|
||||
}
|
||||
return make([]byte, size)
|
||||
}
|
||||
|
||||
func (js *jobState) putOverlapBuf(b []byte) {
|
||||
if cap(b) > 0 {
|
||||
b = b[:0]
|
||||
js.overlapPool.Put(&b)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Encoder) jobFlusher() {
|
||||
js := &e.state.jobs
|
||||
defer js.flusherWg.Done()
|
||||
for job := range js.resultCh {
|
||||
<-job.done
|
||||
// Worker has fully exited compressJob, so the prefix is no longer
|
||||
// in use. Return it to the pool regardless of outcome.
|
||||
if job.prefix != nil {
|
||||
js.putOverlapBuf(job.prefix)
|
||||
job.prefix = nil
|
||||
}
|
||||
if job.err != nil {
|
||||
js.mu.Lock()
|
||||
js.flusherErr = job.err
|
||||
js.cond.Broadcast()
|
||||
js.mu.Unlock()
|
||||
for range js.resultCh {
|
||||
}
|
||||
return
|
||||
}
|
||||
if len(job.output) > 0 {
|
||||
_, err := e.state.w.Write(job.output)
|
||||
if err != nil {
|
||||
js.mu.Lock()
|
||||
js.flusherErr = err
|
||||
js.cond.Broadcast()
|
||||
js.mu.Unlock()
|
||||
for range js.resultCh {
|
||||
}
|
||||
return
|
||||
}
|
||||
e.state.nWritten += int64(len(job.output))
|
||||
}
|
||||
// Return buffers to pools.
|
||||
js.putInputBuf(job.input)
|
||||
js.putOutputBuf(job.output)
|
||||
job.input = nil
|
||||
job.output = nil
|
||||
|
||||
js.mu.Lock()
|
||||
js.flushedSeq++
|
||||
js.cond.Broadcast()
|
||||
js.mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Encoder) shutdownJobWorkers() {
|
||||
js := &e.state.jobs
|
||||
if !js.started {
|
||||
return
|
||||
}
|
||||
close(js.jobCh)
|
||||
js.workerWg.Wait()
|
||||
close(js.resultCh)
|
||||
js.flusherWg.Wait()
|
||||
js.started = false
|
||||
}
|
||||
|
||||
// waitAllJobs blocks until all dispatched jobs have been flushed.
|
||||
func (e *Encoder) waitAllJobs() {
|
||||
js := &e.state.jobs
|
||||
if !js.started {
|
||||
return
|
||||
}
|
||||
js.mu.Lock()
|
||||
for js.flushedSeq < js.jobSeq && js.flusherErr == nil {
|
||||
js.cond.Wait()
|
||||
}
|
||||
js.mu.Unlock()
|
||||
}
|
||||
|
||||
func (e *Encoder) dispatchJob(final bool) error {
|
||||
s := &e.state
|
||||
js := &s.jobs
|
||||
|
||||
js.mu.Lock()
|
||||
fErr := js.flusherErr
|
||||
js.mu.Unlock()
|
||||
if fErr != nil {
|
||||
return fErr
|
||||
}
|
||||
|
||||
if !s.headerWritten {
|
||||
// Single-block optimization: fall through to encodeAll path.
|
||||
if final && len(js.filling) > 0 && len(js.filling) <= e.o.blockSize {
|
||||
s.current = e.encodeAll(s.encoder, js.filling, s.current[:0])
|
||||
var n2 int
|
||||
n2, s.err = s.w.Write(s.current)
|
||||
if s.err != nil {
|
||||
return s.err
|
||||
}
|
||||
s.nWritten += int64(n2)
|
||||
s.nInput += int64(len(js.filling))
|
||||
s.current = s.current[:0]
|
||||
js.filling = js.filling[:0]
|
||||
s.headerWritten = true
|
||||
s.fullFrameWritten = true
|
||||
s.eofWritten = true
|
||||
return nil
|
||||
}
|
||||
if final && len(js.filling) == 0 && !e.o.fullZero {
|
||||
s.headerWritten = true
|
||||
s.fullFrameWritten = true
|
||||
s.eofWritten = true
|
||||
return nil
|
||||
}
|
||||
|
||||
var tmp [maxHeaderSize]byte
|
||||
fh := frameHeader{
|
||||
ContentSize: uint64(s.frameContentSize),
|
||||
WindowSize: uint32(s.encoder.WindowSize(s.frameContentSize)),
|
||||
SingleSegment: false,
|
||||
Checksum: e.o.crc,
|
||||
DictID: 0,
|
||||
}
|
||||
dst := fh.appendTo(tmp[:0])
|
||||
var n2 int
|
||||
n2, s.err = s.w.Write(dst)
|
||||
if s.err != nil {
|
||||
return s.err
|
||||
}
|
||||
s.nWritten += int64(n2)
|
||||
s.headerWritten = true
|
||||
}
|
||||
|
||||
if len(js.filling) == 0 && !final {
|
||||
return nil
|
||||
}
|
||||
|
||||
if !js.started {
|
||||
e.startJobWorkers()
|
||||
}
|
||||
|
||||
// Estimate output size for pooled buffer.
|
||||
outputEst := max(len(js.filling)/2, 512)
|
||||
|
||||
job := &encJob{
|
||||
last: final,
|
||||
done: make(chan struct{}),
|
||||
output: js.getOutputBuf(outputEst),
|
||||
}
|
||||
|
||||
// Each job owns its prefix slice; the flusher returns it to the pool
|
||||
// after <-job.done, so workers and dispatch never share a buffer.
|
||||
if js.nextPrefix != nil {
|
||||
job.prefix = js.nextPrefix
|
||||
js.nextPrefix = nil
|
||||
}
|
||||
|
||||
// Build the next job's prefix from the tail of this job's input.
|
||||
if !final && len(js.filling) > 0 {
|
||||
overlapLen := min(js.overlapSize, len(js.filling))
|
||||
np := js.getOverlapBuf(overlapLen)
|
||||
copy(np, js.filling[len(js.filling)-overlapLen:])
|
||||
js.nextPrefix = np
|
||||
}
|
||||
|
||||
// Swap filling buffer into job — zero-copy for the input data.
|
||||
job.input = js.filling
|
||||
js.filling = js.getInputBuf(js.jobSize)
|
||||
|
||||
s.nInput += int64(len(job.input))
|
||||
js.jobSeq++
|
||||
|
||||
if final {
|
||||
s.eofWritten = true
|
||||
}
|
||||
|
||||
js.resultCh <- job
|
||||
js.jobCh <- job
|
||||
|
||||
return nil
|
||||
}
|
||||
+206
-4
@@ -38,6 +38,7 @@ type encoder interface {
|
||||
WindowSize(size int64) int32
|
||||
UseBlock(*blockEnc)
|
||||
Reset(d *dict, singleBlock bool)
|
||||
ResetPrefix(prefix []byte)
|
||||
}
|
||||
|
||||
type encoderState struct {
|
||||
@@ -60,6 +61,9 @@ type encoderState struct {
|
||||
wg sync.WaitGroup
|
||||
// This waitgroup indicates we have a block encoding/writing.
|
||||
wWg sync.WaitGroup
|
||||
|
||||
// Parallel job state (used when concurrentBlocks is enabled).
|
||||
jobs jobState
|
||||
}
|
||||
|
||||
// NewWriter will create a new Zstandard encoder.
|
||||
@@ -74,6 +78,9 @@ func NewWriter(w io.Writer, opts ...EOption) (*Encoder, error) {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if e.o.concurrentBlocks && (e.o.dict != nil || e.o.concurrent <= 1) {
|
||||
e.o.concurrentBlocks = false
|
||||
}
|
||||
if w != nil {
|
||||
e.Reset(w)
|
||||
}
|
||||
@@ -95,12 +102,31 @@ func (e *Encoder) initialize() {
|
||||
// as a new, independent stream.
|
||||
func (e *Encoder) Reset(w io.Writer) {
|
||||
s := &e.state
|
||||
|
||||
if e.o.concurrentBlocks {
|
||||
e.shutdownJobWorkers()
|
||||
js := &s.jobs
|
||||
js.jobSize = e.o.jobSize()
|
||||
js.overlapSize = e.o.overlapSize()
|
||||
// js.filling is allocated lazily on first Write/ReadFrom so callers
|
||||
// that only use EncodeAll don't pay the (up to ~32 MB) jobSize cost.
|
||||
js.filling = js.filling[:0]
|
||||
if js.nextPrefix != nil {
|
||||
js.putOverlapBuf(js.nextPrefix)
|
||||
js.nextPrefix = nil
|
||||
}
|
||||
js.jobSeq = 0
|
||||
js.flushedSeq = 0
|
||||
js.flusherErr = nil
|
||||
js.started = false
|
||||
}
|
||||
|
||||
s.wg.Wait()
|
||||
s.wWg.Wait()
|
||||
if cap(s.filling) == 0 {
|
||||
s.filling = make([]byte, 0, e.o.blockSize)
|
||||
}
|
||||
if e.o.concurrent > 1 {
|
||||
if e.o.concurrent > 1 && !e.o.concurrentBlocks {
|
||||
if cap(s.current) == 0 {
|
||||
s.current = make([]byte, 0, e.o.blockSize)
|
||||
}
|
||||
@@ -145,6 +171,9 @@ func (e *Encoder) ResetWithOptions(w io.Writer, opts ...EOption) error {
|
||||
}
|
||||
}
|
||||
hasDict := e.o.dict != nil
|
||||
if e.o.concurrentBlocks && hasDict {
|
||||
e.o.concurrentBlocks = false
|
||||
}
|
||||
if hadDict != hasDict {
|
||||
// Dict presence changed — encoder type must be recreated.
|
||||
e.state.encoder = nil
|
||||
@@ -176,6 +205,49 @@ func (e *Encoder) Write(p []byte) (n int, err error) {
|
||||
if s.eofWritten {
|
||||
return 0, ErrEncoderClosed
|
||||
}
|
||||
if e.o.concurrentBlocks {
|
||||
return e.writeJobs(p)
|
||||
}
|
||||
return e.writeBlocks(p)
|
||||
}
|
||||
|
||||
func (e *Encoder) writeJobs(p []byte) (n int, err error) {
|
||||
s := &e.state
|
||||
js := &s.jobs
|
||||
jobSize := js.jobSize
|
||||
if cap(js.filling) == 0 && len(p) > 0 {
|
||||
js.filling = make([]byte, 0, jobSize)
|
||||
}
|
||||
for len(p) > 0 {
|
||||
if len(p)+len(js.filling) < jobSize {
|
||||
if e.o.crc {
|
||||
_, _ = s.encoder.CRC().Write(p)
|
||||
}
|
||||
js.filling = append(js.filling, p...)
|
||||
return n + len(p), nil
|
||||
}
|
||||
add := p
|
||||
if len(p)+len(js.filling) > jobSize {
|
||||
add = add[:jobSize-len(js.filling)]
|
||||
}
|
||||
if e.o.crc {
|
||||
_, _ = s.encoder.CRC().Write(add)
|
||||
}
|
||||
js.filling = append(js.filling, add...)
|
||||
p = p[len(add):]
|
||||
n += len(add)
|
||||
if len(js.filling) < jobSize {
|
||||
return n, nil
|
||||
}
|
||||
if err := e.dispatchJob(false); err != nil {
|
||||
return n, err
|
||||
}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (e *Encoder) writeBlocks(p []byte) (n int, err error) {
|
||||
s := &e.state
|
||||
for len(p) > 0 {
|
||||
if len(p)+len(s.filling) < e.o.blockSize {
|
||||
if e.o.crc {
|
||||
@@ -374,6 +446,10 @@ func (e *Encoder) ReadFrom(r io.Reader) (n int64, err error) {
|
||||
println("Using ReadFrom")
|
||||
}
|
||||
|
||||
if e.o.concurrentBlocks {
|
||||
return e.readFromJobs(r)
|
||||
}
|
||||
|
||||
// Flush any current writes.
|
||||
if len(e.state.filling) > 0 {
|
||||
if err := e.nextBlock(false); err != nil {
|
||||
@@ -387,7 +463,6 @@ func (e *Encoder) ReadFrom(r io.Reader) (n int64, err error) {
|
||||
if e.o.crc {
|
||||
_, _ = e.state.encoder.CRC().Write(src[:n2])
|
||||
}
|
||||
// src is now the unfilled part...
|
||||
src = src[n2:]
|
||||
n += int64(n2)
|
||||
switch err {
|
||||
@@ -420,15 +495,63 @@ func (e *Encoder) ReadFrom(r io.Reader) (n int64, err error) {
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Encoder) readFromJobs(r io.Reader) (n int64, err error) {
|
||||
js := &e.state.jobs
|
||||
jobSize := js.jobSize
|
||||
|
||||
// Flush any current filling.
|
||||
if len(js.filling) > 0 {
|
||||
if err := e.dispatchJob(false); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
|
||||
if cap(js.filling) < jobSize {
|
||||
js.filling = make([]byte, 0, jobSize)
|
||||
}
|
||||
js.filling = js.filling[:jobSize]
|
||||
src := js.filling
|
||||
for {
|
||||
n2, err := r.Read(src)
|
||||
if e.o.crc {
|
||||
_, _ = e.state.encoder.CRC().Write(src[:n2])
|
||||
}
|
||||
src = src[n2:]
|
||||
n += int64(n2)
|
||||
switch err {
|
||||
case io.EOF:
|
||||
js.filling = js.filling[:len(js.filling)-len(src)]
|
||||
return n, nil
|
||||
case nil:
|
||||
default:
|
||||
e.state.err = err
|
||||
return n, err
|
||||
}
|
||||
if len(src) > 0 {
|
||||
continue
|
||||
}
|
||||
if err = e.dispatchJob(false); err != nil {
|
||||
return n, err
|
||||
}
|
||||
if cap(js.filling) < jobSize {
|
||||
js.filling = make([]byte, 0, jobSize)
|
||||
}
|
||||
js.filling = js.filling[:jobSize]
|
||||
src = js.filling
|
||||
}
|
||||
}
|
||||
|
||||
// Flush will send the currently written data to output
|
||||
// and block until everything has been written.
|
||||
// This should only be used on rare occasions where pushing the currently queued data is critical.
|
||||
func (e *Encoder) Flush() error {
|
||||
s := &e.state
|
||||
if e.o.concurrentBlocks {
|
||||
return e.flushJobs()
|
||||
}
|
||||
if len(s.filling) > 0 {
|
||||
err := e.nextBlock(false)
|
||||
if err != nil {
|
||||
// Ignore Flush after Close.
|
||||
if errors.Is(s.err, ErrEncoderClosed) {
|
||||
return nil
|
||||
}
|
||||
@@ -438,7 +561,6 @@ func (e *Encoder) Flush() error {
|
||||
s.wg.Wait()
|
||||
s.wWg.Wait()
|
||||
if s.err != nil {
|
||||
// Ignore Flush after Close.
|
||||
if errors.Is(s.err, ErrEncoderClosed) {
|
||||
return nil
|
||||
}
|
||||
@@ -447,6 +569,20 @@ func (e *Encoder) Flush() error {
|
||||
return s.writeErr
|
||||
}
|
||||
|
||||
func (e *Encoder) flushJobs() error {
|
||||
js := &e.state.jobs
|
||||
if len(js.filling) > 0 {
|
||||
if err := e.dispatchJob(false); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
e.waitAllJobs()
|
||||
js.mu.Lock()
|
||||
fErr := js.flusherErr
|
||||
js.mu.Unlock()
|
||||
return fErr
|
||||
}
|
||||
|
||||
// Close will flush the final output and close the stream.
|
||||
// The function will block until everything has been written.
|
||||
// The Encoder can still be re-used after calling this.
|
||||
@@ -455,12 +591,16 @@ func (e *Encoder) Close() error {
|
||||
if s.encoder == nil {
|
||||
return nil
|
||||
}
|
||||
if e.o.concurrentBlocks {
|
||||
return e.closeJobs()
|
||||
}
|
||||
if s.w == nil {
|
||||
if len(s.filling) == 0 && !s.headerWritten && !s.eofWritten && s.nInput == 0 {
|
||||
return nil
|
||||
}
|
||||
return errors.New("zstd: encoder has no writer")
|
||||
}
|
||||
|
||||
err := e.nextBlock(true)
|
||||
if err != nil {
|
||||
if errors.Is(s.err, ErrEncoderClosed) {
|
||||
@@ -511,6 +651,68 @@ func (e *Encoder) Close() error {
|
||||
return s.err
|
||||
}
|
||||
|
||||
func (e *Encoder) closeJobs() error {
|
||||
s := &e.state
|
||||
js := &s.jobs
|
||||
|
||||
if errors.Is(s.err, ErrEncoderClosed) {
|
||||
return nil
|
||||
}
|
||||
|
||||
if s.w == nil {
|
||||
if len(js.filling) == 0 && !s.headerWritten && !s.eofWritten && s.nInput == 0 {
|
||||
return nil
|
||||
}
|
||||
return errors.New("zstd: encoder has no writer")
|
||||
}
|
||||
|
||||
if err := e.dispatchJob(true); err != nil {
|
||||
e.shutdownJobWorkers()
|
||||
if errors.Is(s.err, ErrEncoderClosed) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
if s.frameContentSize > 0 && s.nInput != s.frameContentSize {
|
||||
e.shutdownJobWorkers()
|
||||
return fmt.Errorf("frame content size %d given, but %d bytes was written", s.frameContentSize, s.nInput)
|
||||
}
|
||||
|
||||
if s.fullFrameWritten {
|
||||
e.shutdownJobWorkers()
|
||||
s.err = ErrEncoderClosed
|
||||
return nil
|
||||
}
|
||||
|
||||
e.shutdownJobWorkers()
|
||||
if js.flusherErr != nil {
|
||||
return js.flusherErr
|
||||
}
|
||||
|
||||
// Write CRC
|
||||
if e.o.crc {
|
||||
var tmp [4]byte
|
||||
_, s.err = s.w.Write(s.encoder.AppendCRC(tmp[:0]))
|
||||
s.nWritten += 4
|
||||
}
|
||||
|
||||
// Add padding
|
||||
if s.err == nil && e.o.pad > 0 {
|
||||
add := calcSkippableFrame(s.nWritten, int64(e.o.pad))
|
||||
frame, err := skippableFrame(js.filling[:0], add, rand.Reader)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, s.err = s.w.Write(frame)
|
||||
}
|
||||
if s.err == nil {
|
||||
s.err = ErrEncoderClosed
|
||||
return nil
|
||||
}
|
||||
return s.err
|
||||
}
|
||||
|
||||
// EncodeAll will encode all input in src and append it to dst.
|
||||
// This function can be called concurrently, but each call will only run on a single goroutine.
|
||||
// If empty input is given, nothing is returned, unless WithZeroFrames is specified.
|
||||
|
||||
+53
-16
@@ -14,22 +14,23 @@ type EOption func(*encoderOptions) error
|
||||
|
||||
// options retains accumulated state of multiple options.
|
||||
type encoderOptions struct {
|
||||
resetOpt bool
|
||||
concurrent int
|
||||
level EncoderLevel
|
||||
single *bool
|
||||
pad int
|
||||
blockSize int
|
||||
windowSize int
|
||||
crc bool
|
||||
fullZero bool
|
||||
noEntropy bool
|
||||
allLitEntropy bool
|
||||
customWindow bool
|
||||
customALEntropy bool
|
||||
customBlockSize bool
|
||||
lowMem bool
|
||||
dict *dict
|
||||
resetOpt bool
|
||||
concurrent int
|
||||
level EncoderLevel
|
||||
single *bool
|
||||
pad int
|
||||
blockSize int
|
||||
windowSize int
|
||||
crc bool
|
||||
fullZero bool
|
||||
noEntropy bool
|
||||
allLitEntropy bool
|
||||
customWindow bool
|
||||
customALEntropy bool
|
||||
customBlockSize bool
|
||||
lowMem bool
|
||||
dict *dict
|
||||
concurrentBlocks bool
|
||||
}
|
||||
|
||||
func (o *encoderOptions) setDefault() {
|
||||
@@ -333,6 +334,42 @@ func WithLowerEncoderMem(b bool) EOption {
|
||||
}
|
||||
}
|
||||
|
||||
// WithConcurrentBlocks enables job-based parallel compression for streams.
|
||||
// When enabled and concurrent > 1, input is split into large sections (jobs)
|
||||
// that are compressed simultaneously by multiple goroutines.
|
||||
// Each non-first job receives an overlap prefix from the previous job for match context.
|
||||
// Output is flushed in order, producing a valid single-frame zstd stream.
|
||||
//
|
||||
// Currently disabled when used with dictionary encoding.
|
||||
// Cannot be changed with ResetWithOptions.
|
||||
func WithConcurrentBlocks(b bool) EOption {
|
||||
return func(o *encoderOptions) error {
|
||||
if o.resetOpt && b != o.concurrentBlocks {
|
||||
return errors.New("WithConcurrentBlocks cannot be changed on Reset")
|
||||
}
|
||||
o.concurrentBlocks = b
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// jobSize returns the input section size per parallel job.
|
||||
func (o *encoderOptions) jobSize() int {
|
||||
s := max(o.windowSize*4, 512<<10)
|
||||
return s
|
||||
}
|
||||
|
||||
// overlapSize returns the overlap prefix size for parallel jobs.
|
||||
func (o *encoderOptions) overlapSize() int {
|
||||
switch o.level {
|
||||
case SpeedBestCompression:
|
||||
return o.windowSize / 2
|
||||
case SpeedBetterCompression:
|
||||
return o.windowSize / 4
|
||||
default:
|
||||
return o.windowSize / 8
|
||||
}
|
||||
}
|
||||
|
||||
// WithEncoderDict allows to register a dictionary that will be used for the encode.
|
||||
//
|
||||
// The slice dict must be in the [dictionary format] produced by
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
// Code generated by command: go run gen_fse.go -out ../fse_decoder_amd64.s -pkg=zstd. DO NOT EDIT.
|
||||
// Code generated by command: go run gen_fse.go -out ../fse_decoder.s -arch amd64,arm64 -pkg=zstd. DO NOT EDIT.
|
||||
|
||||
//go:build !appengine && !noasm && gc && !noasm
|
||||
|
||||
|
||||
+153
@@ -0,0 +1,153 @@
|
||||
// Code generated by command: go run gen_fse.go -out ../fse_decoder.s -arch amd64,arm64 -pkg=zstd. DO NOT EDIT.
|
||||
// EXPERIMENTAL arm64 output lowered from an amd64 avo program.
|
||||
|
||||
//go:build arm64 && !appengine && !noasm && gc && !noasm
|
||||
|
||||
// func buildDtable_asm(s *fseDecoder, ctx *buildDtableAsmContext) int
|
||||
TEXT ·buildDtable_asm(SB), $0-24
|
||||
MOVD ctx+8(FP), R1
|
||||
MOVD s+0(FP), R6
|
||||
|
||||
// Load values
|
||||
MOVBU 4098(R6), R2
|
||||
MOVD $0, R0
|
||||
MOVD $1, R16
|
||||
LSL R2, R16, R16
|
||||
ORR R16, R0, R0
|
||||
MOVD (R1), R3
|
||||
MOVD 16(R1), R5
|
||||
SUB $1, R0, R7
|
||||
MOVD 8(R1), R1
|
||||
MOVHU 4096(R6), R6
|
||||
|
||||
// End load values
|
||||
// Init, lay down lowprob symbols
|
||||
MOVD $0, R8
|
||||
JMP init_main_loop_condition
|
||||
|
||||
init_main_loop:
|
||||
ADD R8<<1, R1, R15
|
||||
MOVH (R15), R9
|
||||
AND $0xffff, R9, R15
|
||||
MOVD $-1, R16
|
||||
AND $0xffff, R16, R16
|
||||
CMP R16, R15
|
||||
BNE do_not_update_high_threshold
|
||||
ADD R7<<3, R5, R15
|
||||
MOVB R8, 1(R15)
|
||||
SUB $1, R7, R7
|
||||
MOVD $0x0000000000000001, R9
|
||||
|
||||
do_not_update_high_threshold:
|
||||
ADD R8<<1, R3, R15
|
||||
MOVH R9, (R15)
|
||||
ADD $1, R8, R8
|
||||
|
||||
init_main_loop_condition:
|
||||
CMP R6, R8
|
||||
BLT init_main_loop
|
||||
|
||||
// Spread symbols
|
||||
// Calculate table step
|
||||
MOVD R0, R8
|
||||
LSR $0x01, R8, R8
|
||||
MOVD R0, R9
|
||||
LSR $0x03, R9, R9
|
||||
ADD R9, R8, R8
|
||||
ADD $3, R8, R8
|
||||
|
||||
// Fill add bits values
|
||||
SUB $1, R0, R9
|
||||
MOVD $0, R10
|
||||
MOVD $0, R11
|
||||
JMP spread_main_loop_condition
|
||||
|
||||
spread_main_loop:
|
||||
MOVD $0, R12
|
||||
ADD R11<<1, R1, R15
|
||||
MOVH (R15), R13
|
||||
JMP spread_inner_loop_condition
|
||||
|
||||
spread_inner_loop:
|
||||
ADD R10<<3, R5, R15
|
||||
MOVB R11, 1(R15)
|
||||
|
||||
adjust_position:
|
||||
ADD R8, R10, R10
|
||||
AND R9, R10, R10
|
||||
CMP R7, R10
|
||||
BGT adjust_position
|
||||
ADD $1, R12, R12
|
||||
|
||||
spread_inner_loop_condition:
|
||||
CMP R13, R12
|
||||
BLT spread_inner_loop
|
||||
ADD $1, R11, R11
|
||||
|
||||
spread_main_loop_condition:
|
||||
CMP R6, R11
|
||||
BLT spread_main_loop
|
||||
TST R10, R10
|
||||
BEQ spread_check_ok
|
||||
MOVD ctx+8(FP), R0
|
||||
MOVD R10, 24(R0)
|
||||
MOVD $+1, R16
|
||||
MOVD R16, ret+16(FP)
|
||||
RET
|
||||
|
||||
spread_check_ok:
|
||||
// Build Decoding table
|
||||
MOVD $0, R6
|
||||
|
||||
build_table_main_table:
|
||||
ADD R6<<3, R5, R15
|
||||
MOVBU 1(R15), R1
|
||||
ADD R1<<1, R3, R15
|
||||
MOVHU (R15), R7
|
||||
ADD $1, R7, R8
|
||||
ADD R1<<1, R3, R15
|
||||
MOVH R8, (R15)
|
||||
MOVD R7, R8
|
||||
CLZ R8, R16
|
||||
MOVD $63, R8
|
||||
SUB R16, R8, R8
|
||||
MOVD R2, R1
|
||||
SUB R8, R1, R1
|
||||
LSL R1, R7, R7
|
||||
SUB R0, R7, R7
|
||||
ADD R6<<3, R5, R15
|
||||
MOVB R1, (R15)
|
||||
ADD R6<<3, R5, R15
|
||||
MOVH R7, 2(R15)
|
||||
CMP R0, R7
|
||||
BLE build_table_check1_ok
|
||||
MOVD ctx+8(FP), R1
|
||||
MOVD R7, 24(R1)
|
||||
MOVD R0, 32(R1)
|
||||
MOVD $+2, R16
|
||||
MOVD R16, ret+16(FP)
|
||||
RET
|
||||
|
||||
build_table_check1_ok:
|
||||
AND $0xff, R1, R15
|
||||
AND $0xff, R1, R16
|
||||
TST R16, R15
|
||||
BNE build_table_check2_ok
|
||||
AND $0xffff, R7, R15
|
||||
AND $0xffff, R6, R16
|
||||
CMP R16, R15
|
||||
BNE build_table_check2_ok
|
||||
MOVD ctx+8(FP), R0
|
||||
MOVD R7, 24(R0)
|
||||
MOVD R6, 32(R0)
|
||||
MOVD $+3, R16
|
||||
MOVD R16, ret+16(FP)
|
||||
RET
|
||||
|
||||
build_table_check2_ok:
|
||||
ADD $1, R6, R6
|
||||
CMP R0, R6
|
||||
BLT build_table_main_table
|
||||
MOVD $+0, R16
|
||||
MOVD R16, ret+16(FP)
|
||||
RET
|
||||
Generated
Vendored
+6
-2
@@ -1,4 +1,4 @@
|
||||
//go:build amd64 && !appengine && !noasm && gc
|
||||
//go:build (amd64 || arm64) && !appengine && !noasm && gc
|
||||
|
||||
package zstd
|
||||
|
||||
@@ -6,6 +6,10 @@ import (
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// buildDtable_asm is generated by _generate/gen_fse.go and lowered to each
|
||||
// architecture (amd64 by goasm, arm64 by the avo arm64 lowering printer). The
|
||||
// Go side is identical across architectures, so it lives here.
|
||||
|
||||
type buildDtableAsmContext struct {
|
||||
// inputs
|
||||
stateTable *uint16
|
||||
@@ -18,7 +22,7 @@ type buildDtableAsmContext struct {
|
||||
errParam2 uint64
|
||||
}
|
||||
|
||||
// buildDtable_asm is an x86 assembly implementation of fseDecoder.buildDtable.
|
||||
// buildDtable_asm is an assembly implementation of fseDecoder.buildDtable.
|
||||
// Function returns non-zero exit code on error.
|
||||
//
|
||||
//go:noescape
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
//go:build !amd64 || appengine || !gc || noasm
|
||||
//go:build (!amd64 && !arm64) || appengine || !gc || noasm
|
||||
|
||||
package zstd
|
||||
|
||||
|
||||
+51
-333
@@ -3,30 +3,47 @@
|
||||
package zstd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/klauspost/compress/internal/cpuinfo"
|
||||
)
|
||||
|
||||
type decodeSyncAsmContext struct {
|
||||
llTable []decSymbol
|
||||
mlTable []decSymbol
|
||||
ofTable []decSymbol
|
||||
llState uint64
|
||||
mlState uint64
|
||||
ofState uint64
|
||||
iteration int
|
||||
litRemain int
|
||||
out []byte
|
||||
outPosition int
|
||||
literals []byte
|
||||
litPosition int
|
||||
history []byte
|
||||
windowSize int
|
||||
ll int // set on error (not for all errors, please refer to _generate/gen.go)
|
||||
ml int // set on error (not for all errors, please refer to _generate/gen.go)
|
||||
mo int // set on error (not for all errors, please refer to _generate/gen.go)
|
||||
// The shared decode/decodeSync/executeSimple wrappers and context structs live
|
||||
// in seqdec_asm.go; this file only declares the amd64 asm routines and the
|
||||
// dispatch helpers that pick the BMI2 / non-BMI2 (and 56-bit / safe) variant.
|
||||
|
||||
// sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm.
|
||||
//
|
||||
// Please refer to seqdec_generic.go for the reference implementation.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_amd64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode_56_amd64 implements the main loop of sequenceDecs in x86 asm.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_56_amd64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode_bmi2 implements the main loop of sequenceDecs in x86 asm with BMI2 extensions.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode_56_bmi2 implements the main loop of sequenceDecs in x86 asm with BMI2 extensions.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_56_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// decodeAsm runs the sequenceDecs decode loop, choosing the BMI2 / 56-bit variant.
|
||||
func decodeAsm(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext, lte56bits bool) int {
|
||||
if cpuinfo.HasBMI2() {
|
||||
if lte56bits {
|
||||
return sequenceDecs_decode_56_bmi2(s, br, ctx)
|
||||
}
|
||||
return sequenceDecs_decode_bmi2(s, br, ctx)
|
||||
}
|
||||
if lte56bits {
|
||||
return sequenceDecs_decode_56_amd64(s, br, ctx)
|
||||
}
|
||||
return sequenceDecs_decode_amd64(s, br, ctx)
|
||||
}
|
||||
|
||||
// sequenceDecs_decodeSync_amd64 implements the main loop of sequenceDecs.decodeSync in x86 asm.
|
||||
@@ -51,273 +68,18 @@ func sequenceDecs_decodeSync_safe_amd64(s *sequenceDecs, br *bitReader, ctx *dec
|
||||
//go:noescape
|
||||
func sequenceDecs_decodeSync_safe_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
|
||||
|
||||
// decode sequences from the stream with the provided history but without a dictionary.
|
||||
func (s *sequenceDecs) decodeSyncSimple(hist []byte) (bool, error) {
|
||||
if len(s.dict) > 0 {
|
||||
return false, nil
|
||||
}
|
||||
if s.maxSyncLen == 0 && cap(s.out)-len(s.out) < maxCompressedBlockSize {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// FIXME: Using unsafe memory copies leads to rare, random crashes
|
||||
// with fuzz testing. It is therefore disabled for now.
|
||||
const useSafe = true
|
||||
/*
|
||||
useSafe := false
|
||||
if s.maxSyncLen == 0 && cap(s.out)-len(s.out) < maxCompressedBlockSizeAlloc {
|
||||
useSafe = true
|
||||
}
|
||||
if s.maxSyncLen > 0 && cap(s.out)-len(s.out)-compressedBlockOverAlloc < int(s.maxSyncLen) {
|
||||
useSafe = true
|
||||
}
|
||||
if cap(s.literals) < len(s.literals)+compressedBlockOverAlloc {
|
||||
useSafe = true
|
||||
}
|
||||
*/
|
||||
|
||||
br := s.br
|
||||
|
||||
maxBlockSize := min(s.windowSize, maxCompressedBlockSize)
|
||||
|
||||
ctx := decodeSyncAsmContext{
|
||||
llTable: s.litLengths.fse.dt[:maxTablesize],
|
||||
mlTable: s.matchLengths.fse.dt[:maxTablesize],
|
||||
ofTable: s.offsets.fse.dt[:maxTablesize],
|
||||
llState: uint64(s.litLengths.state.state),
|
||||
mlState: uint64(s.matchLengths.state.state),
|
||||
ofState: uint64(s.offsets.state.state),
|
||||
iteration: s.nSeqs - 1,
|
||||
litRemain: len(s.literals),
|
||||
out: s.out,
|
||||
outPosition: len(s.out),
|
||||
literals: s.literals,
|
||||
windowSize: s.windowSize,
|
||||
history: hist,
|
||||
}
|
||||
|
||||
s.seqSize = 0
|
||||
startSize := len(s.out)
|
||||
|
||||
var errCode int
|
||||
// decodeSyncAsm runs the decodeSync loop, choosing the BMI2 / safe variant.
|
||||
func decodeSyncAsm(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext, safe bool) int {
|
||||
if cpuinfo.HasBMI2() {
|
||||
if useSafe {
|
||||
errCode = sequenceDecs_decodeSync_safe_bmi2(s, br, &ctx)
|
||||
} else {
|
||||
errCode = sequenceDecs_decodeSync_bmi2(s, br, &ctx)
|
||||
}
|
||||
} else {
|
||||
if useSafe {
|
||||
errCode = sequenceDecs_decodeSync_safe_amd64(s, br, &ctx)
|
||||
} else {
|
||||
errCode = sequenceDecs_decodeSync_amd64(s, br, &ctx)
|
||||
if safe {
|
||||
return sequenceDecs_decodeSync_safe_bmi2(s, br, ctx)
|
||||
}
|
||||
return sequenceDecs_decodeSync_bmi2(s, br, ctx)
|
||||
}
|
||||
switch errCode {
|
||||
case noError:
|
||||
break
|
||||
|
||||
case errorMatchLenOfsMismatch:
|
||||
return true, fmt.Errorf("zero matchoff and matchlen (%d) > 0", ctx.ml)
|
||||
|
||||
case errorMatchLenTooBig:
|
||||
return true, fmt.Errorf("match len (%d) bigger than max allowed length", ctx.ml)
|
||||
|
||||
case errorMatchOffTooBig:
|
||||
return true, fmt.Errorf("match offset (%d) bigger than current history (%d)",
|
||||
ctx.mo, ctx.outPosition+len(hist)-startSize)
|
||||
|
||||
case errorNotEnoughLiterals:
|
||||
return true, fmt.Errorf("unexpected literal count, want %d bytes, but only %d is available",
|
||||
ctx.ll, ctx.litRemain+ctx.ll)
|
||||
|
||||
case errorOverread:
|
||||
return true, io.ErrUnexpectedEOF
|
||||
|
||||
case errorNotEnoughSpace:
|
||||
size := ctx.outPosition + ctx.ll + ctx.ml
|
||||
if debugDecoder {
|
||||
println("msl:", s.maxSyncLen, "cap", cap(s.out), "bef:", startSize, "sz:", size-startSize, "mbs:", maxBlockSize, "outsz:", cap(s.out)-startSize)
|
||||
}
|
||||
return true, fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
|
||||
|
||||
default:
|
||||
return true, fmt.Errorf("sequenceDecs_decode returned erroneous code %d", errCode)
|
||||
if safe {
|
||||
return sequenceDecs_decodeSync_safe_amd64(s, br, ctx)
|
||||
}
|
||||
|
||||
s.seqSize += ctx.litRemain
|
||||
if s.seqSize > maxBlockSize {
|
||||
return true, fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
|
||||
}
|
||||
err := br.close()
|
||||
if err != nil {
|
||||
printf("Closing sequences: %v, %+v\n", err, *br)
|
||||
return true, err
|
||||
}
|
||||
|
||||
s.literals = s.literals[ctx.litPosition:]
|
||||
t := ctx.outPosition
|
||||
s.out = s.out[:t]
|
||||
|
||||
// Add final literals
|
||||
s.out = append(s.out, s.literals...)
|
||||
if debugDecoder {
|
||||
t += len(s.literals)
|
||||
if t != len(s.out) {
|
||||
panic(fmt.Errorf("length mismatch, want %d, got %d", len(s.out), t))
|
||||
}
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------------
|
||||
|
||||
type decodeAsmContext struct {
|
||||
llTable []decSymbol
|
||||
mlTable []decSymbol
|
||||
ofTable []decSymbol
|
||||
llState uint64
|
||||
mlState uint64
|
||||
ofState uint64
|
||||
iteration int
|
||||
seqs []seqVals
|
||||
litRemain int
|
||||
}
|
||||
|
||||
const noError = 0
|
||||
|
||||
// error reported when mo == 0 && ml > 0
|
||||
const errorMatchLenOfsMismatch = 1
|
||||
|
||||
// error reported when ml > maxMatchLen
|
||||
const errorMatchLenTooBig = 2
|
||||
|
||||
// error reported when mo > available history or mo > s.windowSize
|
||||
const errorMatchOffTooBig = 3
|
||||
|
||||
// error reported when the sum of literal lengths exeeceds the literal buffer size
|
||||
const errorNotEnoughLiterals = 4
|
||||
|
||||
// error reported when capacity of `out` is too small
|
||||
const errorNotEnoughSpace = 5
|
||||
|
||||
// error reported when bits are overread.
|
||||
const errorOverread = 6
|
||||
|
||||
// sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm.
|
||||
//
|
||||
// Please refer to seqdec_generic.go for the reference implementation.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_amd64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm.
|
||||
//
|
||||
// Please refer to seqdec_generic.go for the reference implementation.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_56_amd64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm with BMI2 extensions.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode implements the main loop of sequenceDecs in x86 asm with BMI2 extensions.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_56_bmi2(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// decode sequences from the stream without the provided history.
|
||||
func (s *sequenceDecs) decode(seqs []seqVals) error {
|
||||
br := s.br
|
||||
|
||||
maxBlockSize := min(s.windowSize, maxCompressedBlockSize)
|
||||
|
||||
ctx := decodeAsmContext{
|
||||
llTable: s.litLengths.fse.dt[:maxTablesize],
|
||||
mlTable: s.matchLengths.fse.dt[:maxTablesize],
|
||||
ofTable: s.offsets.fse.dt[:maxTablesize],
|
||||
llState: uint64(s.litLengths.state.state),
|
||||
mlState: uint64(s.matchLengths.state.state),
|
||||
ofState: uint64(s.offsets.state.state),
|
||||
seqs: seqs,
|
||||
iteration: len(seqs) - 1,
|
||||
litRemain: len(s.literals),
|
||||
}
|
||||
|
||||
if debugDecoder {
|
||||
println("decode: decoding", len(seqs), "sequences", br.remain(), "bits remain on stream")
|
||||
}
|
||||
|
||||
s.seqSize = 0
|
||||
lte56bits := s.maxBits+s.offsets.fse.actualTableLog+s.matchLengths.fse.actualTableLog+s.litLengths.fse.actualTableLog <= 56
|
||||
var errCode int
|
||||
if cpuinfo.HasBMI2() {
|
||||
if lte56bits {
|
||||
errCode = sequenceDecs_decode_56_bmi2(s, br, &ctx)
|
||||
} else {
|
||||
errCode = sequenceDecs_decode_bmi2(s, br, &ctx)
|
||||
}
|
||||
} else {
|
||||
if lte56bits {
|
||||
errCode = sequenceDecs_decode_56_amd64(s, br, &ctx)
|
||||
} else {
|
||||
errCode = sequenceDecs_decode_amd64(s, br, &ctx)
|
||||
}
|
||||
}
|
||||
if errCode != 0 {
|
||||
i := len(seqs) - ctx.iteration - 1
|
||||
switch errCode {
|
||||
case errorMatchLenOfsMismatch:
|
||||
ml := ctx.seqs[i].ml
|
||||
return fmt.Errorf("zero matchoff and matchlen (%d) > 0", ml)
|
||||
|
||||
case errorMatchLenTooBig:
|
||||
ml := ctx.seqs[i].ml
|
||||
return fmt.Errorf("match len (%d) bigger than max allowed length", ml)
|
||||
|
||||
case errorNotEnoughLiterals:
|
||||
ll := ctx.seqs[i].ll
|
||||
return fmt.Errorf("unexpected literal count, want %d bytes, but only %d is available", ll, ctx.litRemain+ll)
|
||||
case errorOverread:
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
|
||||
return fmt.Errorf("sequenceDecs_decode_amd64 returned erroneous code %d", errCode)
|
||||
}
|
||||
|
||||
if ctx.litRemain < 0 {
|
||||
return fmt.Errorf("literal count is too big: total available %d, total requested %d",
|
||||
len(s.literals), len(s.literals)-ctx.litRemain)
|
||||
}
|
||||
|
||||
s.seqSize += ctx.litRemain
|
||||
if s.seqSize > maxBlockSize {
|
||||
return fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
|
||||
}
|
||||
if debugDecoder {
|
||||
println("decode: ", br.remain(), "bits remain on stream. code:", errCode)
|
||||
}
|
||||
err := br.close()
|
||||
if err != nil {
|
||||
printf("Closing sequences: %v, %+v\n", err, *br)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------------
|
||||
|
||||
type executeAsmContext struct {
|
||||
seqs []seqVals
|
||||
seqIndex int
|
||||
out []byte
|
||||
history []byte
|
||||
literals []byte
|
||||
outPosition int
|
||||
litPosition int
|
||||
windowSize int
|
||||
return sequenceDecs_decodeSync_amd64(s, br, ctx)
|
||||
}
|
||||
|
||||
// sequenceDecs_executeSimple_amd64 implements the main loop of sequenceDecs.executeSimple in x86 asm.
|
||||
@@ -334,54 +96,10 @@ func sequenceDecs_executeSimple_amd64(ctx *executeAsmContext) bool
|
||||
//go:noescape
|
||||
func sequenceDecs_executeSimple_safe_amd64(ctx *executeAsmContext) bool
|
||||
|
||||
// executeSimple handles cases when dictionary is not used.
|
||||
func (s *sequenceDecs) executeSimple(seqs []seqVals, hist []byte) error {
|
||||
// Ensure we have enough output size...
|
||||
if len(s.out)+s.seqSize+compressedBlockOverAlloc > cap(s.out) {
|
||||
addBytes := s.seqSize + len(s.out) + compressedBlockOverAlloc
|
||||
s.out = append(s.out, make([]byte, addBytes)...)
|
||||
s.out = s.out[:len(s.out)-addBytes]
|
||||
// executeSimpleAsm runs the executeSimple loop, choosing the safe variant.
|
||||
func executeSimpleAsm(ctx *executeAsmContext, safe bool) bool {
|
||||
if safe {
|
||||
return sequenceDecs_executeSimple_safe_amd64(ctx)
|
||||
}
|
||||
|
||||
if debugDecoder {
|
||||
printf("Execute %d seqs with literals: %d into %d bytes\n", len(seqs), len(s.literals), s.seqSize)
|
||||
}
|
||||
|
||||
var t = len(s.out)
|
||||
out := s.out[:t+s.seqSize]
|
||||
|
||||
ctx := executeAsmContext{
|
||||
seqs: seqs,
|
||||
seqIndex: 0,
|
||||
out: out,
|
||||
history: hist,
|
||||
outPosition: t,
|
||||
litPosition: 0,
|
||||
literals: s.literals,
|
||||
windowSize: s.windowSize,
|
||||
}
|
||||
var ok bool
|
||||
if cap(s.literals) < len(s.literals)+compressedBlockOverAlloc {
|
||||
ok = sequenceDecs_executeSimple_safe_amd64(&ctx)
|
||||
} else {
|
||||
ok = sequenceDecs_executeSimple_amd64(&ctx)
|
||||
}
|
||||
if !ok {
|
||||
return fmt.Errorf("match offset (%d) bigger than current history (%d)",
|
||||
seqs[ctx.seqIndex].mo, ctx.outPosition+len(hist))
|
||||
}
|
||||
s.literals = s.literals[ctx.litPosition:]
|
||||
t = ctx.outPosition
|
||||
|
||||
// Add final literals
|
||||
copy(out[t:], s.literals)
|
||||
if debugDecoder {
|
||||
t += len(s.literals)
|
||||
if t != len(out) {
|
||||
panic(fmt.Errorf("length mismatch, want %d, got %d, ss: %d", len(out), t, s.seqSize))
|
||||
}
|
||||
}
|
||||
s.out = out
|
||||
|
||||
return nil
|
||||
return sequenceDecs_executeSimple_amd64(ctx)
|
||||
}
|
||||
|
||||
+3
-3
@@ -1,4 +1,4 @@
|
||||
// Code generated by command: go run gen.go -out ../seqdec_amd64.s -pkg=zstd. DO NOT EDIT.
|
||||
// Code generated by command: go run gen.go -out ../seqdec.s -arch amd64,arm64 -pkg=zstd. DO NOT EDIT.
|
||||
|
||||
//go:build !appengine && !noasm && gc && !noasm
|
||||
|
||||
@@ -2081,7 +2081,7 @@ sequenceDecs_decodeSync_amd64_match_len_ofs_ok:
|
||||
MOVQ 16(SP), R13
|
||||
|
||||
// Check if we have enough space in s.out
|
||||
LEAQ (AX)(R13*1), R14
|
||||
LEAQ 16(AX)(R13*1), R14
|
||||
ADDQ R10, R14
|
||||
CMPQ R14, 32(SP)
|
||||
JA error_not_enough_space
|
||||
@@ -2601,7 +2601,7 @@ sequenceDecs_decodeSync_bmi2_match_len_ofs_ok:
|
||||
MOVQ 16(SP), R13
|
||||
|
||||
// Check if we have enough space in s.out
|
||||
LEAQ (CX)(R13*1), R14
|
||||
LEAQ 16(CX)(R13*1), R14
|
||||
ADDQ R9, R14
|
||||
CMPQ R14, 32(SP)
|
||||
JA error_not_enough_space
|
||||
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
//go:build arm64 && !appengine && !noasm && gc
|
||||
|
||||
package zstd
|
||||
|
||||
// The shared decode/decodeSync/executeSimple wrappers and context structs live
|
||||
// in seqdec_asm.go; this file only declares the arm64 asm routines (generated
|
||||
// by the avo arm64 lowering printer) and the dispatch helpers. arm64 has no
|
||||
// BMI2, so each helper selects only between the 56-bit / safe variants.
|
||||
|
||||
// sequenceDecs_decode_arm64 implements the main loop of sequenceDecs in arm64 asm.
|
||||
//
|
||||
// Please refer to seqdec_generic.go for the reference implementation.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_arm64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// sequenceDecs_decode_56_arm64 implements the main loop of sequenceDecs in arm64 asm.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decode_56_arm64(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext) int
|
||||
|
||||
// decodeAsm runs the sequenceDecs decode loop, choosing the 56-bit variant.
|
||||
func decodeAsm(s *sequenceDecs, br *bitReader, ctx *decodeAsmContext, lte56bits bool) int {
|
||||
if lte56bits {
|
||||
return sequenceDecs_decode_56_arm64(s, br, ctx)
|
||||
}
|
||||
return sequenceDecs_decode_arm64(s, br, ctx)
|
||||
}
|
||||
|
||||
// sequenceDecs_decodeSync_arm64 implements the main loop of sequenceDecs.decodeSync in arm64 asm.
|
||||
//
|
||||
// Please refer to seqdec_generic.go for the reference implementation.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decodeSync_arm64(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
|
||||
|
||||
// sequenceDecs_decodeSync_safe_arm64 does the same as above, but does not write more than output buffer.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_decodeSync_safe_arm64(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext) int
|
||||
|
||||
// decodeSyncAsm runs the decodeSync loop, choosing the safe variant.
|
||||
func decodeSyncAsm(s *sequenceDecs, br *bitReader, ctx *decodeSyncAsmContext, safe bool) int {
|
||||
if safe {
|
||||
return sequenceDecs_decodeSync_safe_arm64(s, br, ctx)
|
||||
}
|
||||
return sequenceDecs_decodeSync_arm64(s, br, ctx)
|
||||
}
|
||||
|
||||
// sequenceDecs_executeSimple_arm64 implements the main loop of sequenceDecs.executeSimple in arm64 asm.
|
||||
//
|
||||
// Returns false if a match offset is too big.
|
||||
//
|
||||
// Please refer to seqdec_generic.go for the reference implementation.
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_executeSimple_arm64(ctx *executeAsmContext) bool
|
||||
|
||||
// Same as above, but with safe memcopies
|
||||
//
|
||||
//go:noescape
|
||||
func sequenceDecs_executeSimple_safe_arm64(ctx *executeAsmContext) bool
|
||||
|
||||
// executeSimpleAsm runs the executeSimple loop, choosing the safe variant.
|
||||
func executeSimpleAsm(ctx *executeAsmContext, safe bool) bool {
|
||||
if safe {
|
||||
return sequenceDecs_executeSimple_safe_arm64(ctx)
|
||||
}
|
||||
return sequenceDecs_executeSimple_arm64(ctx)
|
||||
}
|
||||
+2750
File diff suppressed because it is too large
Load Diff
+316
@@ -0,0 +1,316 @@
|
||||
//go:build (amd64 || arm64) && !appengine && !noasm && gc
|
||||
|
||||
package zstd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
// This file holds the parts of the assembly sequence decoder that are identical
|
||||
// across architectures: the context structs exchanged with the asm, the error
|
||||
// codes, and the decode/decodeSync/executeSimple wrappers. Each architecture
|
||||
// supplies the small dispatch helpers (decodeAsm, decodeSyncAsm,
|
||||
// executeSimpleAsm) that select the concrete asm routine — amd64 also chooses a
|
||||
// BMI2 variant, arm64 has a single implementation.
|
||||
|
||||
type decodeSyncAsmContext struct {
|
||||
llTable []decSymbol
|
||||
mlTable []decSymbol
|
||||
ofTable []decSymbol
|
||||
llState uint64
|
||||
mlState uint64
|
||||
ofState uint64
|
||||
iteration int
|
||||
litRemain int
|
||||
out []byte
|
||||
outPosition int
|
||||
literals []byte
|
||||
litPosition int
|
||||
history []byte
|
||||
windowSize int
|
||||
ll int // set on error (not for all errors, please refer to _generate/gen.go)
|
||||
ml int // set on error (not for all errors, please refer to _generate/gen.go)
|
||||
mo int // set on error (not for all errors, please refer to _generate/gen.go)
|
||||
}
|
||||
|
||||
type decodeAsmContext struct {
|
||||
llTable []decSymbol
|
||||
mlTable []decSymbol
|
||||
ofTable []decSymbol
|
||||
llState uint64
|
||||
mlState uint64
|
||||
ofState uint64
|
||||
iteration int
|
||||
seqs []seqVals
|
||||
litRemain int
|
||||
}
|
||||
|
||||
type executeAsmContext struct {
|
||||
seqs []seqVals
|
||||
seqIndex int
|
||||
out []byte
|
||||
history []byte
|
||||
literals []byte
|
||||
outPosition int
|
||||
litPosition int
|
||||
windowSize int
|
||||
}
|
||||
|
||||
const noError = 0
|
||||
|
||||
// error reported when mo == 0 && ml > 0
|
||||
const errorMatchLenOfsMismatch = 1
|
||||
|
||||
// error reported when ml > maxMatchLen
|
||||
const errorMatchLenTooBig = 2
|
||||
|
||||
// error reported when mo > available history or mo > s.windowSize
|
||||
const errorMatchOffTooBig = 3
|
||||
|
||||
// error reported when the sum of literal lengths exeeceds the literal buffer size
|
||||
const errorNotEnoughLiterals = 4
|
||||
|
||||
// error reported when capacity of `out` is too small
|
||||
const errorNotEnoughSpace = 5
|
||||
|
||||
// error reported when bits are overread.
|
||||
const errorOverread = 6
|
||||
|
||||
// useSafeDecodeSync reports whether decodeSyncSimple must use the bounds-exact
|
||||
// ("safe") copy variants for the current buffer geometry.
|
||||
//
|
||||
// When the output and literal buffers have compressedBlockOverAlloc (16)
|
||||
// bytes of slack past their logical use, the assembly may use extended
|
||||
// memory copies that read and write in 16-byte blocks, overrunning the end
|
||||
// of a literal run or match by up to 15 bytes. Otherwise it must use the
|
||||
// bounds-exact ("safe") copies. This mirrors the analogous, always-dynamic
|
||||
// selection in executeSimple below.
|
||||
//
|
||||
// The unsafe copies were disabled in #644 (2022) as a mitigation for a
|
||||
// crash, but that crash's root cause — an unguarded bitReader overread that
|
||||
// produced out-of-range match offsets/lengths — was fixed three days later
|
||||
// in #645, which also added the fuzz corpus that has guarded this path since.
|
||||
// See #1168. An asan-instrumented fuzz job (see .github/workflows/go.yml)
|
||||
// covers the extended-copy path, which -race and plain fuzzing cannot.
|
||||
func (s *sequenceDecs) useSafeDecodeSync() bool {
|
||||
if s.maxSyncLen == 0 && cap(s.out)-len(s.out) < maxCompressedBlockSizeAlloc {
|
||||
return true
|
||||
}
|
||||
if s.maxSyncLen > 0 && cap(s.out)-len(s.out)-compressedBlockOverAlloc < int(s.maxSyncLen) {
|
||||
return true
|
||||
}
|
||||
if cap(s.literals) < len(s.literals)+compressedBlockOverAlloc {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// decode sequences from the stream with the provided history but without a dictionary.
|
||||
func (s *sequenceDecs) decodeSyncSimple(hist []byte) (bool, error) {
|
||||
if len(s.dict) > 0 {
|
||||
return false, nil
|
||||
}
|
||||
if s.maxSyncLen == 0 && cap(s.out)-len(s.out) < maxCompressedBlockSize {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
useSafe := s.useSafeDecodeSync()
|
||||
|
||||
br := s.br
|
||||
|
||||
maxBlockSize := min(s.windowSize, maxCompressedBlockSize)
|
||||
|
||||
ctx := decodeSyncAsmContext{
|
||||
llTable: s.litLengths.fse.dt[:maxTablesize],
|
||||
mlTable: s.matchLengths.fse.dt[:maxTablesize],
|
||||
ofTable: s.offsets.fse.dt[:maxTablesize],
|
||||
llState: uint64(s.litLengths.state.state),
|
||||
mlState: uint64(s.matchLengths.state.state),
|
||||
ofState: uint64(s.offsets.state.state),
|
||||
iteration: s.nSeqs - 1,
|
||||
litRemain: len(s.literals),
|
||||
out: s.out,
|
||||
outPosition: len(s.out),
|
||||
literals: s.literals,
|
||||
windowSize: s.windowSize,
|
||||
history: hist,
|
||||
}
|
||||
|
||||
s.seqSize = 0
|
||||
startSize := len(s.out)
|
||||
|
||||
errCode := decodeSyncAsm(s, br, &ctx, useSafe)
|
||||
switch errCode {
|
||||
case noError:
|
||||
break
|
||||
|
||||
case errorMatchLenOfsMismatch:
|
||||
return true, fmt.Errorf("zero matchoff and matchlen (%d) > 0", ctx.ml)
|
||||
|
||||
case errorMatchLenTooBig:
|
||||
return true, fmt.Errorf("match len (%d) bigger than max allowed length", ctx.ml)
|
||||
|
||||
case errorMatchOffTooBig:
|
||||
return true, fmt.Errorf("match offset (%d) bigger than current history (%d)",
|
||||
ctx.mo, ctx.outPosition+len(hist)-startSize)
|
||||
|
||||
case errorNotEnoughLiterals:
|
||||
return true, fmt.Errorf("unexpected literal count, want %d bytes, but only %d is available",
|
||||
ctx.ll, ctx.litRemain+ctx.ll)
|
||||
|
||||
case errorOverread:
|
||||
return true, io.ErrUnexpectedEOF
|
||||
|
||||
case errorNotEnoughSpace:
|
||||
size := ctx.outPosition + ctx.ll + ctx.ml
|
||||
if debugDecoder {
|
||||
println("msl:", s.maxSyncLen, "cap", cap(s.out), "bef:", startSize, "sz:", size-startSize, "mbs:", maxBlockSize, "outsz:", cap(s.out)-startSize)
|
||||
}
|
||||
return true, fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
|
||||
|
||||
default:
|
||||
return true, fmt.Errorf("sequenceDecs_decode returned erroneous code %d", errCode)
|
||||
}
|
||||
|
||||
s.seqSize += ctx.litRemain
|
||||
if s.seqSize > maxBlockSize {
|
||||
return true, fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
|
||||
}
|
||||
err := br.close()
|
||||
if err != nil {
|
||||
printf("Closing sequences: %v, %+v\n", err, *br)
|
||||
return true, err
|
||||
}
|
||||
|
||||
s.literals = s.literals[ctx.litPosition:]
|
||||
t := ctx.outPosition
|
||||
s.out = s.out[:t]
|
||||
|
||||
// Add final literals
|
||||
s.out = append(s.out, s.literals...)
|
||||
if debugDecoder {
|
||||
t += len(s.literals)
|
||||
if t != len(s.out) {
|
||||
panic(fmt.Errorf("length mismatch, want %d, got %d", len(s.out), t))
|
||||
}
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// decode sequences from the stream without the provided history.
|
||||
func (s *sequenceDecs) decode(seqs []seqVals) error {
|
||||
br := s.br
|
||||
|
||||
maxBlockSize := min(s.windowSize, maxCompressedBlockSize)
|
||||
|
||||
ctx := decodeAsmContext{
|
||||
llTable: s.litLengths.fse.dt[:maxTablesize],
|
||||
mlTable: s.matchLengths.fse.dt[:maxTablesize],
|
||||
ofTable: s.offsets.fse.dt[:maxTablesize],
|
||||
llState: uint64(s.litLengths.state.state),
|
||||
mlState: uint64(s.matchLengths.state.state),
|
||||
ofState: uint64(s.offsets.state.state),
|
||||
seqs: seqs,
|
||||
iteration: len(seqs) - 1,
|
||||
litRemain: len(s.literals),
|
||||
}
|
||||
|
||||
if debugDecoder {
|
||||
println("decode: decoding", len(seqs), "sequences", br.remain(), "bits remain on stream")
|
||||
}
|
||||
|
||||
s.seqSize = 0
|
||||
lte56bits := s.maxBits+s.offsets.fse.actualTableLog+s.matchLengths.fse.actualTableLog+s.litLengths.fse.actualTableLog <= 56
|
||||
errCode := decodeAsm(s, br, &ctx, lte56bits)
|
||||
if errCode != 0 {
|
||||
i := len(seqs) - ctx.iteration - 1
|
||||
switch errCode {
|
||||
case errorMatchLenOfsMismatch:
|
||||
ml := ctx.seqs[i].ml
|
||||
return fmt.Errorf("zero matchoff and matchlen (%d) > 0", ml)
|
||||
|
||||
case errorMatchLenTooBig:
|
||||
ml := ctx.seqs[i].ml
|
||||
return fmt.Errorf("match len (%d) bigger than max allowed length", ml)
|
||||
|
||||
case errorNotEnoughLiterals:
|
||||
ll := ctx.seqs[i].ll
|
||||
return fmt.Errorf("unexpected literal count, want %d bytes, but only %d is available", ll, ctx.litRemain+ll)
|
||||
case errorOverread:
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
|
||||
return fmt.Errorf("sequenceDecs_decode_amd64 returned erroneous code %d", errCode)
|
||||
}
|
||||
|
||||
if ctx.litRemain < 0 {
|
||||
return fmt.Errorf("literal count is too big: total available %d, total requested %d",
|
||||
len(s.literals), len(s.literals)-ctx.litRemain)
|
||||
}
|
||||
|
||||
s.seqSize += ctx.litRemain
|
||||
if s.seqSize > maxBlockSize {
|
||||
return fmt.Errorf("output bigger than max block size (%d)", maxBlockSize)
|
||||
}
|
||||
if debugDecoder {
|
||||
println("decode: ", br.remain(), "bits remain on stream. code:", errCode)
|
||||
}
|
||||
err := br.close()
|
||||
if err != nil {
|
||||
printf("Closing sequences: %v, %+v\n", err, *br)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// executeSimple handles cases when dictionary is not used.
|
||||
func (s *sequenceDecs) executeSimple(seqs []seqVals, hist []byte) error {
|
||||
// Ensure we have enough output size...
|
||||
if len(s.out)+s.seqSize+compressedBlockOverAlloc > cap(s.out) {
|
||||
addBytes := s.seqSize + len(s.out) + compressedBlockOverAlloc
|
||||
s.out = append(s.out, make([]byte, addBytes)...)
|
||||
s.out = s.out[:len(s.out)-addBytes]
|
||||
}
|
||||
|
||||
if debugDecoder {
|
||||
printf("Execute %d seqs with literals: %d into %d bytes\n", len(seqs), len(s.literals), s.seqSize)
|
||||
}
|
||||
|
||||
var t = len(s.out)
|
||||
out := s.out[:t+s.seqSize]
|
||||
|
||||
ctx := executeAsmContext{
|
||||
seqs: seqs,
|
||||
seqIndex: 0,
|
||||
out: out,
|
||||
history: hist,
|
||||
outPosition: t,
|
||||
litPosition: 0,
|
||||
literals: s.literals,
|
||||
windowSize: s.windowSize,
|
||||
}
|
||||
// useSafe avoids overwriting the output buffer when the literals slice has
|
||||
// not been allocated with the required over-allocation slack.
|
||||
useSafe := cap(s.literals) < len(s.literals)+compressedBlockOverAlloc
|
||||
|
||||
ok := executeSimpleAsm(&ctx, useSafe)
|
||||
if !ok {
|
||||
return fmt.Errorf("match offset (%d) bigger than current history (%d)",
|
||||
seqs[ctx.seqIndex].mo, ctx.outPosition+len(hist))
|
||||
}
|
||||
s.literals = s.literals[ctx.litPosition:]
|
||||
t = ctx.outPosition
|
||||
|
||||
// Add final literals
|
||||
copy(out[t:], s.literals)
|
||||
if debugDecoder {
|
||||
t += len(s.literals)
|
||||
if t != len(out) {
|
||||
panic(fmt.Errorf("length mismatch, want %d, got %d, ss: %d", len(out), t, s.seqSize))
|
||||
}
|
||||
}
|
||||
s.out = out
|
||||
|
||||
return nil
|
||||
}
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
//go:build !amd64 || appengine || !gc || noasm
|
||||
//go:build (!amd64 && !arm64) || appengine || !gc || noasm
|
||||
|
||||
package zstd
|
||||
|
||||
|
||||
+4
-3
@@ -334,9 +334,10 @@ func decodeSnappy(blk *blockEnc, src []byte) error {
|
||||
|
||||
return errUnsupportedLiteralLength
|
||||
}
|
||||
//if length > snappyMaxBlockSize-d || uint32(length) > len(src)-s {
|
||||
// return ErrSnappyCorrupt
|
||||
//}
|
||||
if length > len(src)-s {
|
||||
println("length > len(src)-s", length, len(src)-s)
|
||||
return ErrSnappyCorrupt
|
||||
}
|
||||
|
||||
blk.literals = append(blk.literals, src[s:s+length]...)
|
||||
//println(length, "litLen")
|
||||
|
||||
-1
@@ -1 +0,0 @@
|
||||
test
|
||||
@@ -2,6 +2,5 @@ coverage.out
|
||||
coverage.txt
|
||||
release-notes.txt
|
||||
.directory
|
||||
.chglog
|
||||
.vscode
|
||||
.DS_Store
|
||||
+99
-1
@@ -2,6 +2,94 @@
|
||||
## [Unreleased]
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a name="v0.12.4"></a>
|
||||
## [v0.12.4] - 2026-08-17
|
||||
### Fix
|
||||
- Make Interp robust to extreme values and exact knot hits ([#132](https://github.com/montanaflynn/stats/issues/132))
|
||||
- Reject NaN percent in PercentileWeighted ([#131](https://github.com/montanaflynn/stats/issues/131))
|
||||
- Reject NaN percent in Percentile and PercentileNearestRank ([#130](https://github.com/montanaflynn/stats/issues/130))
|
||||
- Interp panics when x or xp contains a NaN ([#129](https://github.com/montanaflynn/stats/issues/129))
|
||||
|
||||
|
||||
<a name="v0.12.3"></a>
|
||||
## [v0.12.3] - 2026-08-10
|
||||
### Fix
|
||||
- ProbGeom off-by-one that drops the first interval term ([#127](https://github.com/montanaflynn/stats/issues/127))
|
||||
- NormIsf ignores loc, NormPpf uncorrected above the median ([#126](https://github.com/montanaflynn/stats/issues/126))
|
||||
- Normal tail collapses to 0/-Inf past ~8 sigma ([#125](https://github.com/montanaflynn/stats/issues/125))
|
||||
|
||||
|
||||
<a name="v0.12.2"></a>
|
||||
## [v0.12.2] - 2026-07-17
|
||||
### Fix
|
||||
- Regression stability and invalid domains ([#124](https://github.com/montanaflynn/stats/issues/124))
|
||||
|
||||
|
||||
<a name="v0.12.1"></a>
|
||||
## [v0.12.1] - 2026-07-16
|
||||
### Fix
|
||||
- Stop Entropy from mutating its input slice ([#123](https://github.com/montanaflynn/stats/issues/123))
|
||||
|
||||
|
||||
<a name="v0.12.0"></a>
|
||||
## [v0.12.0] - 2026-07-16
|
||||
|
||||
<a name="v0.11.0"></a>
|
||||
## [v0.11.0] - 2026-07-13
|
||||
### Add
|
||||
- Add Interp for piecewise-linear interpolation ([#121](https://github.com/montanaflynn/stats/issues/121))
|
||||
- Add Histogram with equal-width bins ([#120](https://github.com/montanaflynn/stats/issues/120))
|
||||
- Add KendallTau rank correlation coefficient ([#119](https://github.com/montanaflynn/stats/issues/119))
|
||||
- Add SEM, RMS, Product, and PercentileOfScore ([#118](https://github.com/montanaflynn/stats/issues/118))
|
||||
- Add MovingMedian, MovingMin, MovingMax, MovingSum, and EWMA ([#117](https://github.com/montanaflynn/stats/issues/117))
|
||||
- Add TrimmedMean and Winsorize robust statistics ([#116](https://github.com/montanaflynn/stats/issues/116))
|
||||
- Add Kurtosis, PopulationKurtosis, and SampleKurtosis ([#115](https://github.com/montanaflynn/stats/issues/115))
|
||||
- Add Clip and Rescale elementwise transforms ([#114](https://github.com/montanaflynn/stats/issues/114))
|
||||
|
||||
|
||||
<a name="v0.10.0"></a>
|
||||
## [v0.10.0] - 2026-07-10
|
||||
### Add
|
||||
- Add MovingAverage and MovingStdDev ([#112](https://github.com/montanaflynn/stats/issues/112))
|
||||
- Add ZScore and Rank functions ([#111](https://github.com/montanaflynn/stats/issues/111))
|
||||
- Add WeightedMean and CoefficientOfVariation ([#110](https://github.com/montanaflynn/stats/issues/110))
|
||||
- Add ArgMax, ArgMin and Range functions ([#109](https://github.com/montanaflynn/stats/issues/109))
|
||||
- Add CumulativeProduct, CumulativeMax and CumulativeMin ([#108](https://github.com/montanaflynn/stats/issues/108))
|
||||
- Add Diff and PercentChange functions ([#107](https://github.com/montanaflynn/stats/issues/107))
|
||||
- Add weighted percentile function ([#102](https://github.com/montanaflynn/stats/issues/102))
|
||||
- Add NormSample function for normal distribution sampling ([#100](https://github.com/montanaflynn/stats/issues/100))
|
||||
- Add Z-test and T-test functions ([#99](https://github.com/montanaflynn/stats/issues/99))
|
||||
- Add Spearman rank correlation function ([#98](https://github.com/montanaflynn/stats/issues/98))
|
||||
|
||||
### Fix
|
||||
- Stabilize GeometricMean and add input validation
|
||||
- Use math.Round to avoid ARM64 FMA fusion miscompile ([#97](https://github.com/montanaflynn/stats/issues/97))
|
||||
- Correct AutoCorrelation lag handling ([#83](https://github.com/montanaflynn/stats/issues/83)) ([#95](https://github.com/montanaflynn/stats/issues/95))
|
||||
|
||||
|
||||
<a name="v0.9.0"></a>
|
||||
## [v0.9.0] - 2026-03-24
|
||||
### Add
|
||||
- Add Skewness, PopulationSkewness, and SampleSkewness functions ([#91](https://github.com/montanaflynn/stats/issues/91))
|
||||
|
||||
### Fix
|
||||
- Restore 100% test coverage for skewness
|
||||
- Remove unused sum[4] in LinearRegression
|
||||
|
||||
|
||||
<a name="v0.8.2"></a>
|
||||
## [v0.8.2] - 2026-03-11
|
||||
|
||||
<a name="v0.8.1"></a>
|
||||
## [v0.8.1] - 2026-03-11
|
||||
|
||||
<a name="v0.8.0"></a>
|
||||
## [v0.8.0] - 2026-03-11
|
||||
### Fix
|
||||
@@ -527,7 +615,17 @@
|
||||
- Merge pull request [#4](https://github.com/montanaflynn/stats/issues/4) from saromanov/sample
|
||||
|
||||
|
||||
[Unreleased]: https://github.com/montanaflynn/stats/compare/v0.8.0...HEAD
|
||||
[Unreleased]: https://github.com/montanaflynn/stats/compare/v0.12.4...HEAD
|
||||
[v0.12.4]: https://github.com/montanaflynn/stats/compare/v0.12.3...v0.12.4
|
||||
[v0.12.3]: https://github.com/montanaflynn/stats/compare/v0.12.2...v0.12.3
|
||||
[v0.12.2]: https://github.com/montanaflynn/stats/compare/v0.12.1...v0.12.2
|
||||
[v0.12.1]: https://github.com/montanaflynn/stats/compare/v0.12.0...v0.12.1
|
||||
[v0.12.0]: https://github.com/montanaflynn/stats/compare/v0.11.0...v0.12.0
|
||||
[v0.11.0]: https://github.com/montanaflynn/stats/compare/v0.10.0...v0.11.0
|
||||
[v0.10.0]: https://github.com/montanaflynn/stats/compare/v0.9.0...v0.10.0
|
||||
[v0.9.0]: https://github.com/montanaflynn/stats/compare/v0.8.2...v0.9.0
|
||||
[v0.8.2]: https://github.com/montanaflynn/stats/compare/v0.8.1...v0.8.2
|
||||
[v0.8.1]: https://github.com/montanaflynn/stats/compare/v0.8.0...v0.8.1
|
||||
[v0.8.0]: https://github.com/montanaflynn/stats/compare/v0.7.1...v0.8.0
|
||||
[v0.7.1]: https://github.com/montanaflynn/stats/compare/v0.7.0...v0.7.1
|
||||
[v0.7.0]: https://github.com/montanaflynn/stats/compare/v0.6.6...v0.7.0
|
||||
|
||||
+956
-69
File diff suppressed because it is too large
Load Diff
+4
-5
@@ -24,11 +24,10 @@ docs:
|
||||
godoc2md github.com/montanaflynn/stats | sed -e s#src/target/##g > DOCUMENTATION.md
|
||||
|
||||
release:
|
||||
git-chglog --output CHANGELOG.md --next-tag ${TAG}
|
||||
@test -n "${TAG}" || { echo "TAG is required, e.g. make release TAG=v0.10.0"; exit 1; }
|
||||
sh scripts/update-changelog.sh ${TAG}
|
||||
git add CHANGELOG.md
|
||||
git commit -m "Update changelog with ${TAG} changes"
|
||||
git tag ${TAG}
|
||||
git-chglog $(TAG) | tail -n +4 | gsed '1s/^/$(TAG)\n/gm' > release-notes.txt
|
||||
git commit -m "chore: update changelog for ${TAG}"
|
||||
git tag -a ${TAG} -m "${TAG}"
|
||||
git push origin master ${TAG}
|
||||
hub release create --copy -F release-notes.txt ${TAG}
|
||||
|
||||
|
||||
+53
-21
@@ -1,6 +1,6 @@
|
||||
# Stats - Golang Statistics Package
|
||||
|
||||
[![][action-svg]][action-url] [![][codecov-svg]][codecov-url] [![][goreport-svg]][goreport-url] [![][godoc-svg]][godoc-url] [![][pkggodev-svg]][pkggodev-url] [![][license-svg]][license-url]
|
||||
[![][action-svg]][action-url] [![][codecov-svg]][codecov-url] [![][pkggodev-svg]][pkggodev-url] [![][license-svg]][license-url]
|
||||
|
||||
A well tested and comprehensive Golang statistics library / package / module with no dependencies.
|
||||
|
||||
@@ -14,7 +14,7 @@ go get github.com/montanaflynn/stats
|
||||
|
||||
## Example Usage
|
||||
|
||||
All the functions can be seen in [examples/main.go](examples/main.go) but here's a little taste:
|
||||
All the functions can be seen in [examples/functions/main.go](examples/functions/main.go) but here's a little taste:
|
||||
|
||||
```go
|
||||
// start with some source data to use
|
||||
@@ -34,20 +34,19 @@ fmt.Println(roundedMedian) // 4
|
||||
|
||||
## Documentation
|
||||
|
||||
The entire API documentation is available on [GoDoc.org](http://godoc.org/github.com/montanaflynn/stats) or [pkg.go.dev](https://pkg.go.dev/github.com/montanaflynn/stats).
|
||||
The entire API documentation is available on [pkg.go.dev](https://pkg.go.dev/github.com/montanaflynn/stats).
|
||||
|
||||
You can also view docs offline with the following commands:
|
||||
|
||||
```
|
||||
# Command line
|
||||
godoc . # show all exported apis
|
||||
godoc . Median # show a single function
|
||||
godoc -ex . Round # show function with example
|
||||
godoc . Float64Data # show the type and methods
|
||||
go doc -all . # show all exported apis
|
||||
go doc Median # show a single function
|
||||
go doc Float64Data # show the type and methods
|
||||
|
||||
# Local website
|
||||
godoc -http=:4444 # start the godoc server on port 4444
|
||||
open http://localhost:4444/pkg/github.com/montanaflynn/stats/
|
||||
# Local website (go install golang.org/x/pkgsite/cmd/pkgsite@latest)
|
||||
pkgsite -http=:4444 # start the pkgsite server on port 4444
|
||||
open http://localhost:4444/github.com/montanaflynn/stats
|
||||
```
|
||||
|
||||
The exported API is as follows:
|
||||
@@ -70,19 +69,32 @@ type Float64Data []float64
|
||||
|
||||
func LoadRawData(raw interface{}) (f Float64Data) {}
|
||||
|
||||
func ArgMax(input Float64Data) (int, error) {}
|
||||
func ArgMin(input Float64Data) (int, error) {}
|
||||
func AutoCorrelation(data Float64Data, lags int) (float64, error) {}
|
||||
func ChebyshevDistance(dataPointX, dataPointY Float64Data) (distance float64, err error) {}
|
||||
func Clip(input Float64Data, min, max float64) ([]float64, error) {}
|
||||
func CoefficientOfVariation(input Float64Data) (float64, error) {}
|
||||
func Correlation(data1, data2 Float64Data) (float64, error) {}
|
||||
func Covariance(data1, data2 Float64Data) (float64, error) {}
|
||||
func CovariancePopulation(data1, data2 Float64Data) (float64, error) {}
|
||||
func CumulativeMax(input Float64Data) ([]float64, error) {}
|
||||
func CumulativeMin(input Float64Data) ([]float64, error) {}
|
||||
func CumulativeProduct(input Float64Data) ([]float64, error) {}
|
||||
func CumulativeSum(input Float64Data) ([]float64, error) {}
|
||||
func Describe(input Float64Data, allowNaN bool, percentiles *[]float64) (*Description, error) {}
|
||||
func DescribePercentileFunc(input Float64Data, allowNaN bool, percentiles *[]float64, percentileFunc func(Float64Data, float64) (float64, error)) (*Description, error) {}
|
||||
func Diff(input Float64Data) ([]float64, error) {}
|
||||
func EWMA(input Float64Data, alpha float64) ([]float64, error) {}
|
||||
func Entropy(input Float64Data) (float64, error) {}
|
||||
func EuclideanDistance(dataPointX, dataPointY Float64Data) (distance float64, err error) {}
|
||||
func GeometricMean(input Float64Data) (float64, error) {}
|
||||
func HarmonicMean(input Float64Data) (float64, error) {}
|
||||
func Histogram(input Float64Data, bins int) ([]int, []float64, error) {}
|
||||
func InterQuartileRange(input Float64Data) (float64, error) {}
|
||||
func Interp(x, xp, fp Float64Data) ([]float64, error) {}
|
||||
func KendallTau(data1, data2 Float64Data) (float64, error) {}
|
||||
func Kurtosis(input Float64Data) (float64, error) {}
|
||||
func ManhattanDistance(dataPointX, dataPointY Float64Data) (distance float64, err error) {}
|
||||
func Max(input Float64Data) (max float64, err error) {}
|
||||
func Mean(input Float64Data) (float64, error) {}
|
||||
@@ -93,6 +105,12 @@ func Midhinge(input Float64Data) (float64, error) {}
|
||||
func Min(input Float64Data) (min float64, err error) {}
|
||||
func MinkowskiDistance(dataPointX, dataPointY Float64Data, lambda float64) (distance float64, err error) {}
|
||||
func Mode(input Float64Data) (mode []float64, err error) {}
|
||||
func MovingAverage(input Float64Data, window int) ([]float64, error) {}
|
||||
func MovingMax(input Float64Data, window int) ([]float64, error) {}
|
||||
func MovingMedian(input Float64Data, window int) ([]float64, error) {}
|
||||
func MovingMin(input Float64Data, window int) ([]float64, error) {}
|
||||
func MovingStdDev(input Float64Data, window int) ([]float64, error) {}
|
||||
func MovingSum(input Float64Data, window int) ([]float64, error) {}
|
||||
func NormBoxMullerRvs(loc float64, scale float64, size int) []float64 {}
|
||||
func NormCdf(x float64, loc float64, scale float64) float64 {}
|
||||
func NormEntropy(loc float64, scale float64) float64 {}
|
||||
@@ -107,20 +125,33 @@ func NormMedian(loc float64, scale float64) float64 {}
|
||||
func NormMoment(n int, loc float64, scale float64) float64 {}
|
||||
func NormPdf(x float64, loc float64, scale float64) float64 {}
|
||||
func NormPpf(p float64, loc float64, scale float64) (x float64) {}
|
||||
func NormSample(loc float64, scale float64, size int) []float64 {}
|
||||
func NormPpfRvs(loc float64, scale float64, size int) []float64 {}
|
||||
func NormSf(x float64, loc float64, scale float64) float64 {}
|
||||
func NormStats(loc float64, scale float64, moments string) []float64 {}
|
||||
func NormStd(loc float64, scale float64) float64 {}
|
||||
func NormVar(loc float64, scale float64) float64 {}
|
||||
func Pearson(data1, data2 Float64Data) (float64, error) {}
|
||||
func PercentChange(input Float64Data) ([]float64, error) {}
|
||||
func Percentile(input Float64Data, percent float64) (percentile float64, err error) {}
|
||||
func PercentileNearestRank(input Float64Data, percent float64) (percentile float64, err error) {}
|
||||
func PercentileOfScore(input Float64Data, score float64) (float64, error) {}
|
||||
func PercentileWeighted(data, weights Float64Data, percent float64) (percentile float64, err error) {}
|
||||
func PopulationKurtosis(input Float64Data) (float64, error) {}
|
||||
func PopulationSkewness(input Float64Data) (float64, error) {}
|
||||
func PopulationVariance(input Float64Data) (pvar float64, err error) {}
|
||||
func Product(input Float64Data) (float64, error) {}
|
||||
func RMS(input Float64Data) (float64, error) {}
|
||||
func Range(input Float64Data) (float64, error) {}
|
||||
func Rank(input Float64Data) ([]float64, error) {}
|
||||
func Rescale(input Float64Data) ([]float64, error) {}
|
||||
func SEM(input Float64Data) (float64, error) {}
|
||||
func Sample(input Float64Data, takenum int, replacement bool) ([]float64, error) {}
|
||||
func SampleKurtosis(input Float64Data) (float64, error) {}
|
||||
func SampleSkewness(input Float64Data) (float64, error) {}
|
||||
func SampleVariance(input Float64Data) (svar float64, err error) {}
|
||||
func Skewness(input Float64Data) (float64, error) {}
|
||||
func Spearman(data1, data2 Float64Data) (float64, error) {}
|
||||
func Sigmoid(input Float64Data) ([]float64, error) {}
|
||||
func SoftMax(input Float64Data) ([]float64, error) {}
|
||||
func StableSample(input Float64Data, takenum int) ([]float64, error) {}
|
||||
@@ -130,10 +161,16 @@ func StandardDeviationSample(input Float64Data) (sdev float64, err error) {}
|
||||
func StdDevP(input Float64Data) (sdev float64, err error) {}
|
||||
func StdDevS(input Float64Data) (sdev float64, err error) {}
|
||||
func Sum(input Float64Data) (sum float64, err error) {}
|
||||
func TTest(data1, data2 Float64Data, populationMean float64) (t float64, pvalue float64, err error) {}
|
||||
func Trimean(input Float64Data) (float64, error) {}
|
||||
func TrimmedMean(input Float64Data, percent float64) (float64, error) {}
|
||||
func VarP(input Float64Data) (sdev float64, err error) {}
|
||||
func VarS(input Float64Data) (sdev float64, err error) {}
|
||||
func Variance(input Float64Data) (sdev float64, err error) {}
|
||||
func WeightedMean(data, weights Float64Data) (float64, error) {}
|
||||
func Winsorize(input Float64Data, percent float64) ([]float64, error) {}
|
||||
func ZScore(input Float64Data) ([]float64, error) {}
|
||||
func ZTest(data1, data2 Float64Data, populationMean, populationStdDev float64) (z float64, pvalue float64, err error) {}
|
||||
func ProbGeom(a int, b int, p float64) (prob float64, err error) {}
|
||||
func ExpGeom(p float64) (exp float64, err error) {}
|
||||
func VarGeom(p float64) (exp float64, err error) {}
|
||||
@@ -178,20 +215,21 @@ Pull request are always welcome no matter how big or small. I've included a [Mak
|
||||
|
||||
To make things as seamless as possible please also consider the following steps:
|
||||
|
||||
- Update `examples/main.go` with a simple example of the new feature
|
||||
- Update `examples/functions/main.go` with a simple example of the new feature
|
||||
- Update `README.md` documentation section with any new exported API
|
||||
- Keep 100% code coverage (you can check with `make coverage`)
|
||||
- Squash commits into single units of work with `git rebase -i new-feature`
|
||||
|
||||
## Releasing
|
||||
|
||||
Releases are automated with [GoReleaser](https://goreleaser.com/) via GitHub Actions. To create a new release, push a version tag:
|
||||
Releases are automated with [GoReleaser](https://goreleaser.com/) via GitHub Actions. To create a new release, run the release target with the next version tag:
|
||||
|
||||
```
|
||||
git tag v0.x.x
|
||||
git push origin v0.x.x
|
||||
make release TAG=v0.x.x
|
||||
```
|
||||
|
||||
This updates `CHANGELOG.md`, commits it, creates an annotated tag, and pushes the commit and tag to GitHub, where the release workflow publishes the release.
|
||||
|
||||
## MIT License
|
||||
|
||||
Copyright (c) 2014-2026 Montana Flynn (https://montanaflynn.com)
|
||||
@@ -200,7 +238,7 @@ Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||
|
||||
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 AUTHORpublicS 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.
|
||||
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.
|
||||
|
||||
[action-url]: https://github.com/montanaflynn/stats/actions
|
||||
[action-svg]: https://img.shields.io/github/actions/workflow/status/montanaflynn/stats/go.yml
|
||||
@@ -208,12 +246,6 @@ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLI
|
||||
[codecov-url]: https://app.codecov.io/gh/montanaflynn/stats
|
||||
[codecov-svg]: https://img.shields.io/codecov/c/github/montanaflynn/stats?token=wnw8dActnH
|
||||
|
||||
[goreport-url]: https://goreportcard.com/report/github.com/montanaflynn/stats
|
||||
[goreport-svg]: https://goreportcard.com/badge/github.com/montanaflynn/stats
|
||||
|
||||
[godoc-url]: https://godoc.org/github.com/montanaflynn/stats
|
||||
[godoc-svg]: https://godoc.org/github.com/montanaflynn/stats?status.svg
|
||||
|
||||
[pkggodev-url]: https://pkg.go.dev/github.com/montanaflynn/stats
|
||||
[pkggodev-svg]: https://gistcdn.githack.com/montanaflynn/b02f1d78d8c0de8435895d7e7cd0d473/raw/17f2a5a69f1323ecd42c00e0683655da96d9ecc8/badge.svg
|
||||
|
||||
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
package stats
|
||||
|
||||
// Clip clamps each value in the input slice into the
|
||||
// inclusive range between min and max.
|
||||
func Clip(input Float64Data, min, max float64) ([]float64, error) {
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if min > max {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
c := make([]float64, len(input))
|
||||
for i, v := range input {
|
||||
switch {
|
||||
case v < min:
|
||||
c[i] = min
|
||||
case v > max:
|
||||
c[i] = max
|
||||
default:
|
||||
c[i] = v
|
||||
}
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// Clip clamps each value in the input slice into the
|
||||
// inclusive range between min and max.
|
||||
func (f Float64Data) Clip(min, max float64) ([]float64, error) {
|
||||
return Clip(f, min, max)
|
||||
}
|
||||
+31
@@ -0,0 +1,31 @@
|
||||
package stats
|
||||
|
||||
import "math"
|
||||
|
||||
// CoefficientOfVariation finds the coefficient of variation of a slice
|
||||
// of floats, defined as the sample standard deviation divided by the
|
||||
// mean. This matches the behavior of Python's scipy.stats.variation
|
||||
// with ddof=1.
|
||||
//
|
||||
// The input must not be empty and its mean must not be zero.
|
||||
func CoefficientOfVariation(input Float64Data) (float64, error) {
|
||||
if input.Len() == 0 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
// Input is known to be non-empty so the mean and sample
|
||||
// standard deviation cannot return an error
|
||||
m, _ := Mean(input)
|
||||
if m == 0 {
|
||||
return math.NaN(), ErrZero
|
||||
}
|
||||
|
||||
sd, _ := StandardDeviationSample(input)
|
||||
|
||||
return sd / m, nil
|
||||
}
|
||||
|
||||
// CoefficientOfVariation finds the sample standard deviation divided by the mean
|
||||
func (f Float64Data) CoefficientOfVariation() (float64, error) {
|
||||
return CoefficientOfVariation(f)
|
||||
}
|
||||
+84
-15
@@ -2,6 +2,7 @@ package stats
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sort"
|
||||
)
|
||||
|
||||
// Correlation describes the degree of relationship between two sets of data
|
||||
@@ -34,27 +35,95 @@ func Pearson(data1, data2 Float64Data) (float64, error) {
|
||||
return Correlation(data1, data2)
|
||||
}
|
||||
|
||||
// Spearman calculates the Spearman rank correlation coefficient between two variables.
|
||||
// It works by ranking the data and then computing the Pearson correlation of the ranks.
|
||||
// This method handles tied values using fractional (average) ranking.
|
||||
func Spearman(data1, data2 Float64Data) (float64, error) {
|
||||
|
||||
l1 := data1.Len()
|
||||
l2 := data2.Len()
|
||||
|
||||
if l1 == 0 || l2 == 0 {
|
||||
return math.NaN(), EmptyInputErr
|
||||
}
|
||||
|
||||
if l1 != l2 {
|
||||
return math.NaN(), SizeErr
|
||||
}
|
||||
|
||||
ranks1 := rankData(data1)
|
||||
ranks2 := rankData(data2)
|
||||
|
||||
return Correlation(ranks1, ranks2)
|
||||
}
|
||||
|
||||
// rankData assigns fractional (average) ranks to the data values.
|
||||
// Tied values receive the average of the ranks they would have been assigned.
|
||||
func rankData(data Float64Data) Float64Data {
|
||||
n := len(data)
|
||||
|
||||
// Create index-value pairs and sort by value
|
||||
type indexedValue struct {
|
||||
index int
|
||||
value float64
|
||||
}
|
||||
|
||||
sorted := make([]indexedValue, n)
|
||||
for i, v := range data {
|
||||
sorted[i] = indexedValue{i, v}
|
||||
}
|
||||
|
||||
sort.SliceStable(sorted, func(i, j int) bool {
|
||||
return sorted[i].value < sorted[j].value
|
||||
})
|
||||
|
||||
ranks := make(Float64Data, n)
|
||||
|
||||
// Assign fractional ranks handling ties
|
||||
for i := 0; i < n; {
|
||||
j := i + 1
|
||||
for j < n && sorted[j].value == sorted[i].value {
|
||||
j++
|
||||
}
|
||||
|
||||
// Average rank for tied values (ranks are 1-based)
|
||||
avgRank := float64(i+j+1) / 2.0
|
||||
for k := i; k < j; k++ {
|
||||
ranks[sorted[k].index] = avgRank
|
||||
}
|
||||
|
||||
i = j
|
||||
}
|
||||
|
||||
return ranks
|
||||
}
|
||||
|
||||
// AutoCorrelation is the correlation of a signal with a delayed copy of itself as a function of delay
|
||||
func AutoCorrelation(data Float64Data, lags int) (float64, error) {
|
||||
if len(data) < 1 {
|
||||
return 0, EmptyInputErr
|
||||
}
|
||||
|
||||
mean, _ := Mean(data)
|
||||
|
||||
var result, q float64
|
||||
|
||||
for i := 0; i < lags; i++ {
|
||||
v := (data[0] - mean) * (data[0] - mean)
|
||||
for i := 1; i < len(data); i++ {
|
||||
delta0 := data[i-1] - mean
|
||||
delta1 := data[i] - mean
|
||||
q += (delta0*delta1 - q) / float64(i+1)
|
||||
v += (delta1*delta1 - v) / float64(i+1)
|
||||
}
|
||||
|
||||
result = q / v
|
||||
if lags < 0 || lags >= len(data) {
|
||||
return 0, BoundsErr
|
||||
}
|
||||
|
||||
return result, nil
|
||||
mean, _ := Mean(data)
|
||||
|
||||
var variance float64
|
||||
for _, v := range data {
|
||||
delta := v - mean
|
||||
variance += delta * delta
|
||||
}
|
||||
|
||||
if variance == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var covariance float64
|
||||
for i := lags; i < len(data); i++ {
|
||||
covariance += (data[i] - mean) * (data[i-lags] - mean)
|
||||
}
|
||||
|
||||
return covariance / variance, nil
|
||||
}
|
||||
|
||||
+76
@@ -0,0 +1,76 @@
|
||||
package stats
|
||||
|
||||
// CumulativeProduct calculates the cumulative product of the input slice
|
||||
func CumulativeProduct(input Float64Data) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return Float64Data{}, ErrEmptyInput
|
||||
}
|
||||
|
||||
cumProduct := make([]float64, input.Len())
|
||||
|
||||
for i, val := range input {
|
||||
if i == 0 {
|
||||
cumProduct[i] = val
|
||||
} else {
|
||||
cumProduct[i] = cumProduct[i-1] * val
|
||||
}
|
||||
}
|
||||
|
||||
return cumProduct, nil
|
||||
}
|
||||
|
||||
// CumulativeMax calculates the cumulative maximum of the input slice
|
||||
func CumulativeMax(input Float64Data) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return Float64Data{}, ErrEmptyInput
|
||||
}
|
||||
|
||||
cumMax := make([]float64, input.Len())
|
||||
|
||||
for i, val := range input {
|
||||
if i == 0 || val > cumMax[i-1] {
|
||||
cumMax[i] = val
|
||||
} else {
|
||||
cumMax[i] = cumMax[i-1]
|
||||
}
|
||||
}
|
||||
|
||||
return cumMax, nil
|
||||
}
|
||||
|
||||
// CumulativeMin calculates the cumulative minimum of the input slice
|
||||
func CumulativeMin(input Float64Data) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return Float64Data{}, ErrEmptyInput
|
||||
}
|
||||
|
||||
cumMin := make([]float64, input.Len())
|
||||
|
||||
for i, val := range input {
|
||||
if i == 0 || val < cumMin[i-1] {
|
||||
cumMin[i] = val
|
||||
} else {
|
||||
cumMin[i] = cumMin[i-1]
|
||||
}
|
||||
}
|
||||
|
||||
return cumMin, nil
|
||||
}
|
||||
|
||||
// CumulativeProduct calculates the cumulative product of the data
|
||||
func (f Float64Data) CumulativeProduct() ([]float64, error) {
|
||||
return CumulativeProduct(f)
|
||||
}
|
||||
|
||||
// CumulativeMax calculates the cumulative maximum of the data
|
||||
func (f Float64Data) CumulativeMax() ([]float64, error) {
|
||||
return CumulativeMax(f)
|
||||
}
|
||||
|
||||
// CumulativeMin calculates the cumulative minimum of the data
|
||||
func (f Float64Data) CumulativeMin() ([]float64, error) {
|
||||
return CumulativeMin(f)
|
||||
}
|
||||
+7
-2
@@ -36,10 +36,10 @@ func (f Float64Data) Median() (float64, error) { return Median(f) }
|
||||
// Mode returns the mode of the data
|
||||
func (f Float64Data) Mode() ([]float64, error) { return Mode(f) }
|
||||
|
||||
// GeometricMean returns the median of the data
|
||||
// GeometricMean returns the geometric mean of the data
|
||||
func (f Float64Data) GeometricMean() (float64, error) { return GeometricMean(f) }
|
||||
|
||||
// HarmonicMean returns the mode of the data
|
||||
// HarmonicMean returns the harmonic mean of the data
|
||||
func (f Float64Data) HarmonicMean() (float64, error) { return HarmonicMean(f) }
|
||||
|
||||
// MedianAbsoluteDeviation the median of the absolute deviations from the dataset median
|
||||
@@ -97,6 +97,11 @@ func (f Float64Data) Pearson(d Float64Data) (float64, error) {
|
||||
return Pearson(f, d)
|
||||
}
|
||||
|
||||
// Spearman calculates the Spearman rank correlation coefficient between two variables.
|
||||
func (f Float64Data) Spearman(d Float64Data) (float64, error) {
|
||||
return Spearman(f, d)
|
||||
}
|
||||
|
||||
// Quartile returns the three quartile points from a slice of data
|
||||
func (f Float64Data) Quartile(d Float64Data) (Quartiles, error) {
|
||||
return Quartile(d)
|
||||
|
||||
+4
@@ -9,6 +9,7 @@ type Description struct {
|
||||
Std float64
|
||||
Max float64
|
||||
Min float64
|
||||
Range float64
|
||||
DescriptionPercentiles []descriptionPercentile
|
||||
AllowedNaN bool
|
||||
}
|
||||
@@ -39,6 +40,7 @@ func DescribePercentileFunc(input Float64Data, allowNaN bool, percentiles *[]flo
|
||||
description.Std, _ = StandardDeviation(input)
|
||||
description.Max, _ = Max(input)
|
||||
description.Min, _ = Min(input)
|
||||
description.Range, _ = Range(input)
|
||||
description.Mean, _ = Mean(input)
|
||||
|
||||
if percentiles != nil {
|
||||
@@ -60,6 +62,7 @@ Represents the Description instance in a string format with specified number of
|
||||
std 0.82
|
||||
max 3.00
|
||||
min 1.00
|
||||
range 2.00
|
||||
25.00% NaN
|
||||
50.00% 1.50
|
||||
75.00% 2.50
|
||||
@@ -73,6 +76,7 @@ func (d *Description) String(decimals int) string {
|
||||
str += fmt.Sprintf("std\t%.*f\n", decimals, d.Std)
|
||||
str += fmt.Sprintf("max\t%.*f\n", decimals, d.Max)
|
||||
str += fmt.Sprintf("min\t%.*f\n", decimals, d.Min)
|
||||
str += fmt.Sprintf("range\t%.*f\n", decimals, d.Range)
|
||||
for _, percentile := range d.DescriptionPercentiles {
|
||||
str += fmt.Sprintf("%.2f%%\t%.*f\n", percentile.Percentile, decimals, percentile.Value)
|
||||
}
|
||||
|
||||
+48
@@ -0,0 +1,48 @@
|
||||
package stats
|
||||
|
||||
// Diff calculates the successive differences of the input slice,
|
||||
// returning input[i] - input[i-1] for each i in 1..len(input)-1.
|
||||
// The output has length len(input) - 1; a single-element input
|
||||
// returns an empty slice.
|
||||
func Diff(input Float64Data) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
diff := make([]float64, input.Len()-1)
|
||||
|
||||
for i := 1; i < input.Len(); i++ {
|
||||
diff[i-1] = input[i] - input[i-1]
|
||||
}
|
||||
|
||||
return diff, nil
|
||||
}
|
||||
|
||||
// PercentChange calculates the fractional change between successive
|
||||
// elements of the input slice, returning
|
||||
// (input[i] - input[i-1]) / input[i-1] for each i in 1..len(input)-1.
|
||||
// The output has length len(input) - 1; a single-element input
|
||||
// returns an empty slice. A zero denominator follows IEEE 754
|
||||
// semantics, yielding +Inf, -Inf, or NaN (for 0/0), matching the
|
||||
// behavior of pandas pct_change.
|
||||
func PercentChange(input Float64Data) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
change := make([]float64, input.Len()-1)
|
||||
|
||||
for i := 1; i < input.Len(); i++ {
|
||||
change[i-1] = (input[i] - input[i-1]) / input[i-1]
|
||||
}
|
||||
|
||||
return change, nil
|
||||
}
|
||||
|
||||
// Diff returns the successive differences of the data
|
||||
func (f Float64Data) Diff() ([]float64, error) { return Diff(f) }
|
||||
|
||||
// PercentChange returns the fractional change between successive elements of the data
|
||||
func (f Float64Data) PercentChange() ([]float64, error) { return PercentChange(f) }
|
||||
+6
-3
@@ -19,13 +19,16 @@ func Entropy(input Float64Data) (float64, error) {
|
||||
return -result, nil
|
||||
}
|
||||
|
||||
// normalize divides each value by the sum of all values,
|
||||
// leaving the input itself unchanged.
|
||||
func normalize(input Float64Data) (Float64Data, error) {
|
||||
sum, err := input.Sum()
|
||||
if err != nil {
|
||||
return Float64Data{}, err
|
||||
}
|
||||
for i := 0; i < input.Len(); i++ {
|
||||
input[i] = input[i] / sum
|
||||
c := copyslice(input)
|
||||
for i := 0; i < c.Len(); i++ {
|
||||
c[i] = c[i] / sum
|
||||
}
|
||||
return input, nil
|
||||
return c, nil
|
||||
}
|
||||
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
package stats
|
||||
|
||||
// EWMA calculates the exponentially weighted moving average of the input
|
||||
// with smoothing factor alpha. The first output equals the first input and
|
||||
// each subsequent entry is alpha*input[i] + (1-alpha)*output[i-1], so the
|
||||
// result has the same length as the input. The alpha must satisfy
|
||||
// 0 < alpha <= 1 or ErrBounds is returned. An empty input returns
|
||||
// ErrEmptyInput.
|
||||
func EWMA(input Float64Data, alpha float64) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if alpha <= 0 || alpha > 1 {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
output := make([]float64, input.Len())
|
||||
|
||||
output[0] = input[0]
|
||||
for i := 1; i < input.Len(); i++ {
|
||||
output[i] = alpha*input[i] + (1-alpha)*output[i-1]
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// EWMA returns the exponentially weighted moving average of the data with smoothing factor alpha
|
||||
func (f Float64Data) EWMA(alpha float64) ([]float64, error) {
|
||||
return EWMA(f, alpha)
|
||||
}
|
||||
+72
@@ -0,0 +1,72 @@
|
||||
package stats
|
||||
|
||||
// ArgMax finds the index of the highest number in a slice,
|
||||
// returning the first occurrence in the case of ties
|
||||
func ArgMax(input Float64Data) (int, error) {
|
||||
|
||||
// Return an error if there are no numbers
|
||||
if input.Len() == 0 {
|
||||
return -1, ErrEmptyInput
|
||||
}
|
||||
|
||||
// Track the index of the highest value seen so far
|
||||
index := 0
|
||||
|
||||
// Loop and replace with strictly higher values only,
|
||||
// which keeps the first occurrence on ties
|
||||
for i := 1; i < input.Len(); i++ {
|
||||
if input.Get(i) > input.Get(index) {
|
||||
index = i
|
||||
}
|
||||
}
|
||||
|
||||
return index, nil
|
||||
}
|
||||
|
||||
// ArgMin finds the index of the lowest number in a slice,
|
||||
// returning the first occurrence in the case of ties
|
||||
func ArgMin(input Float64Data) (int, error) {
|
||||
|
||||
// Return an error if there are no numbers
|
||||
if input.Len() == 0 {
|
||||
return -1, ErrEmptyInput
|
||||
}
|
||||
|
||||
// Track the index of the lowest value seen so far
|
||||
index := 0
|
||||
|
||||
// Loop and replace with strictly lower values only,
|
||||
// which keeps the first occurrence on ties
|
||||
for i := 1; i < input.Len(); i++ {
|
||||
if input.Get(i) < input.Get(index) {
|
||||
index = i
|
||||
}
|
||||
}
|
||||
|
||||
return index, nil
|
||||
}
|
||||
|
||||
// Range finds the difference between the highest and
|
||||
// lowest numbers in a slice
|
||||
func Range(input Float64Data) (float64, error) {
|
||||
|
||||
// Return an error if there are no numbers
|
||||
max, err := Max(input)
|
||||
if err != nil {
|
||||
return max, ErrEmptyInput
|
||||
}
|
||||
|
||||
// Disregard error, since Max would have already returned it
|
||||
min, _ := Min(input)
|
||||
|
||||
return max - min, nil
|
||||
}
|
||||
|
||||
// ArgMax returns the index of the highest number in the data
|
||||
func (f Float64Data) ArgMax() (int, error) { return ArgMax(f) }
|
||||
|
||||
// ArgMin returns the index of the lowest number in the data
|
||||
func (f Float64Data) ArgMin() (int, error) { return ArgMin(f) }
|
||||
|
||||
// Range returns the difference between the highest and lowest numbers in the data
|
||||
func (f Float64Data) Range() (float64, error) { return Range(f) }
|
||||
+5
-4
@@ -12,14 +12,15 @@ func ProbGeom(a int, b int, p float64) (prob float64, err error) {
|
||||
return math.NaN(), ErrBounds
|
||||
}
|
||||
|
||||
prob = 0
|
||||
q := 1 - p // probability of failure
|
||||
|
||||
for k := a + 1; k <= b; k++ {
|
||||
prob = prob + p*math.Pow(q, float64(k-1))
|
||||
if a == b {
|
||||
return p * math.Pow(q, float64(a-1)), nil
|
||||
}
|
||||
|
||||
return prob, nil
|
||||
// closed form of the sum p*q^(k-1) over k = a..b; expm1/log1p keep
|
||||
// 1-q^n accurate where direct subtraction would cancel
|
||||
return math.Pow(q, float64(a-1)) * -math.Expm1(float64(b-a+1)*math.Log1p(-p)), nil
|
||||
}
|
||||
|
||||
// ProbGeom generates the expectation or average number of trials
|
||||
|
||||
+57
@@ -0,0 +1,57 @@
|
||||
package stats
|
||||
|
||||
import "sort"
|
||||
|
||||
// Histogram calculates the histogram of a slice using the given
|
||||
// number of equal-width bins over [min, max], returning the count
|
||||
// of values in each bin along with the bins+1 bin edges. Each bin
|
||||
// is half-open [edges[i], edges[i+1]) except the last, which also
|
||||
// includes the maximum value.
|
||||
func Histogram(input Float64Data, bins int) ([]int, []float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if bins < 1 {
|
||||
return nil, nil, ErrBounds
|
||||
}
|
||||
|
||||
// Disregard errors, since input is not empty
|
||||
min, _ := Min(input)
|
||||
max, _ := Max(input)
|
||||
|
||||
// Expand the range by 0.5 on each side like
|
||||
// numpy when all of the values are equal
|
||||
if min == max {
|
||||
min -= 0.5
|
||||
max += 0.5
|
||||
}
|
||||
|
||||
// Build bins+1 equal-width bin edges from min to max
|
||||
width := (max - min) / float64(bins)
|
||||
edges := make([]float64, bins+1)
|
||||
for i := range edges {
|
||||
edges[i] = min + width*float64(i)
|
||||
}
|
||||
edges[bins] = max
|
||||
|
||||
// Count each value into the last bin whose left
|
||||
// edge does not exceed it, so the maximum value
|
||||
// lands in the final closed bin
|
||||
counts := make([]int, bins)
|
||||
for _, v := range input {
|
||||
i := sort.Search(len(edges), func(i int) bool { return edges[i] > v }) - 1
|
||||
if i == bins {
|
||||
i--
|
||||
}
|
||||
counts[i]++
|
||||
}
|
||||
|
||||
return counts, edges, nil
|
||||
}
|
||||
|
||||
// Histogram returns the counts and equal-width bin edges of the data
|
||||
func (f Float64Data) Histogram(bins int) ([]int, []float64, error) {
|
||||
return Histogram(f, bins)
|
||||
}
|
||||
+67
@@ -0,0 +1,67 @@
|
||||
package stats
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sort"
|
||||
)
|
||||
|
||||
// Interp calculates the one-dimensional piecewise-linear interpolant to a
|
||||
// function with given discrete data points (xp, fp), evaluated at each x.
|
||||
// Values of x below xp[0] return fp[0] and values above xp[len(xp)-1] return
|
||||
// fp[len(xp)-1], so no extrapolation is performed. Unlike numpy's interp,
|
||||
// which silently returns nonsense for unsorted coordinates, xp must be
|
||||
// strictly increasing or ErrBounds is returned. An empty x or xp returns
|
||||
// ErrEmptyInput and xp and fp of different lengths return ErrSize.
|
||||
// A NaN in xp returns ErrBounds and a NaN in x gives a NaN in the output.
|
||||
func Interp(x, xp, fp Float64Data) ([]float64, error) {
|
||||
|
||||
if x.Len() == 0 || xp.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if xp.Len() != fp.Len() {
|
||||
return nil, ErrSize
|
||||
}
|
||||
|
||||
// NaN loses every comparison, so the ordering check can't catch it
|
||||
for i := 0; i < xp.Len(); i++ {
|
||||
if math.IsNaN(xp[i]) || (i > 0 && xp[i] <= xp[i-1]) {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
}
|
||||
|
||||
output := make([]float64, x.Len())
|
||||
|
||||
for i, xv := range x {
|
||||
switch {
|
||||
case math.IsNaN(xv):
|
||||
output[i] = math.NaN()
|
||||
case xv <= xp[0]:
|
||||
output[i] = fp[0]
|
||||
case xv >= xp[xp.Len()-1]:
|
||||
output[i] = fp[fp.Len()-1]
|
||||
default:
|
||||
// The first index with xp[j] >= xv, which the clamping
|
||||
// above guarantees is within [1, len(xp)-1]
|
||||
j := sort.SearchFloat64s(xp, xv)
|
||||
if xv == xp[j] {
|
||||
// An exact knot hit returns fp[j] exactly, with no
|
||||
// interpolation arithmetic that could lose precision
|
||||
output[i] = fp[j]
|
||||
continue
|
||||
}
|
||||
t := (xv - xp[j-1]) / (xp[j] - xp[j-1])
|
||||
if math.IsInf(xp[j]-xp[j-1], 1) {
|
||||
// The knot spacing overflows float64, so halve each
|
||||
// term before dividing; halving is exact for the huge
|
||||
// values that make an overflowing difference possible
|
||||
t = (xv/2 - xp[j-1]/2) / (xp[j]/2 - xp[j-1]/2)
|
||||
}
|
||||
// The symmetric form stays finite for any finite fp where
|
||||
// fp[j]-fp[j-1] would overflow, since 0 < t < 1 here
|
||||
output[i] = fp[j-1]*(1-t) + fp[j]*t
|
||||
}
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
+57
@@ -0,0 +1,57 @@
|
||||
package stats
|
||||
|
||||
import "math"
|
||||
|
||||
// KendallTau calculates Kendall's tau-b rank correlation coefficient
|
||||
// between two variables. Tau-b corrects for ties, matching the values
|
||||
// produced by SciPy's kendalltau and pandas' corr(method="kendall").
|
||||
// Pairs are compared with a simple O(n^2) loop for clarity.
|
||||
func KendallTau(data1, data2 Float64Data) (float64, error) {
|
||||
|
||||
l1 := data1.Len()
|
||||
l2 := data2.Len()
|
||||
|
||||
if l1 == 0 || l2 == 0 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
if l1 != l2 {
|
||||
return math.NaN(), ErrSize
|
||||
}
|
||||
|
||||
// Count concordant and discordant pairs along with pairs
|
||||
// tied only in the first or only in the second variable.
|
||||
var concordant, discordant, tiedX, tiedY float64
|
||||
for i := 0; i < l1; i++ {
|
||||
for j := i + 1; j < l1; j++ {
|
||||
dx := data1[i] - data1[j]
|
||||
dy := data2[i] - data2[j]
|
||||
switch {
|
||||
case dx == 0 && dy == 0:
|
||||
// Pairs tied in both variables are excluded
|
||||
case dx == 0:
|
||||
tiedX++
|
||||
case dy == 0:
|
||||
tiedY++
|
||||
case (dx > 0) == (dy > 0):
|
||||
concordant++
|
||||
default:
|
||||
discordant++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// tau-b = (C - D) / sqrt((C + D + Tx) * (C + D + Ty))
|
||||
denominator := math.Sqrt((concordant + discordant + tiedX) * (concordant + discordant + tiedY))
|
||||
if denominator == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
return (concordant - discordant) / denominator, nil
|
||||
}
|
||||
|
||||
// KendallTau calculates Kendall's tau-b rank correlation coefficient
|
||||
// between two variables.
|
||||
func (f Float64Data) KendallTau(d Float64Data) (float64, error) {
|
||||
return KendallTau(f, d)
|
||||
}
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
package stats
|
||||
|
||||
import "math"
|
||||
|
||||
// Kurtosis computes the population excess kurtosis of the dataset
|
||||
func Kurtosis(input Float64Data) (float64, error) {
|
||||
return PopulationKurtosis(input)
|
||||
}
|
||||
|
||||
// PopulationKurtosis computes the population excess kurtosis (Fisher
|
||||
// definition) using the fourth central moment normalized by the squared
|
||||
// variance, so a normal distribution has a kurtosis of zero.
|
||||
func PopulationKurtosis(input Float64Data) (float64, error) {
|
||||
if input.Len() < 2 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
mean, _ := Mean(input)
|
||||
|
||||
// Compute sums of squared and fourth-power differences from the mean
|
||||
var sumOfSquares, sumOfFourths float64
|
||||
for _, v := range input {
|
||||
d := v - mean
|
||||
d2 := d * d
|
||||
sumOfSquares += d2
|
||||
sumOfFourths += d2 * d2
|
||||
}
|
||||
|
||||
if sumOfSquares == 0 {
|
||||
return math.NaN(), ErrZero
|
||||
}
|
||||
|
||||
n := float64(input.Len())
|
||||
variance := sumOfSquares / n
|
||||
|
||||
return (sumOfFourths/n)/(variance*variance) - 3.0, nil
|
||||
}
|
||||
|
||||
// SampleKurtosis computes the bias-corrected sample excess kurtosis,
|
||||
// matching pandas .kurt() and scipy.stats.kurtosis with bias=False.
|
||||
func SampleKurtosis(input Float64Data) (float64, error) {
|
||||
n := input.Len()
|
||||
if n < 4 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
g2, err := PopulationKurtosis(input)
|
||||
if err != nil {
|
||||
return math.NaN(), err
|
||||
}
|
||||
|
||||
// Bias-corrected: G2 = ((n+1)*g2 + 6) * (n-1) / ((n-2)*(n-3))
|
||||
nf := float64(n)
|
||||
return ((nf+1)*g2 + 6) * (nf - 1) / ((nf - 2) * (nf - 3)), nil
|
||||
}
|
||||
|
||||
// Kurtosis finds the population excess kurtosis of a slice of floats
|
||||
func (f Float64Data) Kurtosis() (float64, error) {
|
||||
return Kurtosis(f)
|
||||
}
|
||||
|
||||
// PopulationKurtosis finds the population excess kurtosis of a slice of floats
|
||||
func (f Float64Data) PopulationKurtosis() (float64, error) {
|
||||
return PopulationKurtosis(f)
|
||||
}
|
||||
|
||||
// SampleKurtosis finds the bias-corrected sample excess kurtosis of a slice of floats
|
||||
func (f Float64Data) SampleKurtosis() (float64, error) {
|
||||
return SampleKurtosis(f)
|
||||
}
|
||||
+8
-7
@@ -22,18 +22,19 @@ func GeometricMean(input Float64Data) (float64, error) {
|
||||
return math.NaN(), EmptyInputErr
|
||||
}
|
||||
|
||||
// Get the product of all the numbers
|
||||
// Get the sum of all the numbers natural logs and return an
|
||||
// error for values that cannot be included in geometric mean
|
||||
var p float64
|
||||
for _, n := range input {
|
||||
if p == 0 {
|
||||
p = n
|
||||
} else {
|
||||
p *= n
|
||||
if n < 0 {
|
||||
return math.NaN(), NegativeErr
|
||||
} else if n == 0 {
|
||||
return math.NaN(), ZeroErr
|
||||
}
|
||||
p += math.Log(n)
|
||||
}
|
||||
|
||||
// Calculate the geometric mean
|
||||
return math.Pow(p, 1/float64(l)), nil
|
||||
return math.Exp(p / float64(l)), nil
|
||||
}
|
||||
|
||||
// HarmonicMean gets the harmonic mean for a slice of numbers
|
||||
|
||||
+125
@@ -0,0 +1,125 @@
|
||||
package stats
|
||||
|
||||
// MovingMedian calculates the rolling median of the input over a trailing
|
||||
// window. Only fully-populated windows produce output, so the result has
|
||||
// len(input)-window+1 entries and entry i is the median of input[i : i+window].
|
||||
// The window must satisfy 1 <= window <= len(input) or ErrBounds is
|
||||
// returned. An empty input returns ErrEmptyInput.
|
||||
func MovingMedian(input Float64Data, window int) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if window < 1 || window > input.Len() {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
output := make([]float64, input.Len()-window+1)
|
||||
|
||||
for i := range output {
|
||||
// Median cannot fail here since every window is non-empty
|
||||
median, _ := Median(input[i : i+window])
|
||||
output[i] = median
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// MovingMin calculates the rolling minimum of the input over a trailing
|
||||
// window. Only fully-populated windows produce output, so the result has
|
||||
// len(input)-window+1 entries and entry i is the minimum of
|
||||
// input[i : i+window]. The window must satisfy 1 <= window <= len(input) or
|
||||
// ErrBounds is returned. An empty input returns ErrEmptyInput.
|
||||
func MovingMin(input Float64Data, window int) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if window < 1 || window > input.Len() {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
output := make([]float64, input.Len()-window+1)
|
||||
|
||||
for i := range output {
|
||||
// Min cannot fail here since every window is non-empty
|
||||
min, _ := Min(input[i : i+window])
|
||||
output[i] = min
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// MovingMax calculates the rolling maximum of the input over a trailing
|
||||
// window. Only fully-populated windows produce output, so the result has
|
||||
// len(input)-window+1 entries and entry i is the maximum of
|
||||
// input[i : i+window]. The window must satisfy 1 <= window <= len(input) or
|
||||
// ErrBounds is returned. An empty input returns ErrEmptyInput.
|
||||
func MovingMax(input Float64Data, window int) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if window < 1 || window > input.Len() {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
output := make([]float64, input.Len()-window+1)
|
||||
|
||||
for i := range output {
|
||||
// Max cannot fail here since every window is non-empty
|
||||
max, _ := Max(input[i : i+window])
|
||||
output[i] = max
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// MovingSum calculates the rolling sum of the input over a trailing
|
||||
// window. Only fully-populated windows produce output, so the result has
|
||||
// len(input)-window+1 entries and entry i is the sum of input[i : i+window].
|
||||
// The window must satisfy 1 <= window <= len(input) or ErrBounds is
|
||||
// returned. An empty input returns ErrEmptyInput.
|
||||
func MovingSum(input Float64Data, window int) ([]float64, error) {
|
||||
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
|
||||
if window < 1 || window > input.Len() {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
output := make([]float64, input.Len()-window+1)
|
||||
|
||||
for i := range output {
|
||||
// Sum cannot fail here since every window is non-empty
|
||||
sum, _ := Sum(input[i : i+window])
|
||||
output[i] = sum
|
||||
}
|
||||
|
||||
return output, nil
|
||||
}
|
||||
|
||||
// MovingMedian returns the rolling median of the data over a trailing window
|
||||
func (f Float64Data) MovingMedian(window int) ([]float64, error) {
|
||||
return MovingMedian(f, window)
|
||||
}
|
||||
|
||||
// MovingMin returns the rolling minimum of the data over a trailing window
|
||||
func (f Float64Data) MovingMin(window int) ([]float64, error) {
|
||||
return MovingMin(f, window)
|
||||
}
|
||||
|
||||
// MovingMax returns the rolling maximum of the data over a trailing window
|
||||
func (f Float64Data) MovingMax(window int) ([]float64, error) {
|
||||
return MovingMax(f, window)
|
||||
}
|
||||
|
||||
// MovingSum returns the rolling sum of the data over a trailing window
|
||||
func (f Float64Data) MovingSum(window int) ([]float64, error) {
|
||||
return MovingSum(f, window)
|
||||
}
|
||||
+54
-16
@@ -7,6 +7,12 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// NormSample generates random samples from a normal distribution
|
||||
// with the given mean (loc) and standard deviation (scale).
|
||||
func NormSample(loc float64, scale float64, size int) []float64 {
|
||||
return NormBoxMullerRvs(loc, scale, size)
|
||||
}
|
||||
|
||||
// NormPpfRvs generates random variates using the Point Percentile Function.
|
||||
// For more information please visit: https://demonstrations.wolfram.com/TheMethodOfInverseTransforms/
|
||||
func NormPpfRvs(loc float64, scale float64, size int) []float64 {
|
||||
@@ -45,27 +51,53 @@ func NormPdf(x float64, loc float64, scale float64) float64 {
|
||||
|
||||
// NormLogPdf is the log of the probability density function.
|
||||
func NormLogPdf(x float64, loc float64, scale float64) float64 {
|
||||
return math.Log((math.Pow(math.E, -(math.Pow(x-loc, 2))/(2*math.Pow(scale, 2)))) / (scale * math.Sqrt(2*math.Pi)))
|
||||
z := (x - loc) / scale
|
||||
return -0.5*z*z - math.Log(scale) - 0.5*math.Log(2*math.Pi)
|
||||
}
|
||||
|
||||
// NormCdf is the cumulative distribution function.
|
||||
func NormCdf(x float64, loc float64, scale float64) float64 {
|
||||
return 0.5 * (1 + math.Erf((x-loc)/(scale*math.Sqrt(2))))
|
||||
return 0.5 * math.Erfc(-(x-loc)/(scale*math.Sqrt2))
|
||||
}
|
||||
|
||||
// NormLogCdf is the log of the cumulative distribution function.
|
||||
func NormLogCdf(x float64, loc float64, scale float64) float64 {
|
||||
return math.Log(0.5 * (1 + math.Erf((x-loc)/(scale*math.Sqrt(2)))))
|
||||
z := (x - loc) / scale
|
||||
if z > 0 {
|
||||
return math.Log1p(-0.5 * math.Erfc(z/math.Sqrt2))
|
||||
}
|
||||
return normLogTail(-z)
|
||||
}
|
||||
|
||||
// NormSf is the survival function (also defined as 1 - cdf, but sf is sometimes more accurate).
|
||||
func NormSf(x float64, loc float64, scale float64) float64 {
|
||||
return 1 - 0.5*(1+math.Erf((x-loc)/(scale*math.Sqrt(2))))
|
||||
return 0.5 * math.Erfc((x-loc)/(scale*math.Sqrt2))
|
||||
}
|
||||
|
||||
// NormLogSf is the log of the survival function.
|
||||
func NormLogSf(x float64, loc float64, scale float64) float64 {
|
||||
return math.Log(1 - 0.5*(1+math.Erf((x-loc)/(scale*math.Sqrt(2)))))
|
||||
z := (x - loc) / scale
|
||||
if z < 0 {
|
||||
return math.Log1p(-0.5 * math.Erfc(-z/math.Sqrt2))
|
||||
}
|
||||
return normLogTail(z)
|
||||
}
|
||||
|
||||
// normSmallestNormal is the smallest positive normal float64; below it math.Erfc
|
||||
// keeps only a handful of significant bits.
|
||||
const normSmallestNormal = 2.2250738585072014e-308
|
||||
|
||||
// normLogTail returns log(sf(z)) for z >= 0.
|
||||
func normLogTail(z float64) float64 {
|
||||
if q := 0.5 * math.Erfc(z/math.Sqrt2); q >= normSmallestNormal {
|
||||
return math.Log(q)
|
||||
}
|
||||
// math.Erfc has decayed into the subnormals, so switch to the Mills ratio
|
||||
// expansion sf(z) = pdf(z)/z * (1 - 1/z^2 + 3/z^4 - 15/z^6 + 105/z^8 - ...),
|
||||
// whose first dropped term is below 1e-12 this far out.
|
||||
r := 1 / (z * z)
|
||||
return -0.5*z*z - math.Log(z) - 0.5*math.Log(2*math.Pi) +
|
||||
math.Log1p(r*(-1+r*(3+r*(-15+r*105))))
|
||||
}
|
||||
|
||||
// NormPpf is the point percentile function.
|
||||
@@ -126,7 +158,15 @@ func NormPpf(p float64, loc float64, scale float64) (x float64) {
|
||||
(((((b1*r+b2)*r+b3)*r+b4)*r+b5)*r + 1)
|
||||
}
|
||||
|
||||
e := 0.5*math.Erfc(-x/math.Sqrt2) - p
|
||||
// Halley correction on cdf(x)-p. Above the median cdf(x) and p have both
|
||||
// already rounded to 1, so the difference is taken between the survival
|
||||
// functions instead; 1-p is exact for p >= 0.5.
|
||||
var e float64
|
||||
if p > 0.5 {
|
||||
e = (1 - p) - 0.5*math.Erfc(x/math.Sqrt2)
|
||||
} else {
|
||||
e = 0.5*math.Erfc(-x/math.Sqrt2) - p
|
||||
}
|
||||
u := e * math.Sqrt(2*math.Pi) * math.Exp(x*x/2)
|
||||
x = x - u/(1+x*u/2)
|
||||
|
||||
@@ -134,11 +174,10 @@ func NormPpf(p float64, loc float64, scale float64) (x float64) {
|
||||
}
|
||||
|
||||
// NormIsf is the inverse survival function (inverse of sf).
|
||||
func NormIsf(p float64, loc float64, scale float64) (x float64) {
|
||||
if -NormPpf(p, loc, scale) == 0 {
|
||||
return 0
|
||||
}
|
||||
return -NormPpf(p, loc, scale)
|
||||
func NormIsf(p float64, loc float64, scale float64) float64 {
|
||||
// isf(p) == ppf(1-p), reached by reflecting the standard normal so that
|
||||
// loc stays out of the negation and 1-p is never formed.
|
||||
return loc - scale*NormPpf(p, 0, 1)
|
||||
}
|
||||
|
||||
// NormMoment approximates the non-central (raw) moment of order n.
|
||||
@@ -219,11 +258,10 @@ func NormStd(loc float64, scale float64) float64 {
|
||||
|
||||
// NormInterval finds endpoints of the range that contains alpha percent of the distribution.
|
||||
func NormInterval(alpha float64, loc float64, scale float64) [2]float64 {
|
||||
q1 := (1.0 - alpha) / 2
|
||||
q2 := (1.0 + alpha) / 2
|
||||
a := NormPpf(q1, loc, scale)
|
||||
b := NormPpf(q2, loc, scale)
|
||||
return [2]float64{a, b}
|
||||
// Derive both endpoints from the lower tail: (1+alpha)/2 rounds to 1 once
|
||||
// alpha is within an ulp of it, which sends the upper endpoint to +Inf.
|
||||
z := NormPpf((1.0-alpha)/2, 0, 1)
|
||||
return [2]float64{loc + scale*z, loc - scale*z}
|
||||
}
|
||||
|
||||
// factorial is the naive factorial algorithm.
|
||||
|
||||
+3
-3
@@ -27,7 +27,7 @@ func Percentile(input Float64Data, percent float64) (percentile float64, err err
|
||||
return input[0], nil
|
||||
}
|
||||
|
||||
if percent <= 0 || percent > 100 {
|
||||
if math.IsNaN(percent) || percent <= 0 || percent > 100 {
|
||||
return math.NaN(), BoundsErr
|
||||
}
|
||||
|
||||
@@ -62,8 +62,8 @@ func PercentileNearestRank(input Float64Data, percent float64) (percentile float
|
||||
return math.NaN(), EmptyInputErr
|
||||
}
|
||||
|
||||
// Return error for less than 0 or greater than 100 percentages
|
||||
if percent < 0 || percent > 100 {
|
||||
// Return error for NaN, less than 0, or greater than 100 percentages
|
||||
if math.IsNaN(percent) || percent < 0 || percent > 100 {
|
||||
return math.NaN(), BoundsErr
|
||||
}
|
||||
|
||||
|
||||
+31
@@ -0,0 +1,31 @@
|
||||
package stats
|
||||
|
||||
import "math"
|
||||
|
||||
// PercentileOfScore calculates the percentile rank of a score
|
||||
// relative to a slice of floats, defined as the percentage of
|
||||
// values strictly below the score plus half the percentage of
|
||||
// values equal to the score. The result is between 0 and 100.
|
||||
// This matches the behavior of Python's
|
||||
// scipy.stats.percentileofscore with kind="mean".
|
||||
func PercentileOfScore(input Float64Data, score float64) (float64, error) {
|
||||
if input.Len() == 0 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
var below, equal float64
|
||||
for _, v := range input {
|
||||
if v < score {
|
||||
below++
|
||||
} else if v == score {
|
||||
equal++
|
||||
}
|
||||
}
|
||||
|
||||
return 100 * (below + 0.5*equal) / float64(input.Len()), nil
|
||||
}
|
||||
|
||||
// PercentileOfScore calculates the percentile rank of a score relative to the data
|
||||
func (f Float64Data) PercentileOfScore(score float64) (float64, error) {
|
||||
return PercentileOfScore(f, score)
|
||||
}
|
||||
+69
@@ -0,0 +1,69 @@
|
||||
package stats
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sort"
|
||||
)
|
||||
|
||||
// PercentileWeighted finds the weighted percentile of a slice of floats
|
||||
// using the weighted empirical CDF (inverse CDF / nearest-rank method).
|
||||
//
|
||||
// For a given percent p, it returns the smallest data value x such that
|
||||
// the cumulative weight of all values <= x is at least p% of the total
|
||||
// weight. This matches the behavior of Python's statsmodels
|
||||
// DescrStatsW.quantile.
|
||||
//
|
||||
// The data and weights slices must be the same length. Weights must be
|
||||
// non-negative and at least one weight must be positive. The percent
|
||||
// parameter must be between 0 and 100 (exclusive).
|
||||
func PercentileWeighted(data, weights Float64Data, percent float64) (percentile float64, err error) {
|
||||
l := data.Len()
|
||||
if l == 0 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
if weights.Len() != l {
|
||||
return math.NaN(), ErrSize
|
||||
}
|
||||
|
||||
if math.IsNaN(percent) || percent <= 0 || percent > 100 {
|
||||
return math.NaN(), ErrBounds
|
||||
}
|
||||
|
||||
// Build sorted pairs by data value
|
||||
type pair struct {
|
||||
value float64
|
||||
weight float64
|
||||
}
|
||||
pairs := make([]pair, l)
|
||||
totalWeight := 0.0
|
||||
for i := 0; i < l; i++ {
|
||||
if weights[i] < 0 {
|
||||
return math.NaN(), ErrNegative
|
||||
}
|
||||
pairs[i] = pair{data[i], weights[i]}
|
||||
totalWeight += weights[i]
|
||||
}
|
||||
|
||||
if totalWeight == 0 {
|
||||
return math.NaN(), ErrBounds
|
||||
}
|
||||
|
||||
sort.Slice(pairs, func(i, j int) bool {
|
||||
return pairs[i].value < pairs[j].value
|
||||
})
|
||||
|
||||
// Find the smallest value where cumulative weight >= target
|
||||
target := (percent / 100) * totalWeight
|
||||
cumWeight := 0.0
|
||||
result := pairs[l-1].value
|
||||
for i := 0; i < l; i++ {
|
||||
cumWeight += pairs[i].weight
|
||||
if cumWeight >= target {
|
||||
result = pairs[i].value
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
+26
@@ -0,0 +1,26 @@
|
||||
package stats
|
||||
|
||||
import "math"
|
||||
|
||||
// Product calculates the product of a slice of floats by
|
||||
// multiplying the values from left to right. It is the scalar
|
||||
// counterpart of CumulativeProduct. Large inputs can overflow
|
||||
// to Inf; use GeometricMean for an overflow-safe summary of
|
||||
// multiplicative data.
|
||||
func Product(input Float64Data) (float64, error) {
|
||||
if input.Len() == 0 {
|
||||
return math.NaN(), ErrEmptyInput
|
||||
}
|
||||
|
||||
product := 1.0
|
||||
for _, v := range input {
|
||||
product *= v
|
||||
}
|
||||
|
||||
return product, nil
|
||||
}
|
||||
|
||||
// Product calculates the product of the data
|
||||
func (f Float64Data) Product() (float64, error) {
|
||||
return Product(f)
|
||||
}
|
||||
+16
@@ -0,0 +1,16 @@
|
||||
package stats
|
||||
|
||||
// Rank assigns fractional (average) ranks to the input values.
|
||||
// Ranks are 1-based and tied values receive the average of the
|
||||
// ranks they would have been assigned.
|
||||
func Rank(input Float64Data) ([]float64, error) {
|
||||
if input.Len() == 0 {
|
||||
return nil, ErrEmptyInput
|
||||
}
|
||||
return rankData(input), nil
|
||||
}
|
||||
|
||||
// Rank assigns fractional (average) ranks to the input values
|
||||
func (f Float64Data) Rank() ([]float64, error) {
|
||||
return Rank(f)
|
||||
}
|
||||
+73
-44
@@ -10,101 +10,130 @@ type Coordinate struct {
|
||||
X, Y float64
|
||||
}
|
||||
|
||||
// LinearRegression finds the least squares linear regression on data series
|
||||
// LinearRegression finds the least squares linear regression on data series.
|
||||
// A series without at least two distinct X values returns ErrBounds.
|
||||
func LinearRegression(s Series) (regressions Series, err error) {
|
||||
|
||||
if len(s) == 0 {
|
||||
return nil, EmptyInputErr
|
||||
}
|
||||
|
||||
// Placeholder for the math to be done
|
||||
var sum [4]float64
|
||||
|
||||
// Loop over data keeping index in place
|
||||
i := 0
|
||||
for ; i < len(s); i++ {
|
||||
sum[0] += s[i].X
|
||||
sum[1] += s[i].Y
|
||||
sum[2] += s[i].X * s[i].X
|
||||
sum[3] += s[i].X * s[i].Y
|
||||
var sumX, sumY float64
|
||||
for _, coordinate := range s {
|
||||
sumX += coordinate.X
|
||||
sumY += coordinate.Y
|
||||
}
|
||||
meanX := sumX / float64(len(s))
|
||||
meanY := sumY / float64(len(s))
|
||||
|
||||
// Find gradient and intercept
|
||||
f := float64(i)
|
||||
gradient := (f*sum[3] - sum[0]*sum[1]) / (f*sum[2] - sum[0]*sum[0])
|
||||
intercept := (sum[1] / f) - (gradient * sum[0] / f)
|
||||
var covariance, variance float64
|
||||
for _, coordinate := range s {
|
||||
dx := coordinate.X - meanX
|
||||
covariance += dx * (coordinate.Y - meanY)
|
||||
variance += dx * dx
|
||||
}
|
||||
if variance == 0 {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
gradient := covariance / variance
|
||||
|
||||
// Create the new regression series
|
||||
for j := 0; j < len(s); j++ {
|
||||
regressions = append(regressions, Coordinate{
|
||||
X: s[j].X,
|
||||
Y: s[j].X*gradient + intercept,
|
||||
Y: meanY + gradient*(s[j].X-meanX),
|
||||
})
|
||||
}
|
||||
|
||||
return regressions, nil
|
||||
}
|
||||
|
||||
// ExponentialRegression returns an exponential regression on data series
|
||||
// ExponentialRegression returns an exponential regression on data series.
|
||||
// A non-positive Y value returns ErrYCoord, and a series without at least two
|
||||
// distinct X values returns ErrBounds.
|
||||
func ExponentialRegression(s Series) (regressions Series, err error) {
|
||||
|
||||
if len(s) == 0 {
|
||||
return nil, EmptyInputErr
|
||||
}
|
||||
|
||||
var sum [6]float64
|
||||
var sumY, sumDeltaXY, sumYLogY float64
|
||||
referenceX := s[0].X
|
||||
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i].Y < 0 {
|
||||
return nil, YCoordErr
|
||||
if s[i].Y <= 0 {
|
||||
return nil, ErrYCoord
|
||||
}
|
||||
sum[0] += s[i].X
|
||||
sum[1] += s[i].Y
|
||||
sum[2] += s[i].X * s[i].X * s[i].Y
|
||||
sum[3] += s[i].Y * math.Log(s[i].Y)
|
||||
sum[4] += s[i].X * s[i].Y * math.Log(s[i].Y)
|
||||
sum[5] += s[i].X * s[i].Y
|
||||
sumY += s[i].Y
|
||||
sumDeltaXY += (s[i].X - referenceX) * s[i].Y
|
||||
sumYLogY += s[i].Y * math.Log(s[i].Y)
|
||||
}
|
||||
meanDeltaX := sumDeltaXY / sumY
|
||||
meanLogY := sumYLogY / sumY
|
||||
|
||||
denominator := (sum[1]*sum[2] - sum[5]*sum[5])
|
||||
a := math.Pow(math.E, (sum[2]*sum[3]-sum[5]*sum[4])/denominator)
|
||||
b := (sum[1]*sum[4] - sum[5]*sum[3]) / denominator
|
||||
var covariance, variance float64
|
||||
for _, coordinate := range s {
|
||||
dx := coordinate.X - referenceX - meanDeltaX
|
||||
covariance += coordinate.Y * dx * (math.Log(coordinate.Y) - meanLogY)
|
||||
variance += coordinate.Y * dx * dx
|
||||
}
|
||||
if variance == 0 {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
b := covariance / variance
|
||||
|
||||
for j := 0; j < len(s); j++ {
|
||||
regressions = append(regressions, Coordinate{
|
||||
X: s[j].X,
|
||||
Y: a * math.Exp(b*s[j].X),
|
||||
Y: math.Exp(meanLogY + b*(s[j].X-referenceX-meanDeltaX)),
|
||||
})
|
||||
}
|
||||
|
||||
return regressions, nil
|
||||
}
|
||||
|
||||
// LogarithmicRegression returns an logarithmic regression on data series
|
||||
// LogarithmicRegression returns a logarithmic regression on data series.
|
||||
// A non-positive X value or a series without at least two distinct X values
|
||||
// returns ErrBounds.
|
||||
func LogarithmicRegression(s Series) (regressions Series, err error) {
|
||||
|
||||
if len(s) == 0 {
|
||||
return nil, EmptyInputErr
|
||||
}
|
||||
|
||||
var sum [4]float64
|
||||
|
||||
i := 0
|
||||
for ; i < len(s); i++ {
|
||||
sum[0] += math.Log(s[i].X)
|
||||
sum[1] += s[i].Y * math.Log(s[i].X)
|
||||
sum[2] += s[i].Y
|
||||
sum[3] += math.Pow(math.Log(s[i].X), 2)
|
||||
if s[0].X <= 0 {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
|
||||
f := float64(i)
|
||||
a := (f*sum[1] - sum[2]*sum[0]) / (f*sum[3] - sum[0]*sum[0])
|
||||
b := (sum[2] - a*sum[0]) / f
|
||||
logX := make([]float64, len(s))
|
||||
referenceLogX := math.Log(s[0].X)
|
||||
var sumDeltaLogX, sumY float64
|
||||
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i].X <= 0 {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
logX[i] = math.Log(s[i].X)
|
||||
sumDeltaLogX += logX[i] - referenceLogX
|
||||
sumY += s[i].Y
|
||||
}
|
||||
meanDeltaLogX := sumDeltaLogX / float64(len(s))
|
||||
meanY := sumY / float64(len(s))
|
||||
|
||||
var covariance, variance float64
|
||||
for i, coordinate := range s {
|
||||
dx := logX[i] - referenceLogX - meanDeltaLogX
|
||||
covariance += dx * (coordinate.Y - meanY)
|
||||
variance += dx * dx
|
||||
}
|
||||
if variance == 0 {
|
||||
return nil, ErrBounds
|
||||
}
|
||||
a := covariance / variance
|
||||
|
||||
for j := 0; j < len(s); j++ {
|
||||
regressions = append(regressions, Coordinate{
|
||||
X: s[j].X,
|
||||
Y: b + a*math.Log(s[j].X),
|
||||
Y: meanY + a*(logX[j]-referenceLogX-meanDeltaLogX),
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user