mirror of
https://github.com/yunionio/cloudpods.git
synced 2026-08-30 17:13:08 +08:00
update vendor
This commit is contained in:
@@ -69,6 +69,9 @@ require (
|
||||
github.com/jinzhu/now v1.0.1 // indirect
|
||||
github.com/koding/websocketproxy v0.0.0-20181220232114-7ed82d81a28c
|
||||
github.com/kr/pty v1.1.5
|
||||
github.com/lestrrat-go/jwx v1.0.2
|
||||
github.com/lestrrat/go-jwx v0.0.0-20180221005942-b7d4802280ae
|
||||
github.com/lestrrat/go-pdebug v0.0.0-20180220043741-569c97477ae8 // indirect
|
||||
github.com/lib/pq v1.2.0 // indirect
|
||||
github.com/libvirt/libvirt-go-xml v5.2.0+incompatible
|
||||
github.com/ma314smith/signedxml v0.0.0-20200410192636-c342a2d0ae60
|
||||
@@ -97,7 +100,7 @@ require (
|
||||
github.com/skip2/go-qrcode v0.0.0-20190110000554-dc11ecdae0a9
|
||||
github.com/smartystreets/goconvey v1.6.4
|
||||
github.com/spaolacci/murmur3 v1.1.0 // indirect
|
||||
github.com/stretchr/testify v1.4.0
|
||||
github.com/stretchr/testify v1.5.1
|
||||
github.com/tencentcloud/tencentcloud-sdk-go v3.0.135+incompatible
|
||||
github.com/tencentyun/cos-go-sdk-v5 v0.0.0-20191108095731-8ca4b370cde4
|
||||
github.com/tinylib/msgp v1.1.0 // indirect
|
||||
|
||||
@@ -501,6 +501,15 @@ github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
|
||||
github.com/kylelemons/godebug v0.0.0-20170820004349-d65d576e9348/go.mod h1:B69LEHPfb2qLo0BaaOLcbitczOKLWTsrBG9LczfCD4k=
|
||||
github.com/leodido/go-urn v1.1.0 h1:Sm1gr51B1kKyfD2BlRcLSiEkffoG96g6TPv6eRoEiB8=
|
||||
github.com/leodido/go-urn v1.1.0/go.mod h1:+cyI34gQWZcE1eQU7NVgKkkzdXDQHr1dBMtdAPozLkw=
|
||||
github.com/lestrrat-go/iter v0.0.0-20200422075355-fc1769541911 h1:FvnrqecqX4zT0wOIbYK1gNgTm0677INEWiFY8UEYggY=
|
||||
github.com/lestrrat-go/iter v0.0.0-20200422075355-fc1769541911/go.mod h1:zIdgO1mRKhn8l9vrZJZz9TUMMFbQbLeTsbqPDrJ/OJc=
|
||||
github.com/lestrrat-go/jwx v1.0.2 h1:FsbZg/v979RikHWhSu/7BRHh2Z1Z8byPleURRb1Y0XI=
|
||||
github.com/lestrrat-go/jwx v1.0.2/go.mod h1:TPF17WiSFegZo+c20fdpw49QD+/7n4/IsGvEmCSWwT0=
|
||||
github.com/lestrrat-go/pdebug v0.0.0-20200204225717-4d6bd78da58d/go.mod h1:B06CSso/AWxiPejj+fheUINGeBKeeEZNt8w+EoU7+L8=
|
||||
github.com/lestrrat/go-jwx v0.0.0-20180221005942-b7d4802280ae h1:XoMPFIGibcPKgLrgIxzif36Zs/2yOEeGYc/7nitjzNM=
|
||||
github.com/lestrrat/go-jwx v0.0.0-20180221005942-b7d4802280ae/go.mod h1:T+yHdCP6MJKtzoVQMHvVCeam5VFwX1+rWzn5zZgKYMI=
|
||||
github.com/lestrrat/go-pdebug v0.0.0-20180220043741-569c97477ae8 h1:ttJD8hTqvrPEUBoAG5hJKbDOJ84u7zmbnZsUL4V9430=
|
||||
github.com/lestrrat/go-pdebug v0.0.0-20180220043741-569c97477ae8/go.mod h1:VXFH11P7fHn2iPBsfSW1JacR59rttTcafJnwYcI/IdY=
|
||||
github.com/lib/pq v1.2.0 h1:LXpIM/LZ5xGFhOpXAQUIMM1HdyqzVYM13zNdjCEEcA0=
|
||||
github.com/lib/pq v1.2.0/go.mod h1:5WUZQaWbwv1U+lTReE5YruASi9Al49XbQIvNi/34Woo=
|
||||
github.com/libopenstorage/openstorage v1.0.0/go.mod h1:Sp1sIObHjat1BeXhfMqLZ14wnOzEhNx2YQedreMcUyc=
|
||||
@@ -728,6 +737,8 @@ github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXf
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/stretchr/testify v1.4.0 h1:2E4SXV/wtOkTonXsotYi4li6zVWxYlZuYNCXe9XRJyk=
|
||||
github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4=
|
||||
github.com/stretchr/testify v1.5.1 h1:nOGnQDM7FYENwehXlg/kFVnos3rEvtKTjRvOWSzb6H4=
|
||||
github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA=
|
||||
github.com/syncthing/syncthing v0.14.48-rc.4/go.mod h1:nw3siZwHPA6M8iSfjDCWQ402eqvEIasMQOE8nFOxy7M=
|
||||
github.com/syndtr/gocapability v0.0.0-20160928074757-e7cb7fa329f4/go.mod h1:hkRG7XYTFWNJGYcbNJQlaLq0fg1yr4J4t/NcTQtrfww=
|
||||
github.com/tencentcloud/tencentcloud-sdk-go v3.0.135+incompatible h1:QIMoFqKCmNp4HPLiTR+couZbHsIZfoOllncHYvtqse8=
|
||||
@@ -967,6 +978,7 @@ golang.org/x/tools v0.0.0-20191130070609-6e064ea0cf2d/go.mod h1:b+2E5dAYhXwXZwtn
|
||||
golang.org/x/tools v0.0.0-20191216173652-a0e659d51361/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||
golang.org/x/tools v0.0.0-20191227053925-7b8e75db28f4 h1:Toz2IK7k8rbltAXwNAxKcn9OzqyNfMUhUNjz3sL0NMk=
|
||||
golang.org/x/tools v0.0.0-20191227053925-7b8e75db28f4/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||
golang.org/x/tools v0.0.0-20200417140056-c07e33ef3290/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
||||
golang.org/x/tools v0.0.0-20200515220128-d3bf790afa53 h1:vmsb6v0zUdmUlXfwKaYrHPPRCV0lHq/IwNIf0ASGjyQ=
|
||||
golang.org/x/tools v0.0.0-20200515220128-d3bf790afa53/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE=
|
||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
|
||||
@@ -15,12 +15,12 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"crypto/rsa"
|
||||
"testing"
|
||||
|
||||
"github.com/lestrrat/go-jwx/jwa"
|
||||
"github.com/lestrrat/go-jwx/jwk"
|
||||
|
||||
"github.com/lestrrat/go-jwx/jwt"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jwk"
|
||||
"github.com/lestrrat-go/jwx/jwt"
|
||||
|
||||
"yunion.io/x/jsonutils"
|
||||
|
||||
@@ -77,7 +77,8 @@ func TestJWKVerify(t *testing.T) {
|
||||
for i := range keySet.Keys {
|
||||
key := keySet.Keys[i]
|
||||
if key.KeyUsage() == "sig" {
|
||||
oKey, err := key.Materialize()
|
||||
var oKey rsa.PublicKey
|
||||
err := key.Raw(&oKey)
|
||||
if err != nil {
|
||||
t.Fatalf("Meterialize fail %s", err)
|
||||
}
|
||||
|
||||
+21
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2020 lestrrat-go
|
||||
|
||||
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.
|
||||
+192
@@ -0,0 +1,192 @@
|
||||
package arrayiter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"sync"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func Iterate(ctx context.Context, a interface{}) (Iterator, error) {
|
||||
arv := reflect.ValueOf(a)
|
||||
|
||||
switch arv.Kind() {
|
||||
case reflect.Array, reflect.Slice:
|
||||
default:
|
||||
return nil, errors.Errorf(`argument must be an array/slice (%s)`, arv.Type())
|
||||
}
|
||||
|
||||
ch := make(chan *Pair)
|
||||
go func(ctx context.Context, ch chan *Pair, arv reflect.Value) {
|
||||
defer close(ch)
|
||||
|
||||
for i := 0; i < arv.Len(); i++ {
|
||||
value := arv.Index(i)
|
||||
pair := &Pair{
|
||||
Index: i,
|
||||
Value: value.Interface(),
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}(ctx, ch, arv)
|
||||
|
||||
return New(ch), nil
|
||||
}
|
||||
|
||||
// Source represents a array that knows how to create an iterator
|
||||
type Source interface {
|
||||
Iterate(context.Context) Iterator
|
||||
}
|
||||
|
||||
// Pair represents a single pair of key and value from a array
|
||||
type Pair struct {
|
||||
Index int
|
||||
Value interface{}
|
||||
}
|
||||
|
||||
// Iterator iterates through keys and values of a array
|
||||
type Iterator interface {
|
||||
Next(context.Context) bool
|
||||
Pair() *Pair
|
||||
}
|
||||
|
||||
type iter struct {
|
||||
ch chan *Pair
|
||||
mu sync.RWMutex
|
||||
next *Pair
|
||||
}
|
||||
|
||||
// Visitor represents an object that handles each pair in a array
|
||||
type Visitor interface {
|
||||
Visit(int, interface{}) error
|
||||
}
|
||||
|
||||
// VisitorFunc is a type of Visitor based on a function
|
||||
type VisitorFunc func(int, interface{}) error
|
||||
|
||||
func (fn VisitorFunc) Visit(s int, v interface{}) error {
|
||||
return fn(s, v)
|
||||
}
|
||||
|
||||
func New(ch chan *Pair) Iterator {
|
||||
return &iter{
|
||||
ch: ch,
|
||||
}
|
||||
}
|
||||
|
||||
// Next returns true if there are more items to read from the iterator
|
||||
func (i *iter) Next(ctx context.Context) bool {
|
||||
i.mu.RLock()
|
||||
if i.ch == nil {
|
||||
i.mu.RUnlock()
|
||||
return false
|
||||
}
|
||||
i.mu.RUnlock()
|
||||
|
||||
i.mu.Lock()
|
||||
defer i.mu.Unlock()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
i.ch = nil
|
||||
return false
|
||||
case v, ok := <-i.ch:
|
||||
if !ok {
|
||||
i.ch = nil
|
||||
return false
|
||||
}
|
||||
i.next = v
|
||||
return true
|
||||
}
|
||||
|
||||
return false // never reached
|
||||
}
|
||||
|
||||
// Pair returns the currently buffered Pair. Calling Next() will reset its value
|
||||
func (i *iter) Pair() *Pair {
|
||||
i.mu.RLock()
|
||||
defer i.mu.RUnlock()
|
||||
return i.next
|
||||
}
|
||||
|
||||
// Walk walks through each element in the array
|
||||
func Walk(ctx context.Context, s Source, v Visitor) error {
|
||||
for i := s.Iterate(ctx); i.Next(ctx); {
|
||||
pair := i.Pair()
|
||||
if err := v.Visit(pair.Index, pair.Value); err != nil {
|
||||
return errors.Wrapf(err, `failed to visit index %d`, pair.Index)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func AsArray(ctx context.Context, s interface{}, v interface{}) error {
|
||||
var iter Iterator
|
||||
switch reflect.ValueOf(s).Kind() {
|
||||
case reflect.Array, reflect.Slice:
|
||||
x, err := Iterate(ctx, s)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to iterate over array/slice type`)
|
||||
}
|
||||
iter = x
|
||||
default:
|
||||
ssrc, ok := s.(Source)
|
||||
if !ok {
|
||||
return errors.Errorf(`cannot iterate over %T: not a arrayiter.Source type`, s)
|
||||
}
|
||||
iter = ssrc.Iterate(ctx)
|
||||
}
|
||||
|
||||
dst := reflect.ValueOf(v)
|
||||
|
||||
// dst MUST be a pointer to a array type
|
||||
if kind := dst.Kind(); kind != reflect.Ptr {
|
||||
return errors.Errorf(`dst must be a pointer to a array (%s)`, dst.Type())
|
||||
}
|
||||
|
||||
dst = dst.Elem()
|
||||
switch dst.Kind() {
|
||||
case reflect.Array, reflect.Slice:
|
||||
default:
|
||||
return errors.Errorf(`dst must be a pointer to an array or slice (%s)`, dst.Type())
|
||||
}
|
||||
|
||||
var pairs []*Pair
|
||||
for iter.Next(ctx) {
|
||||
pair := iter.Pair()
|
||||
pairs = append(pairs, pair)
|
||||
}
|
||||
|
||||
switch dst.Kind() {
|
||||
case reflect.Array:
|
||||
if len(pairs) < dst.Len() {
|
||||
return errors.Errorf(`dst array does not have enough space for elements (%d, want %d)`, dst.Len(), len(pairs))
|
||||
}
|
||||
case reflect.Slice:
|
||||
if dst.IsNil() {
|
||||
dst.Set(reflect.MakeSlice(dst.Type(), len(pairs), len(pairs)))
|
||||
}
|
||||
}
|
||||
|
||||
// dst must be assignable
|
||||
if !dst.CanSet() {
|
||||
return errors.New(`dst is not writeable`)
|
||||
}
|
||||
|
||||
elemtyp := dst.Type().Elem()
|
||||
for _, pair := range pairs {
|
||||
rvvalue := reflect.ValueOf(pair.Value)
|
||||
|
||||
if !rvvalue.Type().AssignableTo(elemtyp) {
|
||||
return errors.Errorf(`cannot assign key of type %s to map key of type %s`, rvvalue.Type(), elemtyp)
|
||||
}
|
||||
|
||||
dst.Index(pair.Index).Set(rvvalue)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
+188
@@ -0,0 +1,188 @@
|
||||
package mapiter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"sync"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// Iterate creates an iterator from arbitrary map types. This is not
|
||||
// the most efficient tool, but it's the quickest way to create an
|
||||
// iterator for maps.
|
||||
// Also, note that you cannot make any assumptions on the order of
|
||||
// pairs being returned.
|
||||
func Iterate(ctx context.Context, m interface{}) (Iterator, error) {
|
||||
mrv := reflect.ValueOf(m)
|
||||
|
||||
if mrv.Kind() != reflect.Map {
|
||||
return nil, errors.Errorf(`argument must be a map (%s)`, mrv.Type())
|
||||
}
|
||||
|
||||
ch := make(chan *Pair)
|
||||
go func(ctx context.Context, ch chan *Pair, mrv reflect.Value) {
|
||||
defer close(ch)
|
||||
for _, key := range mrv.MapKeys() {
|
||||
value := mrv.MapIndex(key)
|
||||
pair := &Pair{
|
||||
Key: key.Interface(),
|
||||
Value: value.Interface(),
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}(ctx, ch, mrv)
|
||||
|
||||
return New(ch), nil
|
||||
}
|
||||
|
||||
// Source represents a map that knows how to create an iterator
|
||||
type Source interface {
|
||||
Iterate(context.Context) Iterator
|
||||
}
|
||||
|
||||
// Pair represents a single pair of key and value from a map
|
||||
type Pair struct {
|
||||
Key interface{}
|
||||
Value interface{}
|
||||
}
|
||||
|
||||
// Iterator iterates through keys and values of a map
|
||||
type Iterator interface {
|
||||
Next(context.Context) bool
|
||||
Pair() *Pair
|
||||
}
|
||||
|
||||
type iter struct {
|
||||
ch chan *Pair
|
||||
mu sync.RWMutex
|
||||
next *Pair
|
||||
}
|
||||
|
||||
// Visitor represents an object that handles each pair in a map
|
||||
type Visitor interface {
|
||||
Visit(interface{}, interface{}) error
|
||||
}
|
||||
|
||||
// VisitorFunc is a type of Visitor based on a function
|
||||
type VisitorFunc func(interface{}, interface{}) error
|
||||
|
||||
func (fn VisitorFunc) Visit(s interface{}, v interface{}) error {
|
||||
return fn(s, v)
|
||||
}
|
||||
|
||||
func New(ch chan *Pair) Iterator {
|
||||
return &iter{
|
||||
ch: ch,
|
||||
}
|
||||
}
|
||||
|
||||
// Next returns true if there are more items to read from the iterator
|
||||
func (i *iter) Next(ctx context.Context) bool {
|
||||
i.mu.RLock()
|
||||
if i.ch == nil {
|
||||
i.mu.RUnlock()
|
||||
return false
|
||||
}
|
||||
i.mu.RUnlock()
|
||||
|
||||
i.mu.Lock()
|
||||
defer i.mu.Unlock()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
i.ch = nil
|
||||
return false
|
||||
case v, ok := <-i.ch:
|
||||
if !ok {
|
||||
i.ch = nil
|
||||
return false
|
||||
}
|
||||
i.next = v
|
||||
return true
|
||||
}
|
||||
|
||||
return false // never reached
|
||||
}
|
||||
|
||||
// Pair returns the currently buffered Pair. Calling Next() will reset its value
|
||||
func (i *iter) Pair() *Pair {
|
||||
i.mu.RLock()
|
||||
defer i.mu.RUnlock()
|
||||
return i.next
|
||||
}
|
||||
|
||||
// Walk walks through each element in the map
|
||||
func Walk(ctx context.Context, s Source, v Visitor) error {
|
||||
for i := s.Iterate(ctx); i.Next(ctx); {
|
||||
pair := i.Pair()
|
||||
if err := v.Visit(pair.Key, pair.Value); err != nil {
|
||||
return errors.Wrapf(err, `failed to visit key %s`, pair.Key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AsMap returns the values obtained from the source as a map
|
||||
func AsMap(ctx context.Context, s interface{}, v interface{}) error {
|
||||
var iter Iterator
|
||||
switch reflect.ValueOf(s).Kind() {
|
||||
case reflect.Map:
|
||||
x, err := Iterate(ctx, s)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to iterate over map type`)
|
||||
}
|
||||
iter = x
|
||||
default:
|
||||
ssrc, ok := s.(Source)
|
||||
if !ok {
|
||||
return errors.Errorf(`cannot iterate over %T: not a mapiter.Source type`, s)
|
||||
}
|
||||
iter = ssrc.Iterate(ctx)
|
||||
}
|
||||
|
||||
dst := reflect.ValueOf(v)
|
||||
|
||||
// dst MUST be a pointer to a map type
|
||||
if kind := dst.Kind(); kind != reflect.Ptr {
|
||||
return errors.Errorf(`dst must be a pointer to a map (%s)`, dst.Type())
|
||||
}
|
||||
|
||||
dst = dst.Elem()
|
||||
if dst.Kind() != reflect.Map {
|
||||
return errors.Errorf(`dst must be a pointer to a map (%s)`, dst.Type())
|
||||
}
|
||||
|
||||
if dst.IsNil() {
|
||||
dst.Set(reflect.MakeMap(dst.Type()))
|
||||
}
|
||||
|
||||
// dst must be assignable
|
||||
if !dst.CanSet() {
|
||||
return errors.New(`dst is not writeable`)
|
||||
}
|
||||
|
||||
keytyp := dst.Type().Key()
|
||||
valtyp := dst.Type().Elem()
|
||||
|
||||
for iter.Next(ctx) {
|
||||
pair := iter.Pair()
|
||||
|
||||
rvkey := reflect.ValueOf(pair.Key)
|
||||
rvvalue := reflect.ValueOf(pair.Value)
|
||||
|
||||
if !rvkey.Type().AssignableTo(keytyp) {
|
||||
return errors.Errorf(`cannot assign key of type %s to map key of type %s`, rvkey.Type(), keytyp)
|
||||
}
|
||||
if !rvvalue.Type().AssignableTo(valtyp) {
|
||||
return errors.Errorf(`cannot assign value of type %s to map value of type %s`, rvvalue.Type(), valtyp)
|
||||
}
|
||||
|
||||
dst.SetMapIndex(rvkey, rvvalue)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
The MIT License (MIT)
|
||||
|
||||
Copyright (c) 2015 lestrrat
|
||||
|
||||
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.
|
||||
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
package base64
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func EncodeToStringStd(src []byte) string {
|
||||
return base64.RawStdEncoding.EncodeToString(src)
|
||||
}
|
||||
|
||||
func EncodeToString(src []byte) string {
|
||||
return base64.RawURLEncoding.EncodeToString(src)
|
||||
}
|
||||
|
||||
func EncodeUint64ToString(v uint64) string {
|
||||
data := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(data, v)
|
||||
|
||||
i := 0
|
||||
for ; i < len(data); i++ {
|
||||
if data[i] != 0x0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return EncodeToString(data[i:])
|
||||
}
|
||||
|
||||
func DecodeString(src string) ([]byte, error) {
|
||||
var isRaw = !strings.HasSuffix(src, "=")
|
||||
if strings.ContainsAny(src, "+/") {
|
||||
if isRaw {
|
||||
return base64.RawStdEncoding.DecodeString(src)
|
||||
}
|
||||
return base64.StdEncoding.DecodeString(src)
|
||||
}
|
||||
|
||||
if isRaw {
|
||||
return base64.RawURLEncoding.DecodeString(src)
|
||||
}
|
||||
return base64.URLEncoding.DecodeString(src)
|
||||
}
|
||||
+34
@@ -0,0 +1,34 @@
|
||||
package iter
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
)
|
||||
|
||||
// MapVisitor is a specialized visitor for our purposes.
|
||||
// Whereas mapiter.Visitor supports any type of key, this
|
||||
// visitor assumes the key is a string
|
||||
type MapVisitor interface {
|
||||
Visit(string, interface{}) error
|
||||
}
|
||||
|
||||
type MapVisitorFunc func(string, interface{}) error
|
||||
|
||||
func (fn MapVisitorFunc) Visit(s string, v interface{}) error {
|
||||
return fn(s, v)
|
||||
}
|
||||
|
||||
func WalkMap(ctx context.Context, src mapiter.Source, visitor MapVisitor) error {
|
||||
return mapiter.Walk(ctx, src, mapiter.VisitorFunc(func(k, v interface{}) error {
|
||||
return visitor.Visit(k.(string), v)
|
||||
}))
|
||||
}
|
||||
|
||||
func AsMap(ctx context.Context, src mapiter.Source) (map[string]interface{}, error) {
|
||||
var m map[string]interface{}
|
||||
if err := mapiter.AsMap(ctx, src, &m); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
package option
|
||||
|
||||
type Interface interface {
|
||||
Name() string
|
||||
Value() interface{}
|
||||
}
|
||||
|
||||
type Option struct {
|
||||
name string
|
||||
value interface{}
|
||||
}
|
||||
|
||||
func New(name string, value interface{}) *Option {
|
||||
return &Option{
|
||||
name: name,
|
||||
value: value,
|
||||
}
|
||||
}
|
||||
|
||||
func (o *Option) Name() string {
|
||||
return o.name
|
||||
}
|
||||
func (o *Option) Value() interface{} {
|
||||
return o.value
|
||||
}
|
||||
+40
@@ -0,0 +1,40 @@
|
||||
package pool
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"math/big"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var bytesBufferPool = sync.Pool{
|
||||
New: allocBytesBuffer,
|
||||
}
|
||||
|
||||
func allocBytesBuffer() interface{} {
|
||||
return &bytes.Buffer{}
|
||||
}
|
||||
|
||||
func GetBytesBuffer() *bytes.Buffer {
|
||||
return bytesBufferPool.Get().(*bytes.Buffer)
|
||||
}
|
||||
|
||||
func ReleaseBytesBuffer(b *bytes.Buffer) {
|
||||
b.Reset()
|
||||
bytesBufferPool.Put(b)
|
||||
}
|
||||
|
||||
var bigIntPool = sync.Pool{
|
||||
New: allocBigInt,
|
||||
}
|
||||
|
||||
func allocBigInt() interface{} {
|
||||
return &big.Int{}
|
||||
}
|
||||
|
||||
func GetBigInt() *big.Int {
|
||||
return bigIntPool.Get().(*big.Int)
|
||||
}
|
||||
|
||||
func ReleaseBigInt(i *big.Int) {
|
||||
bigIntPool.Put(i.SetInt64(0))
|
||||
}
|
||||
+51
@@ -0,0 +1,51 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// CompressionAlgorithm represents the compression algorithms as described in https://tools.ietf.org/html/rfc7518#section-7.3
|
||||
type CompressionAlgorithm string
|
||||
|
||||
// Supported values for CompressionAlgorithm
|
||||
const (
|
||||
Deflate CompressionAlgorithm = "DEF" // DEFLATE (RFC 1951)
|
||||
NoCompress CompressionAlgorithm = "" // No compression
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *CompressionAlgorithm) Accept(value interface{}) error {
|
||||
var tmp CompressionAlgorithm
|
||||
if x, ok := value.(CompressionAlgorithm); ok {
|
||||
tmp = x
|
||||
} else {
|
||||
var s string
|
||||
switch x := value.(type) {
|
||||
case fmt.Stringer:
|
||||
s = x.String()
|
||||
case string:
|
||||
s = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.CompressionAlgorithm: %T`, value)
|
||||
}
|
||||
tmp = CompressionAlgorithm(s)
|
||||
}
|
||||
switch tmp {
|
||||
case Deflate, NoCompress:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.CompressionAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a CompressionAlgorithm
|
||||
func (v CompressionAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+55
@@ -0,0 +1,55 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// ContentEncryptionAlgorithm represents the various encryption algorithms as described in https://tools.ietf.org/html/rfc7518#section-5
|
||||
type ContentEncryptionAlgorithm string
|
||||
|
||||
// Supported values for ContentEncryptionAlgorithm
|
||||
const (
|
||||
A128CBC_HS256 ContentEncryptionAlgorithm = "A128CBC-HS256" // AES-CBC + HMAC-SHA256 (128)
|
||||
A128GCM ContentEncryptionAlgorithm = "A128GCM" // AES-GCM (128)
|
||||
A192CBC_HS384 ContentEncryptionAlgorithm = "A192CBC-HS384" // AES-CBC + HMAC-SHA384 (192)
|
||||
A192GCM ContentEncryptionAlgorithm = "A192GCM" // AES-GCM (192)
|
||||
A256CBC_HS512 ContentEncryptionAlgorithm = "A256CBC-HS512" // AES-CBC + HMAC-SHA512 (256)
|
||||
A256GCM ContentEncryptionAlgorithm = "A256GCM" // AES-GCM (256)
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *ContentEncryptionAlgorithm) Accept(value interface{}) error {
|
||||
var tmp ContentEncryptionAlgorithm
|
||||
if x, ok := value.(ContentEncryptionAlgorithm); ok {
|
||||
tmp = x
|
||||
} else {
|
||||
var s string
|
||||
switch x := value.(type) {
|
||||
case fmt.Stringer:
|
||||
s = x.String()
|
||||
case string:
|
||||
s = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.ContentEncryptionAlgorithm: %T`, value)
|
||||
}
|
||||
tmp = ContentEncryptionAlgorithm(s)
|
||||
}
|
||||
switch tmp {
|
||||
case A128CBC_HS256, A128GCM, A192CBC_HS384, A192GCM, A256CBC_HS512, A256GCM:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.ContentEncryptionAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a ContentEncryptionAlgorithm
|
||||
func (v ContentEncryptionAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// EllipticCurveAlgorithm represents the algorithms used for EC keys
|
||||
type EllipticCurveAlgorithm string
|
||||
|
||||
// Supported values for EllipticCurveAlgorithm
|
||||
const (
|
||||
InvalidEllipticCurve EllipticCurveAlgorithm = "P-invalid"
|
||||
P256 EllipticCurveAlgorithm = "P-256"
|
||||
P384 EllipticCurveAlgorithm = "P-384"
|
||||
P521 EllipticCurveAlgorithm = "P-521"
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *EllipticCurveAlgorithm) Accept(value interface{}) error {
|
||||
var tmp EllipticCurveAlgorithm
|
||||
if x, ok := value.(EllipticCurveAlgorithm); ok {
|
||||
tmp = x
|
||||
} else {
|
||||
var s string
|
||||
switch x := value.(type) {
|
||||
case fmt.Stringer:
|
||||
s = x.String()
|
||||
case string:
|
||||
s = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.EllipticCurveAlgorithm: %T`, value)
|
||||
}
|
||||
tmp = EllipticCurveAlgorithm(s)
|
||||
}
|
||||
switch tmp {
|
||||
case P256, P384, P521:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.EllipticCurveAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a EllipticCurveAlgorithm
|
||||
func (v EllipticCurveAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+17
@@ -0,0 +1,17 @@
|
||||
//go:generate go run internal/cmd/gentypes/main.go
|
||||
|
||||
// Package jwa defines the various algorithm described in https://tools.ietf.org/html/rfc7518
|
||||
package jwa
|
||||
|
||||
// Size returns the size of the EllipticCurveAlgorithm
|
||||
func (crv EllipticCurveAlgorithm) Size() int {
|
||||
switch crv {
|
||||
case P256:
|
||||
return 32
|
||||
case P384:
|
||||
return 48
|
||||
case P521:
|
||||
return 66
|
||||
}
|
||||
return 0
|
||||
}
|
||||
+66
@@ -0,0 +1,66 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// KeyEncryptionAlgorithm represents the various encryption algorithms as described in https://tools.ietf.org/html/rfc7518#section-4.1
|
||||
type KeyEncryptionAlgorithm string
|
||||
|
||||
// Supported values for KeyEncryptionAlgorithm
|
||||
const (
|
||||
A128GCMKW KeyEncryptionAlgorithm = "A128GCMKW" // AES-GCM key wrap (128)
|
||||
A128KW KeyEncryptionAlgorithm = "A128KW" // AES key wrap (128)
|
||||
A192GCMKW KeyEncryptionAlgorithm = "A192GCMKW" // AES-GCM key wrap (192)
|
||||
A192KW KeyEncryptionAlgorithm = "A192KW" // AES key wrap (192)
|
||||
A256GCMKW KeyEncryptionAlgorithm = "A256GCMKW" // AES-GCM key wrap (256)
|
||||
A256KW KeyEncryptionAlgorithm = "A256KW" // AES key wrap (256)
|
||||
DIRECT KeyEncryptionAlgorithm = "dir" // Direct encryption
|
||||
ECDH_ES KeyEncryptionAlgorithm = "ECDH-ES" // ECDH-ES
|
||||
ECDH_ES_A128KW KeyEncryptionAlgorithm = "ECDH-ES+A128KW" // ECDH-ES + AES key wrap (128)
|
||||
ECDH_ES_A192KW KeyEncryptionAlgorithm = "ECDH-ES+A192KW" // ECDH-ES + AES key wrap (192)
|
||||
ECDH_ES_A256KW KeyEncryptionAlgorithm = "ECDH-ES+A256KW" // ECDH-ES + AES key wrap (256)
|
||||
PBES2_HS256_A128KW KeyEncryptionAlgorithm = "PBES2-HS256+A128KW" // PBES2 + HMAC-SHA256 + AES key wrap (128)
|
||||
PBES2_HS384_A192KW KeyEncryptionAlgorithm = "PBES2-HS384+A192KW" // PBES2 + HMAC-SHA384 + AES key wrap (192)
|
||||
PBES2_HS512_A256KW KeyEncryptionAlgorithm = "PBES2-HS512+A256KW" // PBES2 + HMAC-SHA512 + AES key wrap (256)
|
||||
RSA1_5 KeyEncryptionAlgorithm = "RSA1_5" // RSA-PKCS1v1.5
|
||||
RSA_OAEP KeyEncryptionAlgorithm = "RSA-OAEP" // RSA-OAEP-SHA1
|
||||
RSA_OAEP_256 KeyEncryptionAlgorithm = "RSA-OAEP-256" // RSA-OAEP-SHA256
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *KeyEncryptionAlgorithm) Accept(value interface{}) error {
|
||||
var tmp KeyEncryptionAlgorithm
|
||||
if x, ok := value.(KeyEncryptionAlgorithm); ok {
|
||||
tmp = x
|
||||
} else {
|
||||
var s string
|
||||
switch x := value.(type) {
|
||||
case fmt.Stringer:
|
||||
s = x.String()
|
||||
case string:
|
||||
s = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.KeyEncryptionAlgorithm: %T`, value)
|
||||
}
|
||||
tmp = KeyEncryptionAlgorithm(s)
|
||||
}
|
||||
switch tmp {
|
||||
case A128GCMKW, A128KW, A192GCMKW, A192KW, A256GCMKW, A256KW, DIRECT, ECDH_ES, ECDH_ES_A128KW, ECDH_ES_A192KW, ECDH_ES_A256KW, PBES2_HS256_A128KW, PBES2_HS384_A192KW, PBES2_HS512_A256KW, RSA1_5, RSA_OAEP, RSA_OAEP_256:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.KeyEncryptionAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a KeyEncryptionAlgorithm
|
||||
func (v KeyEncryptionAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// KeyType represents the key type ("kty") that are supported
|
||||
type KeyType string
|
||||
|
||||
// Supported values for KeyType
|
||||
const (
|
||||
EC KeyType = "EC" // Elliptic Curve
|
||||
InvalidKeyType KeyType = "" // Invalid KeyType
|
||||
OctetSeq KeyType = "oct" // Octet sequence (used to represent symmetric keys)
|
||||
RSA KeyType = "RSA" // RSA
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *KeyType) Accept(value interface{}) error {
|
||||
var tmp KeyType
|
||||
if x, ok := value.(KeyType); ok {
|
||||
tmp = x
|
||||
} else {
|
||||
var s string
|
||||
switch x := value.(type) {
|
||||
case fmt.Stringer:
|
||||
s = x.String()
|
||||
case string:
|
||||
s = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.KeyType: %T`, value)
|
||||
}
|
||||
tmp = KeyType(s)
|
||||
}
|
||||
switch tmp {
|
||||
case EC, OctetSeq, RSA:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.KeyType value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a KeyType
|
||||
func (v KeyType) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+62
@@ -0,0 +1,62 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// SignatureAlgorithm represents the various signature algorithms as described in https://tools.ietf.org/html/rfc7518#section-3.1
|
||||
type SignatureAlgorithm string
|
||||
|
||||
// Supported values for SignatureAlgorithm
|
||||
const (
|
||||
ES256 SignatureAlgorithm = "ES256" // ECDSA using P-256 and SHA-256
|
||||
ES384 SignatureAlgorithm = "ES384" // ECDSA using P-384 and SHA-384
|
||||
ES512 SignatureAlgorithm = "ES512" // ECDSA using P-521 and SHA-512
|
||||
HS256 SignatureAlgorithm = "HS256" // HMAC using SHA-256
|
||||
HS384 SignatureAlgorithm = "HS384" // HMAC using SHA-384
|
||||
HS512 SignatureAlgorithm = "HS512" // HMAC using SHA-512
|
||||
NoSignature SignatureAlgorithm = "none"
|
||||
PS256 SignatureAlgorithm = "PS256" // RSASSA-PSS using SHA256 and MGF1-SHA256
|
||||
PS384 SignatureAlgorithm = "PS384" // RSASSA-PSS using SHA384 and MGF1-SHA384
|
||||
PS512 SignatureAlgorithm = "PS512" // RSASSA-PSS using SHA512 and MGF1-SHA512
|
||||
RS256 SignatureAlgorithm = "RS256" // RSASSA-PKCS-v1.5 using SHA-256
|
||||
RS384 SignatureAlgorithm = "RS384" // RSASSA-PKCS-v1.5 using SHA-384
|
||||
RS512 SignatureAlgorithm = "RS512" // RSASSA-PKCS-v1.5 using SHA-512
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *SignatureAlgorithm) Accept(value interface{}) error {
|
||||
var tmp SignatureAlgorithm
|
||||
if x, ok := value.(SignatureAlgorithm); ok {
|
||||
tmp = x
|
||||
} else {
|
||||
var s string
|
||||
switch x := value.(type) {
|
||||
case fmt.Stringer:
|
||||
s = x.String()
|
||||
case string:
|
||||
s = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.SignatureAlgorithm: %T`, value)
|
||||
}
|
||||
tmp = SignatureAlgorithm(s)
|
||||
}
|
||||
switch tmp {
|
||||
case ES256, ES384, ES512, HS256, HS384, HS512, NoSignature, PS256, PS384, PS512, RS256, RS384, RS512:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.SignatureAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a SignatureAlgorithm
|
||||
func (v SignatureAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+77
@@ -0,0 +1,77 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func (c CertificateChain) MarshalJSON() ([]byte, error) {
|
||||
certs := c.Get()
|
||||
encoded := make([]string, len(certs))
|
||||
for i := 0; i < len(certs); i++ {
|
||||
encoded[i] = base64.EncodeToStringStd(certs[i].Raw)
|
||||
}
|
||||
return json.Marshal(encoded)
|
||||
}
|
||||
|
||||
func (c *CertificateChain) UnmarshalJSON(buf []byte) error {
|
||||
var list []string
|
||||
if err := json.Unmarshal(buf, &list); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal JSON into []string`)
|
||||
}
|
||||
|
||||
var tmp CertificateChain
|
||||
if err := tmp.Accept(list); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*c = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c CertificateChain) Get() []*x509.Certificate {
|
||||
return c.certs
|
||||
}
|
||||
|
||||
func (c *CertificateChain) Accept(v interface{}) error {
|
||||
var list []string
|
||||
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
list = []string{x}
|
||||
case []interface{}:
|
||||
list = make([]string, len(x))
|
||||
for i, e := range x {
|
||||
if es, ok := e.(string); ok {
|
||||
list[i] = es
|
||||
continue
|
||||
}
|
||||
return errors.Errorf(`invalid list element type: expected string, got %T at element %d`, e, i)
|
||||
}
|
||||
case []string:
|
||||
list = x
|
||||
default:
|
||||
return errors.Errorf(`invalid tpe for CertificateChain: %T`, v)
|
||||
}
|
||||
|
||||
certs := make([]*x509.Certificate, len(list))
|
||||
for i, e := range list {
|
||||
buf, err := base64.DecodeString(e)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to base64 decode list element`)
|
||||
}
|
||||
cert, err := x509.ParseCertificate(buf)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to parse certificate`)
|
||||
}
|
||||
certs[i] = cert
|
||||
}
|
||||
|
||||
*c = CertificateChain{
|
||||
certs: certs,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"fmt"
|
||||
"math/big"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func NewECDSAPublicKey() ECDSAPublicKey {
|
||||
return newECDSAPublicKey()
|
||||
}
|
||||
|
||||
func newECDSAPublicKey() *ecdsaPublicKey {
|
||||
return &ecdsaPublicKey{
|
||||
privateParams: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
func NewECDSAPrivateKey() ECDSAPrivateKey {
|
||||
return newECDSAPrivateKey()
|
||||
}
|
||||
|
||||
func newECDSAPrivateKey() *ecdsaPrivateKey {
|
||||
return &ecdsaPrivateKey{
|
||||
privateParams: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (k *ecdsaPublicKey) FromRaw(rawKey *ecdsa.PublicKey) error {
|
||||
k.x = rawKey.X.Bytes()
|
||||
k.y = rawKey.Y.Bytes()
|
||||
switch rawKey.Curve {
|
||||
case elliptic.P256():
|
||||
if err := k.Set(ECDSACrvKey, jwa.P256); err != nil {
|
||||
return errors.Wrap(err, `failed to set header`)
|
||||
}
|
||||
case elliptic.P384():
|
||||
if err := k.Set(ECDSACrvKey, jwa.P384); err != nil {
|
||||
return errors.Wrap(err, `failed to set header`)
|
||||
}
|
||||
case elliptic.P521():
|
||||
if err := k.Set(ECDSACrvKey, jwa.P521); err != nil {
|
||||
return errors.Wrap(err, `failed to set header`)
|
||||
}
|
||||
default:
|
||||
return errors.Errorf(`invalid elliptic curve %s`, rawKey.Curve)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *ecdsaPrivateKey) FromRaw(rawKey *ecdsa.PrivateKey) error {
|
||||
k.x = rawKey.X.Bytes()
|
||||
k.y = rawKey.Y.Bytes()
|
||||
switch rawKey.Curve {
|
||||
case elliptic.P256():
|
||||
if err := k.Set(ECDSACrvKey, jwa.P256); err != nil {
|
||||
return errors.Wrap(err, "failed to write header")
|
||||
}
|
||||
case elliptic.P384():
|
||||
if err := k.Set(ECDSACrvKey, jwa.P384); err != nil {
|
||||
return errors.Wrap(err, "failed to write header")
|
||||
}
|
||||
case elliptic.P521():
|
||||
if err := k.Set(ECDSACrvKey, jwa.P521); err != nil {
|
||||
return errors.Wrap(err, "failed to write header")
|
||||
}
|
||||
default:
|
||||
return errors.Errorf(`invalid elliptic curve %s`, rawKey.Curve)
|
||||
}
|
||||
|
||||
k.d = rawKey.D.Bytes()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildECDSAPublicKey(alg jwa.EllipticCurveAlgorithm, xbuf, ybuf []byte) (*ecdsa.PublicKey, error) {
|
||||
var curve elliptic.Curve
|
||||
switch alg {
|
||||
case jwa.P256:
|
||||
curve = elliptic.P256()
|
||||
case jwa.P384:
|
||||
curve = elliptic.P384()
|
||||
case jwa.P521:
|
||||
curve = elliptic.P521()
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid curve algorithm %s`, alg)
|
||||
}
|
||||
|
||||
var x, y big.Int
|
||||
x.SetBytes(xbuf)
|
||||
y.SetBytes(ybuf)
|
||||
|
||||
return &ecdsa.PublicKey{Curve: curve, X: &x, Y: &y}, nil
|
||||
}
|
||||
|
||||
// Raw returns the EC-DSA public key represented by this JWK
|
||||
func (k *ecdsaPublicKey) Raw(v interface{}) error {
|
||||
pubk, err := buildECDSAPublicKey(k.Crv(), k.x, k.y)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to build public key`)
|
||||
}
|
||||
|
||||
return assignRawResult(v, pubk)
|
||||
}
|
||||
|
||||
func (k *ecdsaPrivateKey) Raw(v interface{}) error {
|
||||
pubk, err := buildECDSAPublicKey(k.Crv(), k.x, k.y)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to build public key`)
|
||||
}
|
||||
|
||||
var key ecdsa.PrivateKey
|
||||
var d big.Int
|
||||
d.SetBytes(k.d)
|
||||
key.D = &d
|
||||
key.PublicKey = *pubk
|
||||
|
||||
return assignRawResult(v, &key)
|
||||
}
|
||||
|
||||
func (k *ecdsaPrivateKey) PublicKey() (ECDSAPublicKey, error) {
|
||||
var privk ecdsa.PrivateKey
|
||||
if err := k.Raw(&privk); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize ECDSA private key`)
|
||||
}
|
||||
|
||||
newKey := NewECDSAPublicKey()
|
||||
if err := newKey.FromRaw(&privk.PublicKey); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to initialize ECDSAPublicKey`)
|
||||
}
|
||||
return newKey, nil
|
||||
}
|
||||
|
||||
func ecdsaThumbprint(hash crypto.Hash, crv, x, y string) []byte {
|
||||
h := hash.New()
|
||||
fmt.Fprint(h, `{"crv":"`)
|
||||
fmt.Fprint(h, crv)
|
||||
fmt.Fprint(h, `","kty":"EC","x":"`)
|
||||
fmt.Fprint(h, x)
|
||||
fmt.Fprint(h, `","y":"`)
|
||||
fmt.Fprint(h, y)
|
||||
fmt.Fprint(h, `"}`)
|
||||
return h.Sum(nil)
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (k ecdsaPublicKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
var key ecdsa.PublicKey
|
||||
if err := k.Raw(&key); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize ecdsa.PublicKey for thumbprint generation`)
|
||||
}
|
||||
return ecdsaThumbprint(
|
||||
hash,
|
||||
key.Curve.Params().Name,
|
||||
base64.EncodeToString(key.X.Bytes()),
|
||||
base64.EncodeToString(key.Y.Bytes()),
|
||||
), nil
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (k ecdsaPrivateKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
var key ecdsa.PrivateKey
|
||||
if err := k.Raw(&key); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize ecdsa.PrivateKey for thumbprint generation`)
|
||||
}
|
||||
return ecdsaThumbprint(
|
||||
hash,
|
||||
key.Curve.Params().Name,
|
||||
base64.EncodeToString(key.X.Bytes()),
|
||||
base64.EncodeToString(key.Y.Bytes()),
|
||||
), nil
|
||||
}
|
||||
+936
@@ -0,0 +1,936 @@
|
||||
// This file is auto-generated. DO NOT EDIT
|
||||
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/ecdsa"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
ECDSACrvKey = "crv"
|
||||
ECDSADKey = "d"
|
||||
ECDSAXKey = "x"
|
||||
ECDSAYKey = "y"
|
||||
)
|
||||
|
||||
type ECDSAPrivateKey interface {
|
||||
Key
|
||||
FromRaw(*ecdsa.PrivateKey) error
|
||||
Crv() jwa.EllipticCurveAlgorithm
|
||||
D() []byte
|
||||
X() []byte
|
||||
Y() []byte
|
||||
PublicKey() (ECDSAPublicKey, error)
|
||||
}
|
||||
|
||||
type ecdsaPrivateKey struct {
|
||||
algorithm *string // https://tools.ietf.org/html/rfc7517#section-4.4
|
||||
crv *jwa.EllipticCurveAlgorithm
|
||||
d []byte
|
||||
keyID *string // https://tools.ietf.org/html/rfc7515#section-4.1.4
|
||||
keyUsage *string // https://tools.ietf.org/html/rfc7517#section-4.2
|
||||
keyops *KeyOperationList // https://tools.ietf.org/html/rfc7517#section-4.3
|
||||
x []byte
|
||||
x509CertChain *CertificateChain // https://tools.ietf.org/html/rfc7515#section-4.1.6
|
||||
x509CertThumbprint *string // https://tools.ietf.org/html/rfc7515#section-4.1.7
|
||||
x509CertThumbprintS256 *string // https://tools.ietf.org/html/rfc7515#section-4.1.8
|
||||
x509URL *string // https://tools.ietf.org/html/rfc7515#section-4.1.5
|
||||
y []byte
|
||||
privateParams map[string]interface{}
|
||||
}
|
||||
|
||||
type ecdsaPrivateKeyMarshalProxy struct {
|
||||
XkeyType jwa.KeyType `json:"kty"`
|
||||
Xalgorithm *string `json:"alg,omitempty"`
|
||||
Xcrv *jwa.EllipticCurveAlgorithm `json:"crv,omitempty"`
|
||||
Xd *string `json:"d,omitempty"`
|
||||
XkeyID *string `json:"kid,omitempty"`
|
||||
XkeyUsage *string `json:"use,omitempty"`
|
||||
Xkeyops *KeyOperationList `json:"key_ops,omitempty"`
|
||||
Xx *string `json:"x,omitempty"`
|
||||
Xx509CertChain *CertificateChain `json:"x5c,omitempty"`
|
||||
Xx509CertThumbprint *string `json:"x5t,omitempty"`
|
||||
Xx509CertThumbprintS256 *string `json:"x5t#S256,omitempty"`
|
||||
Xx509URL *string `json:"x5u,omitempty"`
|
||||
Xy *string `json:"y,omitempty"`
|
||||
}
|
||||
|
||||
func (h ecdsaPrivateKey) KeyType() jwa.KeyType {
|
||||
return jwa.EC
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Algorithm() string {
|
||||
if h.algorithm != nil {
|
||||
return *(h.algorithm)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Crv() jwa.EllipticCurveAlgorithm {
|
||||
if h.crv != nil {
|
||||
return *(h.crv)
|
||||
}
|
||||
return jwa.InvalidEllipticCurve
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) D() []byte {
|
||||
return h.d
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) KeyID() string {
|
||||
if h.keyID != nil {
|
||||
return *(h.keyID)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) KeyUsage() string {
|
||||
if h.keyUsage != nil {
|
||||
return *(h.keyUsage)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) KeyOps() KeyOperationList {
|
||||
if h.keyops != nil {
|
||||
return *(h.keyops)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) X() []byte {
|
||||
return h.x
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) X509CertChain() []*x509.Certificate {
|
||||
if h.x509CertChain != nil {
|
||||
return h.x509CertChain.Get()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) X509CertThumbprint() string {
|
||||
if h.x509CertThumbprint != nil {
|
||||
return *(h.x509CertThumbprint)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) X509CertThumbprintS256() string {
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
return *(h.x509CertThumbprintS256)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) X509URL() string {
|
||||
if h.x509URL != nil {
|
||||
return *(h.x509URL)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Y() []byte {
|
||||
return h.y
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) iterate(ctx context.Context, ch chan *HeaderPair) {
|
||||
defer close(ch)
|
||||
|
||||
var pairs []*HeaderPair
|
||||
pairs = append(pairs, &HeaderPair{Key: "kty", Value: jwa.EC})
|
||||
if h.algorithm != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: AlgorithmKey, Value: *(h.algorithm)})
|
||||
}
|
||||
if h.crv != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSACrvKey, Value: *(h.crv)})
|
||||
}
|
||||
if h.d != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSADKey, Value: h.d})
|
||||
}
|
||||
if h.keyID != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyIDKey, Value: *(h.keyID)})
|
||||
}
|
||||
if h.keyUsage != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyUsageKey, Value: *(h.keyUsage)})
|
||||
}
|
||||
if h.keyops != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyOpsKey, Value: *(h.keyops)})
|
||||
}
|
||||
if h.x != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSAXKey, Value: h.x})
|
||||
}
|
||||
if h.x509CertChain != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertChainKey, Value: *(h.x509CertChain)})
|
||||
}
|
||||
if h.x509CertThumbprint != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintKey, Value: *(h.x509CertThumbprint)})
|
||||
}
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintS256Key, Value: *(h.x509CertThumbprintS256)})
|
||||
}
|
||||
if h.x509URL != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509URLKey, Value: *(h.x509URL)})
|
||||
}
|
||||
if h.y != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSAYKey, Value: h.y})
|
||||
}
|
||||
for k, v := range h.privateParams {
|
||||
pairs = append(pairs, &HeaderPair{Key: k, Value: v})
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) PrivateParams() map[string]interface{} {
|
||||
return h.privateParams
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case KeyTypeKey:
|
||||
return h.KeyType(), true
|
||||
case AlgorithmKey:
|
||||
if h.algorithm == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.algorithm), true
|
||||
case ECDSACrvKey:
|
||||
if h.crv == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.crv), true
|
||||
case ECDSADKey:
|
||||
if h.d == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.d, true
|
||||
case KeyIDKey:
|
||||
if h.keyID == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyID), true
|
||||
case KeyUsageKey:
|
||||
if h.keyUsage == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyUsage), true
|
||||
case KeyOpsKey:
|
||||
if h.keyops == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyops), true
|
||||
case ECDSAXKey:
|
||||
if h.x == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.x, true
|
||||
case X509CertChainKey:
|
||||
if h.x509CertChain == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertChain), true
|
||||
case X509CertThumbprintKey:
|
||||
if h.x509CertThumbprint == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprint), true
|
||||
case X509CertThumbprintS256Key:
|
||||
if h.x509CertThumbprintS256 == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprintS256), true
|
||||
case X509URLKey:
|
||||
if h.x509URL == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509URL), true
|
||||
case ECDSAYKey:
|
||||
if h.y == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.y, true
|
||||
default:
|
||||
v, ok := h.privateParams[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case "kty":
|
||||
return nil
|
||||
case AlgorithmKey:
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
h.algorithm = &v
|
||||
case fmt.Stringer:
|
||||
tmp := v.String()
|
||||
h.algorithm = &tmp
|
||||
default:
|
||||
return errors.Errorf(`invalid type for %s key: %T`, AlgorithmKey, value)
|
||||
}
|
||||
return nil
|
||||
case ECDSACrvKey:
|
||||
if v, ok := value.(jwa.EllipticCurveAlgorithm); ok {
|
||||
h.crv = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSACrvKey, value)
|
||||
case ECDSADKey:
|
||||
if v, ok := value.([]byte); ok {
|
||||
h.d = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSADKey, value)
|
||||
case KeyIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyIDKey, value)
|
||||
case KeyUsageKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyUsage = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyUsageKey, value)
|
||||
case KeyOpsKey:
|
||||
var acceptor KeyOperationList
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, KeyOpsKey)
|
||||
}
|
||||
h.keyops = &acceptor
|
||||
return nil
|
||||
case ECDSAXKey:
|
||||
if v, ok := value.([]byte); ok {
|
||||
h.x = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSAXKey, value)
|
||||
case X509CertChainKey:
|
||||
var acceptor CertificateChain
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, X509CertChainKey)
|
||||
}
|
||||
h.x509CertChain = &acceptor
|
||||
return nil
|
||||
case X509CertThumbprintKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprint = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintKey, value)
|
||||
case X509CertThumbprintS256Key:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprintS256 = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintS256Key, value)
|
||||
case X509URLKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509URL = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509URLKey, value)
|
||||
case ECDSAYKey:
|
||||
if v, ok := value.([]byte); ok {
|
||||
h.y = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSAYKey, value)
|
||||
default:
|
||||
if h.privateParams == nil {
|
||||
h.privateParams = map[string]interface{}{}
|
||||
}
|
||||
h.privateParams[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) UnmarshalJSON(buf []byte) error {
|
||||
var proxy ecdsaPrivateKeyMarshalProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal ecdsaPrivateKey`)
|
||||
}
|
||||
if proxy.XkeyType != jwa.EC {
|
||||
return errors.Errorf(`invalid kty value for ECDSAPrivateKey (%s)`, proxy.XkeyType)
|
||||
}
|
||||
h.algorithm = proxy.Xalgorithm
|
||||
h.crv = proxy.Xcrv
|
||||
if proxy.Xd == nil {
|
||||
return errors.New(`required field d is missing`)
|
||||
}
|
||||
if h.d = nil; proxy.Xd != nil {
|
||||
decoded, err := base64.DecodeString(*(proxy.Xd))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to decode base64 value for d`)
|
||||
}
|
||||
h.d = decoded
|
||||
}
|
||||
h.keyID = proxy.XkeyID
|
||||
h.keyUsage = proxy.XkeyUsage
|
||||
h.keyops = proxy.Xkeyops
|
||||
if proxy.Xx == nil {
|
||||
return errors.New(`required field x is missing`)
|
||||
}
|
||||
if h.x = nil; proxy.Xx != nil {
|
||||
decoded, err := base64.DecodeString(*(proxy.Xx))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to decode base64 value for x`)
|
||||
}
|
||||
h.x = decoded
|
||||
}
|
||||
h.x509CertChain = proxy.Xx509CertChain
|
||||
h.x509CertThumbprint = proxy.Xx509CertThumbprint
|
||||
h.x509CertThumbprintS256 = proxy.Xx509CertThumbprintS256
|
||||
h.x509URL = proxy.Xx509URL
|
||||
if proxy.Xy == nil {
|
||||
return errors.New(`required field y is missing`)
|
||||
}
|
||||
if h.y = nil; proxy.Xy != nil {
|
||||
decoded, err := base64.DecodeString(*(proxy.Xy))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to decode base64 value for y`)
|
||||
}
|
||||
h.y = decoded
|
||||
}
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to parse privsate parameters`)
|
||||
}
|
||||
delete(m, `kty`)
|
||||
delete(m, AlgorithmKey)
|
||||
delete(m, ECDSACrvKey)
|
||||
delete(m, ECDSADKey)
|
||||
delete(m, KeyIDKey)
|
||||
delete(m, KeyUsageKey)
|
||||
delete(m, KeyOpsKey)
|
||||
delete(m, ECDSAXKey)
|
||||
delete(m, X509CertChainKey)
|
||||
delete(m, X509CertThumbprintKey)
|
||||
delete(m, X509CertThumbprintS256Key)
|
||||
delete(m, X509URLKey)
|
||||
delete(m, ECDSAYKey)
|
||||
h.privateParams = m
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h ecdsaPrivateKey) MarshalJSON() ([]byte, error) {
|
||||
var proxy ecdsaPrivateKeyMarshalProxy
|
||||
proxy.XkeyType = jwa.EC
|
||||
proxy.Xalgorithm = h.algorithm
|
||||
proxy.Xcrv = h.crv
|
||||
if len(h.d) > 0 {
|
||||
v := base64.EncodeToString(h.d)
|
||||
proxy.Xd = &v
|
||||
}
|
||||
proxy.XkeyID = h.keyID
|
||||
proxy.XkeyUsage = h.keyUsage
|
||||
proxy.Xkeyops = h.keyops
|
||||
if len(h.x) > 0 {
|
||||
v := base64.EncodeToString(h.x)
|
||||
proxy.Xx = &v
|
||||
}
|
||||
proxy.Xx509CertChain = h.x509CertChain
|
||||
proxy.Xx509CertThumbprint = h.x509CertThumbprint
|
||||
proxy.Xx509CertThumbprintS256 = h.x509CertThumbprintS256
|
||||
proxy.Xx509URL = h.x509URL
|
||||
if len(h.y) > 0 {
|
||||
v := base64.EncodeToString(h.y)
|
||||
proxy.Xy = &v
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
if err := enc.Encode(proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to encode proxy to JSON`)
|
||||
}
|
||||
hasContent := buf.Len() > 3 // encoding/json always adds a newline, so "{}\n" is the empty hash
|
||||
if l := len(h.privateParams); l > 0 {
|
||||
buf.Truncate(buf.Len() - 2)
|
||||
keys := make([]string, 0, l)
|
||||
for k := range h.privateParams {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for i, k := range keys {
|
||||
if hasContent || i > 0 {
|
||||
fmt.Fprintf(&buf, `,`)
|
||||
}
|
||||
fmt.Fprintf(&buf, `%s:`, strconv.Quote(k))
|
||||
if err := enc.Encode(h.privateParams[k]); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to encode private param %s`, k)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(&buf, `}`)
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Iterate(ctx context.Context) HeaderIterator {
|
||||
ch := make(chan *HeaderPair)
|
||||
go h.iterate(ctx, ch)
|
||||
return mapiter.New(ch)
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) Walk(ctx context.Context, visitor HeaderVisitor) error {
|
||||
return iter.WalkMap(ctx, h, visitor)
|
||||
}
|
||||
|
||||
func (h *ecdsaPrivateKey) AsMap(ctx context.Context) (map[string]interface{}, error) {
|
||||
return iter.AsMap(ctx, h)
|
||||
}
|
||||
|
||||
type ECDSAPublicKey interface {
|
||||
Key
|
||||
FromRaw(*ecdsa.PublicKey) error
|
||||
Crv() jwa.EllipticCurveAlgorithm
|
||||
X() []byte
|
||||
Y() []byte
|
||||
}
|
||||
|
||||
type ecdsaPublicKey struct {
|
||||
algorithm *string // https://tools.ietf.org/html/rfc7517#section-4.4
|
||||
crv *jwa.EllipticCurveAlgorithm
|
||||
keyID *string // https://tools.ietf.org/html/rfc7515#section-4.1.4
|
||||
keyUsage *string // https://tools.ietf.org/html/rfc7517#section-4.2
|
||||
keyops *KeyOperationList // https://tools.ietf.org/html/rfc7517#section-4.3
|
||||
x []byte
|
||||
x509CertChain *CertificateChain // https://tools.ietf.org/html/rfc7515#section-4.1.6
|
||||
x509CertThumbprint *string // https://tools.ietf.org/html/rfc7515#section-4.1.7
|
||||
x509CertThumbprintS256 *string // https://tools.ietf.org/html/rfc7515#section-4.1.8
|
||||
x509URL *string // https://tools.ietf.org/html/rfc7515#section-4.1.5
|
||||
y []byte
|
||||
privateParams map[string]interface{}
|
||||
}
|
||||
|
||||
type ecdsaPublicKeyMarshalProxy struct {
|
||||
XkeyType jwa.KeyType `json:"kty"`
|
||||
Xalgorithm *string `json:"alg,omitempty"`
|
||||
Xcrv *jwa.EllipticCurveAlgorithm `json:"crv,omitempty"`
|
||||
XkeyID *string `json:"kid,omitempty"`
|
||||
XkeyUsage *string `json:"use,omitempty"`
|
||||
Xkeyops *KeyOperationList `json:"key_ops,omitempty"`
|
||||
Xx *string `json:"x,omitempty"`
|
||||
Xx509CertChain *CertificateChain `json:"x5c,omitempty"`
|
||||
Xx509CertThumbprint *string `json:"x5t,omitempty"`
|
||||
Xx509CertThumbprintS256 *string `json:"x5t#S256,omitempty"`
|
||||
Xx509URL *string `json:"x5u,omitempty"`
|
||||
Xy *string `json:"y,omitempty"`
|
||||
}
|
||||
|
||||
func (h ecdsaPublicKey) KeyType() jwa.KeyType {
|
||||
return jwa.EC
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Algorithm() string {
|
||||
if h.algorithm != nil {
|
||||
return *(h.algorithm)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Crv() jwa.EllipticCurveAlgorithm {
|
||||
if h.crv != nil {
|
||||
return *(h.crv)
|
||||
}
|
||||
return jwa.InvalidEllipticCurve
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) KeyID() string {
|
||||
if h.keyID != nil {
|
||||
return *(h.keyID)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) KeyUsage() string {
|
||||
if h.keyUsage != nil {
|
||||
return *(h.keyUsage)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) KeyOps() KeyOperationList {
|
||||
if h.keyops != nil {
|
||||
return *(h.keyops)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) X() []byte {
|
||||
return h.x
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) X509CertChain() []*x509.Certificate {
|
||||
if h.x509CertChain != nil {
|
||||
return h.x509CertChain.Get()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) X509CertThumbprint() string {
|
||||
if h.x509CertThumbprint != nil {
|
||||
return *(h.x509CertThumbprint)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) X509CertThumbprintS256() string {
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
return *(h.x509CertThumbprintS256)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) X509URL() string {
|
||||
if h.x509URL != nil {
|
||||
return *(h.x509URL)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Y() []byte {
|
||||
return h.y
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) iterate(ctx context.Context, ch chan *HeaderPair) {
|
||||
defer close(ch)
|
||||
|
||||
var pairs []*HeaderPair
|
||||
pairs = append(pairs, &HeaderPair{Key: "kty", Value: jwa.EC})
|
||||
if h.algorithm != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: AlgorithmKey, Value: *(h.algorithm)})
|
||||
}
|
||||
if h.crv != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSACrvKey, Value: *(h.crv)})
|
||||
}
|
||||
if h.keyID != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyIDKey, Value: *(h.keyID)})
|
||||
}
|
||||
if h.keyUsage != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyUsageKey, Value: *(h.keyUsage)})
|
||||
}
|
||||
if h.keyops != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyOpsKey, Value: *(h.keyops)})
|
||||
}
|
||||
if h.x != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSAXKey, Value: h.x})
|
||||
}
|
||||
if h.x509CertChain != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertChainKey, Value: *(h.x509CertChain)})
|
||||
}
|
||||
if h.x509CertThumbprint != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintKey, Value: *(h.x509CertThumbprint)})
|
||||
}
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintS256Key, Value: *(h.x509CertThumbprintS256)})
|
||||
}
|
||||
if h.x509URL != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509URLKey, Value: *(h.x509URL)})
|
||||
}
|
||||
if h.y != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ECDSAYKey, Value: h.y})
|
||||
}
|
||||
for k, v := range h.privateParams {
|
||||
pairs = append(pairs, &HeaderPair{Key: k, Value: v})
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) PrivateParams() map[string]interface{} {
|
||||
return h.privateParams
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case KeyTypeKey:
|
||||
return h.KeyType(), true
|
||||
case AlgorithmKey:
|
||||
if h.algorithm == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.algorithm), true
|
||||
case ECDSACrvKey:
|
||||
if h.crv == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.crv), true
|
||||
case KeyIDKey:
|
||||
if h.keyID == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyID), true
|
||||
case KeyUsageKey:
|
||||
if h.keyUsage == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyUsage), true
|
||||
case KeyOpsKey:
|
||||
if h.keyops == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyops), true
|
||||
case ECDSAXKey:
|
||||
if h.x == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.x, true
|
||||
case X509CertChainKey:
|
||||
if h.x509CertChain == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertChain), true
|
||||
case X509CertThumbprintKey:
|
||||
if h.x509CertThumbprint == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprint), true
|
||||
case X509CertThumbprintS256Key:
|
||||
if h.x509CertThumbprintS256 == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprintS256), true
|
||||
case X509URLKey:
|
||||
if h.x509URL == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509URL), true
|
||||
case ECDSAYKey:
|
||||
if h.y == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.y, true
|
||||
default:
|
||||
v, ok := h.privateParams[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case "kty":
|
||||
return nil
|
||||
case AlgorithmKey:
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
h.algorithm = &v
|
||||
case fmt.Stringer:
|
||||
tmp := v.String()
|
||||
h.algorithm = &tmp
|
||||
default:
|
||||
return errors.Errorf(`invalid type for %s key: %T`, AlgorithmKey, value)
|
||||
}
|
||||
return nil
|
||||
case ECDSACrvKey:
|
||||
if v, ok := value.(jwa.EllipticCurveAlgorithm); ok {
|
||||
h.crv = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSACrvKey, value)
|
||||
case KeyIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyIDKey, value)
|
||||
case KeyUsageKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyUsage = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyUsageKey, value)
|
||||
case KeyOpsKey:
|
||||
var acceptor KeyOperationList
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, KeyOpsKey)
|
||||
}
|
||||
h.keyops = &acceptor
|
||||
return nil
|
||||
case ECDSAXKey:
|
||||
if v, ok := value.([]byte); ok {
|
||||
h.x = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSAXKey, value)
|
||||
case X509CertChainKey:
|
||||
var acceptor CertificateChain
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, X509CertChainKey)
|
||||
}
|
||||
h.x509CertChain = &acceptor
|
||||
return nil
|
||||
case X509CertThumbprintKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprint = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintKey, value)
|
||||
case X509CertThumbprintS256Key:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprintS256 = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintS256Key, value)
|
||||
case X509URLKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509URL = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509URLKey, value)
|
||||
case ECDSAYKey:
|
||||
if v, ok := value.([]byte); ok {
|
||||
h.y = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ECDSAYKey, value)
|
||||
default:
|
||||
if h.privateParams == nil {
|
||||
h.privateParams = map[string]interface{}{}
|
||||
}
|
||||
h.privateParams[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) UnmarshalJSON(buf []byte) error {
|
||||
var proxy ecdsaPublicKeyMarshalProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal ecdsaPublicKey`)
|
||||
}
|
||||
if proxy.XkeyType != jwa.EC {
|
||||
return errors.Errorf(`invalid kty value for ECDSAPublicKey (%s)`, proxy.XkeyType)
|
||||
}
|
||||
h.algorithm = proxy.Xalgorithm
|
||||
h.crv = proxy.Xcrv
|
||||
h.keyID = proxy.XkeyID
|
||||
h.keyUsage = proxy.XkeyUsage
|
||||
h.keyops = proxy.Xkeyops
|
||||
if proxy.Xx == nil {
|
||||
return errors.New(`required field x is missing`)
|
||||
}
|
||||
if h.x = nil; proxy.Xx != nil {
|
||||
decoded, err := base64.DecodeString(*(proxy.Xx))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to decode base64 value for x`)
|
||||
}
|
||||
h.x = decoded
|
||||
}
|
||||
h.x509CertChain = proxy.Xx509CertChain
|
||||
h.x509CertThumbprint = proxy.Xx509CertThumbprint
|
||||
h.x509CertThumbprintS256 = proxy.Xx509CertThumbprintS256
|
||||
h.x509URL = proxy.Xx509URL
|
||||
if proxy.Xy == nil {
|
||||
return errors.New(`required field y is missing`)
|
||||
}
|
||||
if h.y = nil; proxy.Xy != nil {
|
||||
decoded, err := base64.DecodeString(*(proxy.Xy))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to decode base64 value for y`)
|
||||
}
|
||||
h.y = decoded
|
||||
}
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to parse privsate parameters`)
|
||||
}
|
||||
delete(m, `kty`)
|
||||
delete(m, AlgorithmKey)
|
||||
delete(m, ECDSACrvKey)
|
||||
delete(m, KeyIDKey)
|
||||
delete(m, KeyUsageKey)
|
||||
delete(m, KeyOpsKey)
|
||||
delete(m, ECDSAXKey)
|
||||
delete(m, X509CertChainKey)
|
||||
delete(m, X509CertThumbprintKey)
|
||||
delete(m, X509CertThumbprintS256Key)
|
||||
delete(m, X509URLKey)
|
||||
delete(m, ECDSAYKey)
|
||||
h.privateParams = m
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h ecdsaPublicKey) MarshalJSON() ([]byte, error) {
|
||||
var proxy ecdsaPublicKeyMarshalProxy
|
||||
proxy.XkeyType = jwa.EC
|
||||
proxy.Xalgorithm = h.algorithm
|
||||
proxy.Xcrv = h.crv
|
||||
proxy.XkeyID = h.keyID
|
||||
proxy.XkeyUsage = h.keyUsage
|
||||
proxy.Xkeyops = h.keyops
|
||||
if len(h.x) > 0 {
|
||||
v := base64.EncodeToString(h.x)
|
||||
proxy.Xx = &v
|
||||
}
|
||||
proxy.Xx509CertChain = h.x509CertChain
|
||||
proxy.Xx509CertThumbprint = h.x509CertThumbprint
|
||||
proxy.Xx509CertThumbprintS256 = h.x509CertThumbprintS256
|
||||
proxy.Xx509URL = h.x509URL
|
||||
if len(h.y) > 0 {
|
||||
v := base64.EncodeToString(h.y)
|
||||
proxy.Xy = &v
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
if err := enc.Encode(proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to encode proxy to JSON`)
|
||||
}
|
||||
hasContent := buf.Len() > 3 // encoding/json always adds a newline, so "{}\n" is the empty hash
|
||||
if l := len(h.privateParams); l > 0 {
|
||||
buf.Truncate(buf.Len() - 2)
|
||||
keys := make([]string, 0, l)
|
||||
for k := range h.privateParams {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for i, k := range keys {
|
||||
if hasContent || i > 0 {
|
||||
fmt.Fprintf(&buf, `,`)
|
||||
}
|
||||
fmt.Fprintf(&buf, `%s:`, strconv.Quote(k))
|
||||
if err := enc.Encode(h.privateParams[k]); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to encode private param %s`, k)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(&buf, `}`)
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Iterate(ctx context.Context) HeaderIterator {
|
||||
ch := make(chan *HeaderPair)
|
||||
go h.iterate(ctx, ch)
|
||||
return mapiter.New(ch)
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) Walk(ctx context.Context, visitor HeaderVisitor) error {
|
||||
return iter.WalkMap(ctx, h, visitor)
|
||||
}
|
||||
|
||||
func (h *ecdsaPublicKey) AsMap(ctx context.Context) (map[string]interface{}, error) {
|
||||
return iter.AsMap(ctx, h)
|
||||
}
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
|
||||
"github.com/lestrrat-go/iter/arrayiter"
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
)
|
||||
|
||||
// KeyUsageType is used to denote what this key should be used for
|
||||
type KeyUsageType string
|
||||
|
||||
const (
|
||||
// ForSignature is the value used in the headers to indicate that
|
||||
// this key should be used for signatures
|
||||
ForSignature KeyUsageType = "sig"
|
||||
// ForEncryption is the value used in the headers to indicate that
|
||||
// this key should be used for encryptiong
|
||||
ForEncryption KeyUsageType = "enc"
|
||||
)
|
||||
|
||||
type CertificateChain struct {
|
||||
certs []*x509.Certificate
|
||||
}
|
||||
|
||||
type KeyOperation string
|
||||
type KeyOperationList []KeyOperation
|
||||
|
||||
const (
|
||||
KeyOpSign KeyOperation = "sign" // (compute digital signature or MAC)
|
||||
KeyOpVerify KeyOperation = "verify" // (verify digital signature or MAC)
|
||||
KeyOpEncrypt KeyOperation = "encrypt" // (encrypt content)
|
||||
KeyOpDecrypt KeyOperation = "decrypt" // (decrypt content and validate decryption, if applicable)
|
||||
KeyOpWrapKey KeyOperation = "wrapKey" // (encrypt key)
|
||||
KeyOpUnwrapKey KeyOperation = "unwrapKey" // (decrypt key and validate decryption, if applicable)
|
||||
KeyOpDeriveKey KeyOperation = "deriveKey" // (derive key)
|
||||
KeyOpDeriveBits KeyOperation = "deriveBits" // (derive bits not to be used as a key)
|
||||
)
|
||||
|
||||
// Set is a convenience struct to allow generating and parsing
|
||||
// JWK sets as opposed to single JWKs
|
||||
type Set struct {
|
||||
Keys []Key `json:"keys"`
|
||||
}
|
||||
|
||||
type HeaderVisitor = iter.MapVisitor
|
||||
type HeaderVisitorFunc = iter.MapVisitorFunc
|
||||
type HeaderPair = mapiter.Pair
|
||||
type HeaderIterator = mapiter.Iterator
|
||||
type KeyPair = arrayiter.Pair
|
||||
type KeyIterator = arrayiter.Iterator
|
||||
+84
@@ -0,0 +1,84 @@
|
||||
// This file is auto-generated. DO NOT EDIT
|
||||
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto"
|
||||
"crypto/x509"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
)
|
||||
|
||||
const (
|
||||
KeyTypeKey = "kty"
|
||||
KeyUsageKey = "use"
|
||||
KeyOpsKey = "key_ops"
|
||||
AlgorithmKey = "alg"
|
||||
KeyIDKey = "kid"
|
||||
X509URLKey = "x5u"
|
||||
X509CertChainKey = "x5c"
|
||||
X509CertThumbprintKey = "x5t"
|
||||
X509CertThumbprintS256Key = "x5t#S256"
|
||||
)
|
||||
|
||||
// Key defines the minimal interface for each of the
|
||||
// key types. Their use and implementation differ significantly
|
||||
// between each key types, so you should use type assertions
|
||||
// to perform more specific tasks with each key
|
||||
type Key interface {
|
||||
// Get returns the value of a single field. The second boolean return value
|
||||
// will be false if the field is not stored in the source
|
||||
//
|
||||
// This method, which returns an `interface{}`, exists because
|
||||
// these objects can contain extra _arbitrary_ fields that users can
|
||||
// specify, and there is no way of knowing what type they could be
|
||||
Get(string) (interface{}, bool)
|
||||
|
||||
// Set sets the value of a single field. Note that certain fields,
|
||||
// notably "kty" cannot be altered, but will not return an error
|
||||
//
|
||||
// This method, which takes an `interface{}`, exists because
|
||||
// these objects can contain extra _arbitrary_ fields that users can
|
||||
// specify, and there is no way of knowing what type they could be
|
||||
Set(string, interface{}) error
|
||||
|
||||
// Raw creates the corresponding raw key. For example,
|
||||
// EC types would create *ecdsa.PublicKey or *ecdsa.PrivateKey,
|
||||
// and OctetSeq types create a []byte key.
|
||||
//
|
||||
// If you do not know the exact type of a jwk.Key before attempting
|
||||
// to obtain the raw key, you can simply pass a pointer to an
|
||||
// empty interface as the first argument.
|
||||
//
|
||||
// If you already know the exact type, it is recommended that you
|
||||
// pass a pointer to the actual key type (e.g. *rsa.PrivateKey, *ecdsa.PublicKey
|
||||
// for efficiency
|
||||
Raw(interface{}) error
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
Thumbprint(crypto.Hash) ([]byte, error)
|
||||
|
||||
// Iterate returns an iterator that returns all keys and values
|
||||
Iterate(ctx context.Context) HeaderIterator
|
||||
|
||||
// Walk is a utility tool that allows a visitor to iterate all keys and values
|
||||
Walk(context.Context, HeaderVisitor) error
|
||||
|
||||
// AsMap is a utility tool returns a map that contains the same fields as the source
|
||||
AsMap(context.Context) (map[string]interface{}, error)
|
||||
|
||||
// PrivateParams returns the non-standard elements in the source structure
|
||||
PrivateParams() map[string]interface{}
|
||||
|
||||
KeyType() jwa.KeyType
|
||||
KeyUsage() string
|
||||
KeyOps() KeyOperationList
|
||||
Algorithm() string
|
||||
KeyID() string
|
||||
X509URL() string
|
||||
X509CertChain() []*x509.Certificate
|
||||
X509CertThumbprint() string
|
||||
X509CertThumbprintS256() string
|
||||
}
|
||||
+384
@@ -0,0 +1,384 @@
|
||||
//go:generate go run internal/cmd/genheader/main.go
|
||||
|
||||
// Package jwk implements JWK as described in https://tools.ietf.org/html/rfc7517
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/rsa"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/lestrrat-go/iter/arrayiter"
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// New creates a jwk.Key from the given key (RSA/ECDSA/symmetric keys).
|
||||
func New(key interface{}) (Key, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`jwk.New requires a non-nil key`)
|
||||
}
|
||||
|
||||
var ptr interface{}
|
||||
switch v := key.(type) {
|
||||
case rsa.PrivateKey:
|
||||
ptr = &v
|
||||
case rsa.PublicKey:
|
||||
ptr = &v
|
||||
case ecdsa.PrivateKey:
|
||||
ptr = &v
|
||||
case ecdsa.PublicKey:
|
||||
ptr = &v
|
||||
default:
|
||||
ptr = v
|
||||
}
|
||||
|
||||
switch rawKey := ptr.(type) {
|
||||
case *rsa.PrivateKey:
|
||||
k := NewRSAPrivateKey()
|
||||
if err := k.FromRaw(rawKey); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to initialize %T from %T`, k, rawKey)
|
||||
}
|
||||
return k, nil
|
||||
case *rsa.PublicKey:
|
||||
k := NewRSAPublicKey()
|
||||
if err := k.FromRaw(rawKey); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to initialize %T from %T`, k, rawKey)
|
||||
}
|
||||
return k, nil
|
||||
case *ecdsa.PrivateKey:
|
||||
k := NewECDSAPrivateKey()
|
||||
if err := k.FromRaw(rawKey); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to initialize %T from %T`, k, rawKey)
|
||||
}
|
||||
return k, nil
|
||||
case *ecdsa.PublicKey:
|
||||
k := NewECDSAPublicKey()
|
||||
if err := k.FromRaw(rawKey); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to initialize %T from %T`, k, rawKey)
|
||||
}
|
||||
return k, nil
|
||||
case []byte:
|
||||
k := NewSymmetricKey()
|
||||
if err := k.FromRaw(rawKey); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to initialize %T from %T`, k, rawKey)
|
||||
}
|
||||
return k, nil
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid key type '%T' for jwk.New`, key)
|
||||
}
|
||||
}
|
||||
|
||||
// PublicKeyOf returns the corresponding public key of the given
|
||||
// value `v`. For example, if v is a `*rsa.PrivateKey`, then
|
||||
// `*rsa.PublicKey` is returned.
|
||||
//
|
||||
// If given a public key, then the same public key will be returned.
|
||||
// For example, if v is a `*rsa.PublicKey`, then the same value
|
||||
// is returned.
|
||||
//
|
||||
// If v is of a type that we don't support, an error is returned.
|
||||
//
|
||||
// This is useful when you are dealing with the jwk.Key interface
|
||||
// alone and you don't know before hand what the underlying key
|
||||
// type is, but you still want to obtain the corresponding public key
|
||||
func PublicKeyOf(v interface{}) (interface{}, error) {
|
||||
// may be a silly idea, but if the user gave us a non-pointer value...
|
||||
var ptr interface{}
|
||||
switch v := v.(type) {
|
||||
case rsa.PrivateKey:
|
||||
ptr = &v
|
||||
case rsa.PublicKey:
|
||||
ptr = &v
|
||||
case ecdsa.PrivateKey:
|
||||
ptr = &v
|
||||
case ecdsa.PublicKey:
|
||||
ptr = &v
|
||||
default:
|
||||
ptr = v
|
||||
}
|
||||
|
||||
switch x := ptr.(type) {
|
||||
case *rsa.PrivateKey:
|
||||
return &x.PublicKey, nil
|
||||
case *rsa.PublicKey:
|
||||
return x, nil
|
||||
case *ecdsa.PrivateKey:
|
||||
return &x.PublicKey, nil
|
||||
case *ecdsa.PublicKey:
|
||||
return x, nil
|
||||
case []byte:
|
||||
return x, nil
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid key type passed to PublicKeyOf (%T)`, v)
|
||||
}
|
||||
}
|
||||
|
||||
// Fetch fetches a JWK resource specified by a URL
|
||||
func Fetch(urlstring string, options ...Option) (*Set, error) {
|
||||
u, err := url.Parse(urlstring)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to parse url`)
|
||||
}
|
||||
|
||||
switch u.Scheme {
|
||||
case "http", "https":
|
||||
return FetchHTTP(urlstring, options...)
|
||||
case "file":
|
||||
f, err := os.Open(u.Path)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to open jwk file`)
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
return Parse(f)
|
||||
}
|
||||
return nil, errors.Errorf(`invalid url scheme %s`, u.Scheme)
|
||||
}
|
||||
|
||||
// FetchHTTP wraps FetchHTTPWithContext using the background context.
|
||||
func FetchHTTP(jwkurl string, options ...Option) (*Set, error) {
|
||||
return FetchHTTPWithContext(context.Background(), jwkurl, options...)
|
||||
}
|
||||
|
||||
// FetchHTTPWithContext fetches the remote JWK and parses its contents
|
||||
func FetchHTTPWithContext(ctx context.Context, jwkurl string, options ...Option) (*Set, error) {
|
||||
httpcl := http.DefaultClient
|
||||
for _, option := range options {
|
||||
switch option.Name() {
|
||||
case optkeyHTTPClient:
|
||||
httpcl = option.Value().(*http.Client)
|
||||
}
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodGet, jwkurl, nil)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to new request to remote JWK")
|
||||
}
|
||||
|
||||
res, err := httpcl.Do(req.WithContext(ctx))
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to fetch remote JWK")
|
||||
}
|
||||
defer res.Body.Close()
|
||||
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("failed to fetch remote JWK (status = %d)", res.StatusCode)
|
||||
}
|
||||
|
||||
return Parse(res.Body)
|
||||
}
|
||||
|
||||
func ParseKey(data []byte) (Key, error) {
|
||||
var hint struct {
|
||||
Kty string `json:"kty"`
|
||||
D json.RawMessage `json:"d"`
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(data, &hint); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unmarshal JSON into key hint`)
|
||||
}
|
||||
|
||||
var key Key
|
||||
switch jwa.KeyType(hint.Kty) {
|
||||
case jwa.RSA:
|
||||
if len(hint.D) > 0 {
|
||||
key = newRSAPrivateKey()
|
||||
} else {
|
||||
key = newRSAPublicKey()
|
||||
}
|
||||
case jwa.EC:
|
||||
if len(hint.D) > 0 {
|
||||
key = newECDSAPrivateKey()
|
||||
} else {
|
||||
key = newECDSAPublicKey()
|
||||
}
|
||||
case jwa.OctetSeq:
|
||||
key = newSymmetricKey()
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid key type from JSON (%s)`, hint.Kty)
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(data, key); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to unmarshal JSON into key (%T)`, key)
|
||||
}
|
||||
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (s *Set) UnmarshalJSON(data []byte) error {
|
||||
var proxy struct {
|
||||
Keys []json.RawMessage `json:"keys"`
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(data, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal into Key (proxy)`)
|
||||
}
|
||||
|
||||
if len(proxy.Keys) == 0 {
|
||||
k, err := ParseKey(data)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal key from JSON headers`)
|
||||
}
|
||||
s.Keys = append(s.Keys, k)
|
||||
} else {
|
||||
for i, buf := range proxy.Keys {
|
||||
k, err := ParseKey([]byte(buf))
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to unmarshal key #%d (total %d) from multi-key JWK set`, i+1, len(proxy.Keys))
|
||||
}
|
||||
s.Keys = append(s.Keys, k)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Parse parses JWK from the incoming io.Reader. This function can handle
|
||||
// both single-key and multi-key formats. If you know before hand which
|
||||
// format the incoming data is in, you might want to consider using
|
||||
// "encoding/json" directly
|
||||
//
|
||||
// Note that a successful parsing does NOT guarantee a valid key
|
||||
func Parse(in io.Reader) (*Set, error) {
|
||||
var s Set
|
||||
if err := json.NewDecoder(in).Decode(&s); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to unmarshal JWK")
|
||||
}
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
// ParseBytes parses JWK from the incoming byte buffer.
|
||||
//
|
||||
// Note that a successful parsing does NOT guarantee a valid key
|
||||
func ParseBytes(buf []byte) (*Set, error) {
|
||||
return Parse(bytes.NewReader(buf))
|
||||
}
|
||||
|
||||
// ParseString parses JWK from the incoming string.
|
||||
//
|
||||
// Note that a successful parsing does NOT guarantee a valid key
|
||||
func ParseString(s string) (*Set, error) {
|
||||
return Parse(strings.NewReader(s))
|
||||
}
|
||||
|
||||
// LookupKeyID looks for keys matching the given key id. Note that the
|
||||
// Set *may* contain multiple keys with the same key id
|
||||
func (s Set) LookupKeyID(kid string) []Key {
|
||||
var keys []Key
|
||||
for iter := s.Iterate(context.TODO()); iter.Next(context.TODO()); {
|
||||
pair := iter.Pair()
|
||||
key := pair.Value.(Key)
|
||||
if key.KeyID() == kid {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func (s *Set) Len() int {
|
||||
return len(s.Keys)
|
||||
}
|
||||
|
||||
func (s *Set) Iterate(ctx context.Context) KeyIterator {
|
||||
ch := make(chan *KeyPair, s.Len())
|
||||
go iterate(ctx, s.Keys, ch)
|
||||
return arrayiter.New(ch)
|
||||
}
|
||||
|
||||
func iterate(ctx context.Context, keys []Key, ch chan *KeyPair) {
|
||||
defer close(ch)
|
||||
|
||||
for i, key := range keys {
|
||||
pair := &KeyPair{Index: i, Value: key}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// assignRawResult is a convenience function to safely
|
||||
// assign arbitrary values from Raw
|
||||
func assignRawResult(v, t interface{}) error {
|
||||
orv := reflect.ValueOf(t) // save this value for error reporting
|
||||
result := orv
|
||||
|
||||
// t can be a pointer or a slice, and the code will slightly change
|
||||
// depending on this
|
||||
var isSlice bool
|
||||
switch result.Kind() {
|
||||
case reflect.Ptr:
|
||||
// no op
|
||||
case reflect.Slice:
|
||||
isSlice = true
|
||||
default:
|
||||
return errors.Errorf("argument t to assignRawResult must be a pointer or a slice: %T", t)
|
||||
}
|
||||
|
||||
rv := reflect.ValueOf(v)
|
||||
if rv.Kind() != reflect.Ptr {
|
||||
return errors.Errorf(`argument to Raw() must be a pointer: %T`, v)
|
||||
}
|
||||
|
||||
dst := rv.Elem()
|
||||
switch dst.Kind() {
|
||||
case reflect.Interface:
|
||||
// If it's an interface, we can just assign the pointer to the interface{}
|
||||
default:
|
||||
// If it's a pointer to the struct we're looking for, we need to set
|
||||
// the de-referenced struct
|
||||
if !isSlice {
|
||||
result = result.Elem()
|
||||
}
|
||||
}
|
||||
if !result.Type().AssignableTo(dst.Type()) {
|
||||
return errors.Errorf(`argument to Raw() must be compatible with %T (was %T)`, orv.Interface(), v)
|
||||
}
|
||||
|
||||
if !dst.CanSet() {
|
||||
return errors.Errorf(`argument to Raw() must be settable`)
|
||||
}
|
||||
dst.Set(result)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// AssignKeyID is a convenience function to automatically assign the "kid"
|
||||
// section of the key, if it already doesn't have one. It uses Key.Thumbprint
|
||||
// method with crypto.SHA256 as the default hashing algorithm
|
||||
func AssignKeyID(key Key, options ...Option) error {
|
||||
if _, ok := key.Get(KeyIDKey); ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
hash := crypto.SHA256
|
||||
for _, option := range options {
|
||||
switch option.Name() {
|
||||
case optkeyThumbprintHash:
|
||||
hash = option.Value().(crypto.Hash)
|
||||
}
|
||||
}
|
||||
|
||||
h, err := key.Thumbprint(hash)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to generate thumbprint`)
|
||||
}
|
||||
|
||||
if err := key.Set(KeyIDKey, base64.EncodeToString(h)); err != nil {
|
||||
return errors.Wrap(err, `failed to set "kid"`)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
package jwk
|
||||
|
||||
import "github.com/pkg/errors"
|
||||
|
||||
func (ops *KeyOperationList) Get() KeyOperationList {
|
||||
if ops == nil {
|
||||
return nil
|
||||
}
|
||||
return *ops
|
||||
}
|
||||
|
||||
func (ops *KeyOperationList) Accept(v interface{}) error {
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
return ops.Accept([]string{x})
|
||||
case []interface{}:
|
||||
l := make([]string, len(x))
|
||||
for i, e := range x {
|
||||
if es, ok := e.(string); ok {
|
||||
l[i] = es
|
||||
} else {
|
||||
return errors.Errorf(`invalid list element type: expected string, got %T`, v)
|
||||
}
|
||||
}
|
||||
return ops.Accept(l)
|
||||
case []string:
|
||||
list := make(KeyOperationList, len(x))
|
||||
for i, e := range x {
|
||||
switch e := KeyOperation(e); e {
|
||||
case KeyOpSign, KeyOpVerify, KeyOpEncrypt, KeyOpDecrypt, KeyOpWrapKey, KeyOpUnwrapKey, KeyOpDeriveKey, KeyOpDeriveBits:
|
||||
list[i] = e
|
||||
default:
|
||||
return errors.Errorf(`invalid keyoperation %v`, e)
|
||||
}
|
||||
}
|
||||
|
||||
*ops = list
|
||||
return nil
|
||||
case []KeyOperation:
|
||||
list := make(KeyOperationList, len(x))
|
||||
for i, e := range x {
|
||||
switch e {
|
||||
case KeyOpSign, KeyOpVerify, KeyOpEncrypt, KeyOpDecrypt, KeyOpWrapKey, KeyOpUnwrapKey, KeyOpDeriveKey, KeyOpDeriveBits:
|
||||
list[i] = e
|
||||
default:
|
||||
return errors.Errorf(`invalid keyoperation %v`, e)
|
||||
}
|
||||
}
|
||||
|
||||
*ops = list
|
||||
return nil
|
||||
case KeyOperationList:
|
||||
*ops = x
|
||||
return nil
|
||||
default:
|
||||
return errors.Errorf(`invalid value %T`, v)
|
||||
}
|
||||
}
|
||||
+23
@@ -0,0 +1,23 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"net/http"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/option"
|
||||
)
|
||||
|
||||
type Option = option.Interface
|
||||
|
||||
const (
|
||||
optkeyHTTPClient = `http-client`
|
||||
optkeyThumbprintHash = `thumbprint-hash`
|
||||
)
|
||||
|
||||
func WithHTTPClient(cl *http.Client) Option {
|
||||
return option.New(optkeyHTTPClient, cl)
|
||||
}
|
||||
|
||||
func WithThumbprintHash(h crypto.Hash) Option {
|
||||
return option.New(optkeyThumbprintHash, h)
|
||||
}
|
||||
+194
@@ -0,0 +1,194 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto"
|
||||
"crypto/rsa"
|
||||
"encoding/binary"
|
||||
"math/big"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/lestrrat-go/jwx/internal/pool"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func NewRSAPublicKey() RSAPublicKey {
|
||||
return newRSAPublicKey()
|
||||
}
|
||||
|
||||
func newRSAPublicKey() *rsaPublicKey {
|
||||
return &rsaPublicKey{
|
||||
privateParams: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
func NewRSAPrivateKey() RSAPrivateKey {
|
||||
return newRSAPrivateKey()
|
||||
}
|
||||
|
||||
func newRSAPrivateKey() *rsaPrivateKey {
|
||||
return &rsaPrivateKey{
|
||||
privateParams: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (k *rsaPrivateKey) FromRaw(rawKey *rsa.PrivateKey) error {
|
||||
k.d = rawKey.D.Bytes()
|
||||
if len(rawKey.Primes) < 2 {
|
||||
return errors.Errorf(`invalid number of primes in rsa.PrivateKey: need 2, got %d`, len(rawKey.Primes))
|
||||
}
|
||||
|
||||
k.p = rawKey.Primes[0].Bytes()
|
||||
k.q = rawKey.Primes[1].Bytes()
|
||||
|
||||
if v := rawKey.Precomputed.Dp; v != nil {
|
||||
k.dp = v.Bytes()
|
||||
}
|
||||
if v := rawKey.Precomputed.Dq; v != nil {
|
||||
k.dq = v.Bytes()
|
||||
}
|
||||
if v := rawKey.Precomputed.Qinv; v != nil {
|
||||
k.qi = v.Bytes()
|
||||
}
|
||||
|
||||
k.n = rawKey.PublicKey.N.Bytes()
|
||||
data := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(data, uint64(rawKey.PublicKey.E))
|
||||
i := 0
|
||||
for ; i < len(data); i++ {
|
||||
if data[i] != 0x0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
k.e = data[i:]
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *rsaPublicKey) FromRaw(rawKey *rsa.PublicKey) error {
|
||||
k.n = rawKey.N.Bytes()
|
||||
data := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(data, uint64(rawKey.E))
|
||||
i := 0
|
||||
for ; i < len(data); i++ {
|
||||
if data[i] != 0x0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
k.e = data[i:]
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *rsaPrivateKey) Raw(v interface{}) error {
|
||||
var d, q, p big.Int // note: do not use from sync.Pool
|
||||
|
||||
d.SetBytes(k.d)
|
||||
q.SetBytes(k.q)
|
||||
p.SetBytes(k.p)
|
||||
|
||||
// optional fields
|
||||
var dp, dq, qi *big.Int
|
||||
if len(k.dp) > 0 {
|
||||
dp = &big.Int{} // note: do not use from sync.Pool
|
||||
dp.SetBytes(k.dp)
|
||||
}
|
||||
|
||||
if len(k.dq) > 0 {
|
||||
dq = &big.Int{} // note: do not use from sync.Pool
|
||||
dq.SetBytes(k.dq)
|
||||
}
|
||||
|
||||
if len(k.qi) > 0 {
|
||||
qi = &big.Int{} // note: do not use from sync.Pool
|
||||
qi.SetBytes(k.qi)
|
||||
}
|
||||
|
||||
var key rsa.PrivateKey
|
||||
|
||||
pubk := newRSAPublicKey()
|
||||
pubk.n = k.n
|
||||
pubk.e = k.e
|
||||
if err := pubk.Raw(&key.PublicKey); err != nil {
|
||||
return errors.Wrap(err, `failed to materialize RSA public key`)
|
||||
}
|
||||
|
||||
key.D = &d
|
||||
key.Primes = []*big.Int{&p, &q}
|
||||
|
||||
if dp != nil {
|
||||
key.Precomputed.Dp = dp
|
||||
}
|
||||
if dq != nil {
|
||||
key.Precomputed.Dq = dq
|
||||
}
|
||||
if qi != nil {
|
||||
key.Precomputed.Qinv = qi
|
||||
}
|
||||
|
||||
return assignRawResult(v, &key)
|
||||
}
|
||||
|
||||
// Raw takes the values stored in the Key object, and creates the
|
||||
// corresponding *rsa.PublicKey object.
|
||||
func (k *rsaPublicKey) Raw(v interface{}) error {
|
||||
var key rsa.PublicKey
|
||||
|
||||
n := pool.GetBigInt()
|
||||
e := pool.GetBigInt()
|
||||
defer pool.ReleaseBigInt(e)
|
||||
|
||||
n.SetBytes(k.n)
|
||||
e.SetBytes(k.e)
|
||||
|
||||
key.N = n
|
||||
key.E = int(e.Int64())
|
||||
|
||||
return assignRawResult(v, &key)
|
||||
}
|
||||
|
||||
func (k rsaPrivateKey) PublicKey() (RSAPublicKey, error) {
|
||||
var key rsa.PrivateKey
|
||||
if err := k.Raw(&key); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize key to generate public key`)
|
||||
}
|
||||
|
||||
newKey := NewRSAPublicKey()
|
||||
if err := newKey.FromRaw(&key.PublicKey); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to initialize RSAPublicKey`)
|
||||
}
|
||||
return newKey, nil
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (k rsaPrivateKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
var key rsa.PrivateKey
|
||||
if err := k.Raw(&key); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize RSA private key`)
|
||||
}
|
||||
return rsaThumbprint(hash, &key.PublicKey)
|
||||
}
|
||||
|
||||
func (k rsaPublicKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
var key rsa.PublicKey
|
||||
if err := k.Raw(&key); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize RSA public key`)
|
||||
}
|
||||
return rsaThumbprint(hash, &key)
|
||||
}
|
||||
|
||||
func rsaThumbprint(hash crypto.Hash, key *rsa.PublicKey) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
buf.WriteString(`{"e":"`)
|
||||
buf.WriteString(base64.EncodeUint64ToString(uint64(key.E)))
|
||||
buf.WriteString(`","kty":"RSA","n":"`)
|
||||
buf.WriteString(base64.EncodeToString(key.N.Bytes()))
|
||||
buf.WriteString(`"}`)
|
||||
|
||||
h := hash.New()
|
||||
if _, err := buf.WriteTo(h); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to write rsaThumbprint")
|
||||
}
|
||||
return h.Sum(nil), nil
|
||||
}
|
||||
+1057
File diff suppressed because it is too large
Load Diff
+50
@@ -0,0 +1,50 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"fmt"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func NewSymmetricKey() SymmetricKey {
|
||||
return newSymmetricKey()
|
||||
}
|
||||
|
||||
func newSymmetricKey() *symmetricKey {
|
||||
return &symmetricKey{
|
||||
privateParams: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (k *symmetricKey) FromRaw(rawKey []byte) error {
|
||||
if len(rawKey) == 0 {
|
||||
return errors.New(`non-empty []byte key required`)
|
||||
}
|
||||
|
||||
k.octets = rawKey
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Raw returns the octets for this symmetric key.
|
||||
// Since this is a symmetric key, this just calls Octets
|
||||
func (k symmetricKey) Raw(v interface{}) error {
|
||||
return assignRawResult(v, k.octets)
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hakhing algorithm, according to RFC 7638
|
||||
func (k symmetricKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
var octets []byte
|
||||
if err := k.Raw(&octets); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize symmetric key`)
|
||||
}
|
||||
|
||||
h := hash.New()
|
||||
fmt.Fprint(h, `{"k":"`)
|
||||
fmt.Fprint(h, base64.EncodeToString(octets))
|
||||
fmt.Fprint(h, `","kty":"oct"}`)
|
||||
return h.Sum(nil), nil
|
||||
}
|
||||
+396
@@ -0,0 +1,396 @@
|
||||
// This file is auto-generated. DO NOT EDIT
|
||||
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/base64"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
SymmetricOctetsKey = "k"
|
||||
)
|
||||
|
||||
type SymmetricKey interface {
|
||||
Key
|
||||
FromRaw([]byte) error
|
||||
Octets() []byte
|
||||
}
|
||||
|
||||
type symmetricKey struct {
|
||||
algorithm *string // https://tools.ietf.org/html/rfc7517#section-4.4
|
||||
keyID *string // https://tools.ietf.org/html/rfc7515#section-4.1.4
|
||||
keyUsage *string // https://tools.ietf.org/html/rfc7517#section-4.2
|
||||
keyops *KeyOperationList // https://tools.ietf.org/html/rfc7517#section-4.3
|
||||
octets []byte
|
||||
x509CertChain *CertificateChain // https://tools.ietf.org/html/rfc7515#section-4.1.6
|
||||
x509CertThumbprint *string // https://tools.ietf.org/html/rfc7515#section-4.1.7
|
||||
x509CertThumbprintS256 *string // https://tools.ietf.org/html/rfc7515#section-4.1.8
|
||||
x509URL *string // https://tools.ietf.org/html/rfc7515#section-4.1.5
|
||||
privateParams map[string]interface{}
|
||||
}
|
||||
|
||||
type symmetricSymmetricKeyMarshalProxy struct {
|
||||
XkeyType jwa.KeyType `json:"kty"`
|
||||
Xalgorithm *string `json:"alg,omitempty"`
|
||||
XkeyID *string `json:"kid,omitempty"`
|
||||
XkeyUsage *string `json:"use,omitempty"`
|
||||
Xkeyops *KeyOperationList `json:"key_ops,omitempty"`
|
||||
Xoctets *string `json:"k,omitempty"`
|
||||
Xx509CertChain *CertificateChain `json:"x5c,omitempty"`
|
||||
Xx509CertThumbprint *string `json:"x5t,omitempty"`
|
||||
Xx509CertThumbprintS256 *string `json:"x5t#S256,omitempty"`
|
||||
Xx509URL *string `json:"x5u,omitempty"`
|
||||
}
|
||||
|
||||
func (h symmetricKey) KeyType() jwa.KeyType {
|
||||
return jwa.OctetSeq
|
||||
}
|
||||
|
||||
func (h *symmetricKey) Algorithm() string {
|
||||
if h.algorithm != nil {
|
||||
return *(h.algorithm)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *symmetricKey) KeyID() string {
|
||||
if h.keyID != nil {
|
||||
return *(h.keyID)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *symmetricKey) KeyUsage() string {
|
||||
if h.keyUsage != nil {
|
||||
return *(h.keyUsage)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *symmetricKey) KeyOps() KeyOperationList {
|
||||
if h.keyops != nil {
|
||||
return *(h.keyops)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *symmetricKey) Octets() []byte {
|
||||
return h.octets
|
||||
}
|
||||
|
||||
func (h *symmetricKey) X509CertChain() []*x509.Certificate {
|
||||
if h.x509CertChain != nil {
|
||||
return h.x509CertChain.Get()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *symmetricKey) X509CertThumbprint() string {
|
||||
if h.x509CertThumbprint != nil {
|
||||
return *(h.x509CertThumbprint)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *symmetricKey) X509CertThumbprintS256() string {
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
return *(h.x509CertThumbprintS256)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *symmetricKey) X509URL() string {
|
||||
if h.x509URL != nil {
|
||||
return *(h.x509URL)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *symmetricKey) iterate(ctx context.Context, ch chan *HeaderPair) {
|
||||
defer close(ch)
|
||||
|
||||
var pairs []*HeaderPair
|
||||
pairs = append(pairs, &HeaderPair{Key: "kty", Value: jwa.OctetSeq})
|
||||
if h.algorithm != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: AlgorithmKey, Value: *(h.algorithm)})
|
||||
}
|
||||
if h.keyID != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyIDKey, Value: *(h.keyID)})
|
||||
}
|
||||
if h.keyUsage != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyUsageKey, Value: *(h.keyUsage)})
|
||||
}
|
||||
if h.keyops != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyOpsKey, Value: *(h.keyops)})
|
||||
}
|
||||
if h.octets != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: SymmetricOctetsKey, Value: h.octets})
|
||||
}
|
||||
if h.x509CertChain != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertChainKey, Value: *(h.x509CertChain)})
|
||||
}
|
||||
if h.x509CertThumbprint != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintKey, Value: *(h.x509CertThumbprint)})
|
||||
}
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintS256Key, Value: *(h.x509CertThumbprintS256)})
|
||||
}
|
||||
if h.x509URL != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509URLKey, Value: *(h.x509URL)})
|
||||
}
|
||||
for k, v := range h.privateParams {
|
||||
pairs = append(pairs, &HeaderPair{Key: k, Value: v})
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *symmetricKey) PrivateParams() map[string]interface{} {
|
||||
return h.privateParams
|
||||
}
|
||||
|
||||
func (h *symmetricKey) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case KeyTypeKey:
|
||||
return h.KeyType(), true
|
||||
case AlgorithmKey:
|
||||
if h.algorithm == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.algorithm), true
|
||||
case KeyIDKey:
|
||||
if h.keyID == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyID), true
|
||||
case KeyUsageKey:
|
||||
if h.keyUsage == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyUsage), true
|
||||
case KeyOpsKey:
|
||||
if h.keyops == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyops), true
|
||||
case SymmetricOctetsKey:
|
||||
if h.octets == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.octets, true
|
||||
case X509CertChainKey:
|
||||
if h.x509CertChain == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertChain), true
|
||||
case X509CertThumbprintKey:
|
||||
if h.x509CertThumbprint == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprint), true
|
||||
case X509CertThumbprintS256Key:
|
||||
if h.x509CertThumbprintS256 == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprintS256), true
|
||||
case X509URLKey:
|
||||
if h.x509URL == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509URL), true
|
||||
default:
|
||||
v, ok := h.privateParams[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *symmetricKey) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case "kty":
|
||||
return nil
|
||||
case AlgorithmKey:
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
h.algorithm = &v
|
||||
case fmt.Stringer:
|
||||
tmp := v.String()
|
||||
h.algorithm = &tmp
|
||||
default:
|
||||
return errors.Errorf(`invalid type for %s key: %T`, AlgorithmKey, value)
|
||||
}
|
||||
return nil
|
||||
case KeyIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyIDKey, value)
|
||||
case KeyUsageKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyUsage = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyUsageKey, value)
|
||||
case KeyOpsKey:
|
||||
var acceptor KeyOperationList
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, KeyOpsKey)
|
||||
}
|
||||
h.keyops = &acceptor
|
||||
return nil
|
||||
case SymmetricOctetsKey:
|
||||
if v, ok := value.([]byte); ok {
|
||||
h.octets = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, SymmetricOctetsKey, value)
|
||||
case X509CertChainKey:
|
||||
var acceptor CertificateChain
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, X509CertChainKey)
|
||||
}
|
||||
h.x509CertChain = &acceptor
|
||||
return nil
|
||||
case X509CertThumbprintKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprint = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintKey, value)
|
||||
case X509CertThumbprintS256Key:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprintS256 = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintS256Key, value)
|
||||
case X509URLKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509URL = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509URLKey, value)
|
||||
default:
|
||||
if h.privateParams == nil {
|
||||
h.privateParams = map[string]interface{}{}
|
||||
}
|
||||
h.privateParams[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *symmetricKey) UnmarshalJSON(buf []byte) error {
|
||||
var proxy symmetricSymmetricKeyMarshalProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal symmetricKey`)
|
||||
}
|
||||
if proxy.XkeyType != jwa.OctetSeq {
|
||||
return errors.Errorf(`invalid kty value for SymmetricKey (%s)`, proxy.XkeyType)
|
||||
}
|
||||
h.algorithm = proxy.Xalgorithm
|
||||
h.keyID = proxy.XkeyID
|
||||
h.keyUsage = proxy.XkeyUsage
|
||||
h.keyops = proxy.Xkeyops
|
||||
if proxy.Xoctets == nil {
|
||||
return errors.New(`required field k is missing`)
|
||||
}
|
||||
if h.octets = nil; proxy.Xoctets != nil {
|
||||
decoded, err := base64.DecodeString(*(proxy.Xoctets))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to decode base64 value for octets`)
|
||||
}
|
||||
h.octets = decoded
|
||||
}
|
||||
h.x509CertChain = proxy.Xx509CertChain
|
||||
h.x509CertThumbprint = proxy.Xx509CertThumbprint
|
||||
h.x509CertThumbprintS256 = proxy.Xx509CertThumbprintS256
|
||||
h.x509URL = proxy.Xx509URL
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to parse privsate parameters`)
|
||||
}
|
||||
delete(m, `kty`)
|
||||
delete(m, AlgorithmKey)
|
||||
delete(m, KeyIDKey)
|
||||
delete(m, KeyUsageKey)
|
||||
delete(m, KeyOpsKey)
|
||||
delete(m, SymmetricOctetsKey)
|
||||
delete(m, X509CertChainKey)
|
||||
delete(m, X509CertThumbprintKey)
|
||||
delete(m, X509CertThumbprintS256Key)
|
||||
delete(m, X509URLKey)
|
||||
h.privateParams = m
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h symmetricKey) MarshalJSON() ([]byte, error) {
|
||||
var proxy symmetricSymmetricKeyMarshalProxy
|
||||
proxy.XkeyType = jwa.OctetSeq
|
||||
proxy.Xalgorithm = h.algorithm
|
||||
proxy.XkeyID = h.keyID
|
||||
proxy.XkeyUsage = h.keyUsage
|
||||
proxy.Xkeyops = h.keyops
|
||||
if len(h.octets) > 0 {
|
||||
v := base64.EncodeToString(h.octets)
|
||||
proxy.Xoctets = &v
|
||||
}
|
||||
proxy.Xx509CertChain = h.x509CertChain
|
||||
proxy.Xx509CertThumbprint = h.x509CertThumbprint
|
||||
proxy.Xx509CertThumbprintS256 = h.x509CertThumbprintS256
|
||||
proxy.Xx509URL = h.x509URL
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
if err := enc.Encode(proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to encode proxy to JSON`)
|
||||
}
|
||||
hasContent := buf.Len() > 3 // encoding/json always adds a newline, so "{}\n" is the empty hash
|
||||
if l := len(h.privateParams); l > 0 {
|
||||
buf.Truncate(buf.Len() - 2)
|
||||
keys := make([]string, 0, l)
|
||||
for k := range h.privateParams {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for i, k := range keys {
|
||||
if hasContent || i > 0 {
|
||||
fmt.Fprintf(&buf, `,`)
|
||||
}
|
||||
fmt.Fprintf(&buf, `%s:`, strconv.Quote(k))
|
||||
if err := enc.Encode(h.privateParams[k]); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to encode private param %s`, k)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(&buf, `}`)
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func (h *symmetricKey) Iterate(ctx context.Context) HeaderIterator {
|
||||
ch := make(chan *HeaderPair)
|
||||
go h.iterate(ctx, ch)
|
||||
return mapiter.New(ch)
|
||||
}
|
||||
|
||||
func (h *symmetricKey) Walk(ctx context.Context, visitor HeaderVisitor) error {
|
||||
return iter.WalkMap(ctx, h, visitor)
|
||||
}
|
||||
|
||||
func (h *symmetricKey) AsMap(ctx context.Context) (map[string]interface{}, error) {
|
||||
return iter.AsMap(ctx, h)
|
||||
}
|
||||
+24
@@ -0,0 +1,24 @@
|
||||
package jws
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
)
|
||||
|
||||
// Iterate returns a channel that successively returns all the
|
||||
// header name and values.
|
||||
func (h *stdHeaders) Iterate(ctx context.Context) Iterator {
|
||||
ch := make(chan *HeaderPair)
|
||||
go h.iterate(ctx, ch)
|
||||
return mapiter.New(ch)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) Walk(ctx context.Context, visitor Visitor) error {
|
||||
return iter.WalkMap(ctx, h, visitor)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) AsMap(ctx context.Context) (map[string]interface{}, error) {
|
||||
return iter.AsMap(ctx, h)
|
||||
}
|
||||
+432
@@ -0,0 +1,432 @@
|
||||
// This file is auto-generated. DO NOT EDIT
|
||||
package jws
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jwk"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
AlgorithmKey = "alg"
|
||||
ContentTypeKey = "cty"
|
||||
CriticalKey = "crit"
|
||||
JWKKey = "jwk"
|
||||
JWKSetURLKey = "jku"
|
||||
KeyIDKey = "kid"
|
||||
TypeKey = "typ"
|
||||
X509CertChainKey = "x5c"
|
||||
X509CertThumbprintKey = "x5t"
|
||||
X509CertThumbprintS256Key = "x5t#S256"
|
||||
X509URLKey = "x5u"
|
||||
)
|
||||
|
||||
// Headers describe a standard Header set.
|
||||
type Headers interface {
|
||||
Algorithm() jwa.SignatureAlgorithm
|
||||
ContentType() string
|
||||
Critical() []string
|
||||
JWK() jwk.Key
|
||||
JWKSetURL() string
|
||||
KeyID() string
|
||||
Type() string
|
||||
X509CertChain() []string
|
||||
X509CertThumbprint() string
|
||||
X509CertThumbprintS256() string
|
||||
X509URL() string
|
||||
Iterate(ctx context.Context) Iterator
|
||||
Walk(ctx context.Context, v Visitor) error
|
||||
AsMap(ctx context.Context) (map[string]interface{}, error)
|
||||
Get(string) (interface{}, bool)
|
||||
Set(string, interface{}) error
|
||||
}
|
||||
|
||||
type stdHeaders struct {
|
||||
algorithm *jwa.SignatureAlgorithm `json:"alg,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.1
|
||||
contentType *string `json:"cty,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.10
|
||||
critical []string `json:"crit,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.11
|
||||
jwk jwk.Key `json:"jwk,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.3
|
||||
jwkSetURL *string `json:"jku,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.2
|
||||
keyID *string `json:"kid,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.4
|
||||
typ *string `json:"typ,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.9
|
||||
x509CertChain []string `json:"x5c,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.6
|
||||
x509CertThumbprint *string `json:"x5t,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.7
|
||||
x509CertThumbprintS256 *string `json:"x5t#S256,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.8
|
||||
x509URL *string `json:"x5u,omitempty"` // https://tools.ietf.org/html/rfc7515#section-4.1.5
|
||||
privateParams map[string]interface{}
|
||||
}
|
||||
|
||||
type standardHeadersMarshalProxy struct {
|
||||
Xalgorithm *jwa.SignatureAlgorithm `json:"alg,omitempty"`
|
||||
XcontentType *string `json:"cty,omitempty"`
|
||||
Xcritical []string `json:"crit,omitempty"`
|
||||
Xjwk json.RawMessage `json:"jwk,omitempty"`
|
||||
XjwkSetURL *string `json:"jku,omitempty"`
|
||||
XkeyID *string `json:"kid,omitempty"`
|
||||
Xtyp *string `json:"typ,omitempty"`
|
||||
Xx509CertChain []string `json:"x5c,omitempty"`
|
||||
Xx509CertThumbprint *string `json:"x5t,omitempty"`
|
||||
Xx509CertThumbprintS256 *string `json:"x5t#S256,omitempty"`
|
||||
Xx509URL *string `json:"x5u,omitempty"`
|
||||
}
|
||||
|
||||
func NewHeaders() Headers {
|
||||
return &stdHeaders{}
|
||||
}
|
||||
|
||||
func (h *stdHeaders) Algorithm() jwa.SignatureAlgorithm {
|
||||
if h.algorithm == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.algorithm)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) ContentType() string {
|
||||
if h.contentType == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.contentType)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) Critical() []string {
|
||||
return h.critical
|
||||
}
|
||||
|
||||
func (h *stdHeaders) JWK() jwk.Key {
|
||||
return h.jwk
|
||||
}
|
||||
|
||||
func (h *stdHeaders) JWKSetURL() string {
|
||||
if h.jwkSetURL == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.jwkSetURL)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) KeyID() string {
|
||||
if h.keyID == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.keyID)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) Type() string {
|
||||
if h.typ == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.typ)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) X509CertChain() []string {
|
||||
return h.x509CertChain
|
||||
}
|
||||
|
||||
func (h *stdHeaders) X509CertThumbprint() string {
|
||||
if h.x509CertThumbprint == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.x509CertThumbprint)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) X509CertThumbprintS256() string {
|
||||
if h.x509CertThumbprintS256 == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.x509CertThumbprintS256)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) X509URL() string {
|
||||
if h.x509URL == nil {
|
||||
return ""
|
||||
}
|
||||
return *(h.x509URL)
|
||||
}
|
||||
|
||||
func (h *stdHeaders) iterate(ctx context.Context, ch chan *HeaderPair) {
|
||||
defer close(ch)
|
||||
var pairs []*HeaderPair
|
||||
if h.algorithm != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: AlgorithmKey, Value: *(h.algorithm)})
|
||||
}
|
||||
if h.contentType != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: ContentTypeKey, Value: *(h.contentType)})
|
||||
}
|
||||
if h.critical != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: CriticalKey, Value: h.critical})
|
||||
}
|
||||
if h.jwk != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: JWKKey, Value: h.jwk})
|
||||
}
|
||||
if h.jwkSetURL != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: JWKSetURLKey, Value: *(h.jwkSetURL)})
|
||||
}
|
||||
if h.keyID != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: KeyIDKey, Value: *(h.keyID)})
|
||||
}
|
||||
if h.typ != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: TypeKey, Value: *(h.typ)})
|
||||
}
|
||||
if h.x509CertChain != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertChainKey, Value: h.x509CertChain})
|
||||
}
|
||||
if h.x509CertThumbprint != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintKey, Value: *(h.x509CertThumbprint)})
|
||||
}
|
||||
if h.x509CertThumbprintS256 != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509CertThumbprintS256Key, Value: *(h.x509CertThumbprintS256)})
|
||||
}
|
||||
if h.x509URL != nil {
|
||||
pairs = append(pairs, &HeaderPair{Key: X509URLKey, Value: *(h.x509URL)})
|
||||
}
|
||||
for k, v := range h.privateParams {
|
||||
pairs = append(pairs, &HeaderPair{Key: k, Value: v})
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *stdHeaders) PrivateParams() map[string]interface{} {
|
||||
return h.privateParams
|
||||
}
|
||||
|
||||
func (h *stdHeaders) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case AlgorithmKey:
|
||||
if h.algorithm == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.algorithm), true
|
||||
case ContentTypeKey:
|
||||
if h.contentType == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.contentType), true
|
||||
case CriticalKey:
|
||||
if h.critical == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.critical, true
|
||||
case JWKKey:
|
||||
if h.jwk == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.jwk, true
|
||||
case JWKSetURLKey:
|
||||
if h.jwkSetURL == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.jwkSetURL), true
|
||||
case KeyIDKey:
|
||||
if h.keyID == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.keyID), true
|
||||
case TypeKey:
|
||||
if h.typ == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.typ), true
|
||||
case X509CertChainKey:
|
||||
if h.x509CertChain == nil {
|
||||
return nil, false
|
||||
}
|
||||
return h.x509CertChain, true
|
||||
case X509CertThumbprintKey:
|
||||
if h.x509CertThumbprint == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprint), true
|
||||
case X509CertThumbprintS256Key:
|
||||
if h.x509CertThumbprintS256 == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509CertThumbprintS256), true
|
||||
case X509URLKey:
|
||||
if h.x509URL == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(h.x509URL), true
|
||||
default:
|
||||
v, ok := h.privateParams[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *stdHeaders) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case AlgorithmKey:
|
||||
var acceptor jwa.SignatureAlgorithm
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, AlgorithmKey)
|
||||
}
|
||||
h.algorithm = &acceptor
|
||||
return nil
|
||||
case ContentTypeKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.contentType = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ContentTypeKey, value)
|
||||
case CriticalKey:
|
||||
if v, ok := value.([]string); ok {
|
||||
h.critical = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, CriticalKey, value)
|
||||
case JWKKey:
|
||||
if v, ok := value.(jwk.Key); ok {
|
||||
h.jwk = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, JWKKey, value)
|
||||
case JWKSetURLKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.jwkSetURL = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, JWKSetURLKey, value)
|
||||
case KeyIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyIDKey, value)
|
||||
case TypeKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.typ = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, TypeKey, value)
|
||||
case X509CertChainKey:
|
||||
if v, ok := value.([]string); ok {
|
||||
h.x509CertChain = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertChainKey, value)
|
||||
case X509CertThumbprintKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprint = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintKey, value)
|
||||
case X509CertThumbprintS256Key:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprintS256 = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintS256Key, value)
|
||||
case X509URLKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509URL = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509URLKey, value)
|
||||
default:
|
||||
if h.privateParams == nil {
|
||||
h.privateParams = map[string]interface{}{}
|
||||
}
|
||||
h.privateParams[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdHeaders) UnmarshalJSON(buf []byte) error {
|
||||
var proxy standardHeadersMarshalProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal headers`)
|
||||
}
|
||||
|
||||
h.jwk = nil
|
||||
if jwkField := proxy.Xjwk; len(jwkField) > 0 {
|
||||
set, err := jwk.ParseBytes([]byte(proxy.Xjwk))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to parse jwk field`)
|
||||
}
|
||||
h.jwk = set.Keys[0]
|
||||
}
|
||||
h.algorithm = proxy.Xalgorithm
|
||||
h.contentType = proxy.XcontentType
|
||||
h.critical = proxy.Xcritical
|
||||
h.jwkSetURL = proxy.XjwkSetURL
|
||||
h.keyID = proxy.XkeyID
|
||||
h.typ = proxy.Xtyp
|
||||
h.x509CertChain = proxy.Xx509CertChain
|
||||
h.x509CertThumbprint = proxy.Xx509CertThumbprint
|
||||
h.x509CertThumbprintS256 = proxy.Xx509CertThumbprintS256
|
||||
h.x509URL = proxy.Xx509URL
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to parse privsate parameters`)
|
||||
}
|
||||
delete(m, AlgorithmKey)
|
||||
delete(m, ContentTypeKey)
|
||||
delete(m, CriticalKey)
|
||||
delete(m, JWKKey)
|
||||
delete(m, JWKSetURLKey)
|
||||
delete(m, KeyIDKey)
|
||||
delete(m, TypeKey)
|
||||
delete(m, X509CertChainKey)
|
||||
delete(m, X509CertThumbprintKey)
|
||||
delete(m, X509CertThumbprintS256Key)
|
||||
delete(m, X509URLKey)
|
||||
h.privateParams = m
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h stdHeaders) MarshalJSON() ([]byte, error) {
|
||||
var proxy standardHeadersMarshalProxy
|
||||
if h.jwk != nil {
|
||||
jwkbuf, err := json.Marshal(h.jwk)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to marshal jwk field`)
|
||||
}
|
||||
proxy.Xjwk = jwkbuf
|
||||
}
|
||||
proxy.Xalgorithm = h.algorithm
|
||||
proxy.XcontentType = h.contentType
|
||||
proxy.Xcritical = h.critical
|
||||
proxy.XjwkSetURL = h.jwkSetURL
|
||||
proxy.XkeyID = h.keyID
|
||||
proxy.Xtyp = h.typ
|
||||
proxy.Xx509CertChain = h.x509CertChain
|
||||
proxy.Xx509CertThumbprint = h.x509CertThumbprint
|
||||
proxy.Xx509CertThumbprintS256 = h.x509CertThumbprintS256
|
||||
proxy.Xx509URL = h.x509URL
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
if err := enc.Encode(proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to encode proxy to JSON`)
|
||||
}
|
||||
hasContent := buf.Len() > 3 // encoding/json always adds a newline, so "{}\n" is the empty hash
|
||||
if l := len(h.privateParams); l > 0 {
|
||||
buf.Truncate(buf.Len() - 2)
|
||||
keys := make([]string, 0, l)
|
||||
for k := range h.privateParams {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for i, k := range keys {
|
||||
if hasContent || i > 0 {
|
||||
fmt.Fprintf(&buf, `,`)
|
||||
}
|
||||
fmt.Fprintf(&buf, `%s:`, strconv.Quote(k))
|
||||
if err := enc.Encode(h.privateParams[k]); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to encode private param %s`, k)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(&buf, `}`)
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
+77
@@ -0,0 +1,77 @@
|
||||
package jws
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jwk"
|
||||
)
|
||||
|
||||
type encodedSignature struct {
|
||||
Protected string `json:"protected,omitempty"`
|
||||
Headers Headers `json:"header,omitempty"`
|
||||
Signature string `json:"signature,omitempty"`
|
||||
}
|
||||
|
||||
type encodedMessage struct {
|
||||
Payload string `json:"payload"`
|
||||
Signatures []*encodedSignature `json:"signatures,omitempty"`
|
||||
}
|
||||
|
||||
// PayloadSigner generates signature for the given payload
|
||||
type PayloadSigner interface {
|
||||
Sign([]byte) ([]byte, error)
|
||||
Algorithm() jwa.SignatureAlgorithm
|
||||
ProtectedHeader() Headers
|
||||
PublicHeader() Headers
|
||||
}
|
||||
|
||||
// Message represents a full JWS encoded message. Flattened serialization
|
||||
// is not supported as a struct, but rather it's represented as a
|
||||
// Message struct with only one `signature` element.
|
||||
//
|
||||
// Do not expect to use the Message object to verify or construct a
|
||||
// signed payloads with. You should only use this when you want to actually
|
||||
// want to programmatically view the contents for the full JWS payload.
|
||||
//
|
||||
// To sign and verify, use the appropriate `Sign()` nad `Verify()` functions
|
||||
type Message struct {
|
||||
payload []byte
|
||||
signatures []*Signature
|
||||
}
|
||||
|
||||
type Signature struct {
|
||||
headers Headers // Unprotected Headers
|
||||
protected Headers // Protected Headers
|
||||
signature []byte // Signature
|
||||
}
|
||||
|
||||
// JWKAcceptor decides which keys can be accepted
|
||||
// by functions that iterate over a JWK key set.
|
||||
type JWKAcceptor interface {
|
||||
Accept(jwk.Key) bool
|
||||
}
|
||||
|
||||
// JWKAcceptFunc is an implementation of JWKAcceptor
|
||||
// using a plain function
|
||||
type JWKAcceptFunc func(jwk.Key) bool
|
||||
|
||||
// Accept executes the provided function to determine if the
|
||||
// given key can be used
|
||||
func (f JWKAcceptFunc) Accept(key jwk.Key) bool {
|
||||
return f(key)
|
||||
}
|
||||
|
||||
// DefaultJWKAcceptor is the default acceptor that is used
|
||||
// in functions like VerifyWithJWKSet
|
||||
var DefaultJWKAcceptor = JWKAcceptFunc(func(key jwk.Key) bool {
|
||||
if u := key.KeyUsage(); u != "" && u != "enc" && u != "sig" {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
type Visitor = iter.MapVisitor
|
||||
type VisitorFunc = iter.MapVisitorFunc
|
||||
type HeaderPair = mapiter.Pair
|
||||
type Iterator = mapiter.Iterator
|
||||
+609
@@ -0,0 +1,609 @@
|
||||
//go:generate go run internal/cmd/genheader/main.go
|
||||
|
||||
// Package jws implements the digital signature on JSON based data
|
||||
// structures as described in https://tools.ietf.org/html/rfc7515
|
||||
//
|
||||
// If you do not care about the details, the only things that you
|
||||
// would need to use are the following functions:
|
||||
//
|
||||
// jws.Sign(payload, algorithm, key)
|
||||
// jws.Verify(encodedjws, algorithm, key)
|
||||
//
|
||||
// To sign, simply use `jws.Sign`. `payload` is a []byte buffer that
|
||||
// contains whatever data you want to sign. `alg` is one of the
|
||||
// jwa.SignatureAlgorithm constants from package jwa. For RSA and
|
||||
// ECDSA family of algorithms, you will need to prepare a private key.
|
||||
// For HMAC family, you just need a []byte value. The `jws.Sign`
|
||||
// function will return the encoded JWS message on success.
|
||||
//
|
||||
// To verify, use `jws.Verify`. It will parse the `encodedjws` buffer
|
||||
// and verify the result using `algorithm` and `key`. Upon successful
|
||||
// verification, the original payload is returned, so you can work on it.
|
||||
package jws
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/pool"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jwk"
|
||||
"github.com/lestrrat-go/jwx/jws/sign"
|
||||
"github.com/lestrrat-go/jwx/jws/verify"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type payloadSigner struct {
|
||||
signer sign.Signer
|
||||
key interface{}
|
||||
protected Headers
|
||||
public Headers
|
||||
}
|
||||
|
||||
func (s *payloadSigner) Sign(payload []byte) ([]byte, error) {
|
||||
return s.signer.Sign(payload, s.key)
|
||||
}
|
||||
|
||||
func (s *payloadSigner) Algorithm() jwa.SignatureAlgorithm {
|
||||
return s.signer.Algorithm()
|
||||
}
|
||||
|
||||
func (s *payloadSigner) ProtectedHeader() Headers {
|
||||
return s.protected
|
||||
}
|
||||
|
||||
func (s *payloadSigner) PublicHeader() Headers {
|
||||
return s.public
|
||||
}
|
||||
|
||||
// Sign generates a signature for the given payload, and serializes
|
||||
// it in compact serialization format. In this format you may NOT use
|
||||
// multiple signers.
|
||||
//
|
||||
// If you would like to pass custom headers, use the WithHeaders option.
|
||||
func Sign(payload []byte, alg jwa.SignatureAlgorithm, key interface{}, options ...Option) ([]byte, error) {
|
||||
var hdrs Headers = NewHeaders()
|
||||
for _, o := range options {
|
||||
switch o.Name() {
|
||||
case optkeyHeaders:
|
||||
hdrs = o.Value().(Headers)
|
||||
}
|
||||
}
|
||||
|
||||
signer, err := sign.New(alg)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to create signer`)
|
||||
}
|
||||
|
||||
if err := hdrs.Set(AlgorithmKey, signer.Algorithm()); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to set header`)
|
||||
}
|
||||
|
||||
hdrbuf, err := json.Marshal(hdrs)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to marshal headers`)
|
||||
}
|
||||
|
||||
buf := pool.GetBytesBuffer()
|
||||
defer pool.ReleaseBytesBuffer(buf)
|
||||
enc := base64.NewEncoder(base64.RawURLEncoding, buf)
|
||||
if _, err := enc.Write(hdrbuf); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to write headers as base64`)
|
||||
}
|
||||
if err := enc.Close(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to finalize writing headers as base64`)
|
||||
}
|
||||
|
||||
buf.WriteByte('.')
|
||||
enc = base64.NewEncoder(base64.RawURLEncoding, buf)
|
||||
if _, err := enc.Write(payload); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to write payload as base64`)
|
||||
}
|
||||
if err := enc.Close(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to finalize writing payload as base64`)
|
||||
}
|
||||
|
||||
signature, err := signer.Sign(buf.Bytes(), key)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to sign payload`)
|
||||
}
|
||||
|
||||
buf.WriteByte('.')
|
||||
enc = base64.NewEncoder(base64.RawURLEncoding, buf)
|
||||
if _, err := enc.Write(signature); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to write signature as base64`)
|
||||
}
|
||||
if err := enc.Close(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to finalize writing signature as base64`)
|
||||
}
|
||||
|
||||
result := make([]byte, buf.Len())
|
||||
copy(result, buf.Bytes())
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// SignLiteral generates a signature for the given payload and headers, and serializes
|
||||
// it in compact serialization format. In this format you may NOT use
|
||||
// multiple signers.
|
||||
//
|
||||
func SignLiteral(payload []byte, alg jwa.SignatureAlgorithm, key interface{}, headers []byte) ([]byte, error) {
|
||||
signer, err := sign.New(alg)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to create signer`)
|
||||
}
|
||||
|
||||
buf := pool.GetBytesBuffer()
|
||||
defer pool.ReleaseBytesBuffer(buf)
|
||||
|
||||
enc := base64.NewEncoder(base64.RawURLEncoding, buf)
|
||||
if _, err := enc.Write(headers); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to write headers as base64`)
|
||||
}
|
||||
if err := enc.Close(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to finalize writing headers as base64`)
|
||||
}
|
||||
|
||||
buf.WriteByte('.')
|
||||
enc = base64.NewEncoder(base64.RawURLEncoding, buf)
|
||||
if _, err := enc.Write(payload); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to write payload as base64`)
|
||||
}
|
||||
if err := enc.Close(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to finalize writing payload as base64`)
|
||||
}
|
||||
|
||||
signature, err := signer.Sign(buf.Bytes(), key)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to sign payload`)
|
||||
}
|
||||
|
||||
buf.WriteByte('.')
|
||||
enc = base64.NewEncoder(base64.RawURLEncoding, buf)
|
||||
if _, err := enc.Write(signature); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to write signature as base64`)
|
||||
}
|
||||
if err := enc.Close(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to finalize writing signature as base64`)
|
||||
}
|
||||
|
||||
result := make([]byte, buf.Len())
|
||||
copy(result, buf.Bytes())
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// SignMulti accepts multiple signers via the options parameter,
|
||||
// and creates a JWS in JSON serialization format that contains
|
||||
// signatures from applying aforementioned signers.
|
||||
func SignMulti(payload []byte, options ...Option) ([]byte, error) {
|
||||
var signers []PayloadSigner
|
||||
for _, o := range options {
|
||||
switch o.Name() {
|
||||
case optkeyPayloadSigner:
|
||||
signers = append(signers, o.Value().(PayloadSigner))
|
||||
}
|
||||
}
|
||||
|
||||
if len(signers) == 0 {
|
||||
return nil, errors.New(`no signers provided`)
|
||||
}
|
||||
|
||||
var result encodedMessage
|
||||
|
||||
result.Payload = base64.RawURLEncoding.EncodeToString(payload)
|
||||
|
||||
buf := pool.GetBytesBuffer()
|
||||
defer pool.ReleaseBytesBuffer(buf)
|
||||
for _, signer := range signers {
|
||||
protected := signer.ProtectedHeader()
|
||||
if protected == nil {
|
||||
protected = NewHeaders()
|
||||
}
|
||||
|
||||
if err := protected.Set(AlgorithmKey, signer.Algorithm()); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to set header`)
|
||||
}
|
||||
|
||||
hdrbuf, err := json.Marshal(protected)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to marshal headers`)
|
||||
}
|
||||
encodedHeader := base64.RawURLEncoding.EncodeToString(hdrbuf)
|
||||
|
||||
buf.Reset()
|
||||
buf.WriteString(encodedHeader)
|
||||
buf.WriteByte('.')
|
||||
buf.WriteString(result.Payload)
|
||||
signature, err := signer.Sign(buf.Bytes())
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to sign payload`)
|
||||
}
|
||||
|
||||
result.Signatures = append(result.Signatures, &encodedSignature{
|
||||
Headers: signer.PublicHeader(),
|
||||
Protected: encodedHeader,
|
||||
Signature: base64.RawURLEncoding.EncodeToString(signature),
|
||||
})
|
||||
}
|
||||
|
||||
return json.Marshal(result)
|
||||
}
|
||||
|
||||
// Verify checks if the given JWS message is verifiable using `alg` and `key`.
|
||||
// If the verification is successful, `err` is nil, and the content of the
|
||||
// payload that was signed is returned. If you need more fine-grained
|
||||
// control of the verification process, manually call `Parse`, generate a
|
||||
// verifier, and call `Verify` on the parsed JWS message object.
|
||||
func Verify(buf []byte, alg jwa.SignatureAlgorithm, key interface{}) (ret []byte, err error) {
|
||||
verifier, err := verify.New(alg)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to create verifier")
|
||||
}
|
||||
|
||||
buf = bytes.TrimSpace(buf)
|
||||
if len(buf) == 0 {
|
||||
return nil, errors.New(`attempt to verify empty buffer`)
|
||||
}
|
||||
|
||||
if buf[0] == '{' {
|
||||
// FUuuuuuuuuuuuuuuuck // WTF am I doing here.
|
||||
var proxy fullMessageProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unmarshal JWS message`)
|
||||
}
|
||||
|
||||
// There's something wrong if the Message part is not initialized
|
||||
if len(proxy.Payload) == 0 {
|
||||
return nil, errors.New(`invalid JWS message format (missing payload)`)
|
||||
}
|
||||
|
||||
// if we're using the compact serialization format, then m.Signature
|
||||
// will be non-nil
|
||||
if len(proxy.Signature) > 0 {
|
||||
if len(proxy.Signatures) > 0 {
|
||||
return nil, errors.New(`invalid JWS message format (signature and signatures both exist)`)
|
||||
}
|
||||
encodedSig, err := proxy.encodedSignature()
|
||||
if err != nil {
|
||||
return nil, err // don't think we need to wrap this one
|
||||
}
|
||||
|
||||
proxy.Signatures = append(proxy.Signatures, encodedSig)
|
||||
}
|
||||
|
||||
buf := pool.GetBytesBuffer()
|
||||
defer pool.ReleaseBytesBuffer(buf)
|
||||
for _, sig := range proxy.Signatures {
|
||||
buf.Reset()
|
||||
buf.WriteString(sig.Protected)
|
||||
buf.WriteByte('.')
|
||||
buf.WriteString(proxy.Payload)
|
||||
decodedSignature, err := base64.RawURLEncoding.DecodeString(sig.Signature)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
if err := verifier.Verify(buf.Bytes(), decodedSignature, key); err == nil {
|
||||
// verified!
|
||||
decodedPayload, err := base64.RawURLEncoding.DecodeString(proxy.Payload)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `message verified, failed to decode payload`)
|
||||
}
|
||||
return decodedPayload, nil
|
||||
}
|
||||
}
|
||||
return nil, errors.New(`could not verify with any of the signatures`)
|
||||
}
|
||||
|
||||
protected, payload, signature, err := SplitCompact(bytes.NewReader(buf))
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed extract from compact serialization format`)
|
||||
}
|
||||
|
||||
verifyBuf := pool.GetBytesBuffer()
|
||||
defer pool.ReleaseBytesBuffer(verifyBuf)
|
||||
|
||||
verifyBuf.Write(protected)
|
||||
verifyBuf.WriteByte('.')
|
||||
verifyBuf.Write(payload)
|
||||
|
||||
decodedSignature := make([]byte, base64.RawURLEncoding.DecodedLen(len(signature)))
|
||||
if _, err := base64.RawURLEncoding.Decode(decodedSignature, signature); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to decode signature`)
|
||||
}
|
||||
if err := verifier.Verify(verifyBuf.Bytes(), decodedSignature, key); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to verify message`)
|
||||
}
|
||||
|
||||
decodedPayload := make([]byte, base64.RawURLEncoding.DecodedLen(len(payload)))
|
||||
if _, err := base64.RawURLEncoding.Decode(decodedPayload, payload); err != nil {
|
||||
return nil, errors.Wrap(err, `message verified, failed to decode payload`)
|
||||
}
|
||||
return decodedPayload, nil
|
||||
}
|
||||
|
||||
// VerifyWithJKU wraps VerifyWithJKUAndContext using the background context.
|
||||
func VerifyWithJKU(buf []byte, jwkurl string, options ...Option) ([]byte, error) {
|
||||
return VerifyWithJKUAndContext(context.Background(), buf, jwkurl, options...)
|
||||
}
|
||||
|
||||
// VerifyWithJKUAndContext verifies the JWS message using a remote JWK
|
||||
// file represented in the url.
|
||||
func VerifyWithJKUAndContext(ctx context.Context, buf []byte, jwkurl string, options ...Option) ([]byte, error) {
|
||||
key, err := jwk.FetchHTTPWithContext(ctx, jwkurl, options...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to fetch jwk via HTTP`)
|
||||
}
|
||||
|
||||
return VerifyWithJWKSet(buf, key, nil)
|
||||
}
|
||||
|
||||
// VerifyWithJWK verifies the JWS message using the specified JWK
|
||||
func VerifyWithJWK(buf []byte, key jwk.Key) (payload []byte, err error) {
|
||||
var rawkey interface{}
|
||||
if err := key.Raw(&rawkey); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to materialize jwk.Key`)
|
||||
}
|
||||
|
||||
payload, err = Verify(buf, jwa.SignatureAlgorithm(key.Algorithm()), rawkey)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to verify message")
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
// VerifyWithJWKSet verifies the JWS message using JWK key set.
|
||||
// By default it will only pick up keys that have the "use" key
|
||||
// set to either "sig" or "enc", but you can override it by
|
||||
// providing a keyaccept function.
|
||||
func VerifyWithJWKSet(buf []byte, keyset *jwk.Set, keyaccept JWKAcceptFunc) ([]byte, error) {
|
||||
if keyaccept == nil {
|
||||
keyaccept = DefaultJWKAcceptor
|
||||
}
|
||||
|
||||
for _, key := range keyset.Keys {
|
||||
if !keyaccept(key) {
|
||||
continue
|
||||
}
|
||||
|
||||
payload, err := VerifyWithJWK(buf, key)
|
||||
if err == nil {
|
||||
return payload, nil
|
||||
}
|
||||
}
|
||||
|
||||
// refs #140, #141
|
||||
//
|
||||
// We should not be Wrap()'ing the error here, because of various
|
||||
// reasons -- but the fundamental one is that the only value we can get
|
||||
// here is the "last error" seen in the above loop, when the symptom
|
||||
// that we want to report is that none of the keys worked.
|
||||
//
|
||||
// Here, we just return that fact, and we do not rely on the value of
|
||||
// previous errors.
|
||||
return nil, errors.New("failed to verify with any of the keys")
|
||||
}
|
||||
|
||||
// Parse parses contents from the given source and creates a jws.Message
|
||||
// struct. The input can be in either compact or full JSON serialization.
|
||||
func Parse(src io.Reader) (m *Message, err error) {
|
||||
rdr := bufio.NewReader(src)
|
||||
var first rune
|
||||
for {
|
||||
r, _, err := rdr.ReadRune()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to read rune`)
|
||||
}
|
||||
if !unicode.IsSpace(r) {
|
||||
first = r
|
||||
if err := rdr.UnreadRune(); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unread rune`)
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
var parser func(io.Reader) (*Message, error)
|
||||
if first == '{' {
|
||||
parser = parseJSON
|
||||
} else {
|
||||
parser = parseCompact
|
||||
}
|
||||
|
||||
m, err = parser(rdr)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to parse jws message`)
|
||||
}
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// ParseString is the same as Parse, but take in a string
|
||||
func ParseString(s string) (*Message, error) {
|
||||
return Parse(strings.NewReader(s))
|
||||
}
|
||||
|
||||
type fullMessageProxy struct {
|
||||
// encoded signature fields
|
||||
Signature json.RawMessage `json:"signature"`
|
||||
Headers json.RawMessage `json:"header"` // あれ、"s"いらないんだっけ
|
||||
Protected json.RawMessage `json:"protected"`
|
||||
|
||||
// encoded message fields
|
||||
Signatures []*encodedSignature `json:"signatures"`
|
||||
Payload string `json:"payload"`
|
||||
}
|
||||
|
||||
func (proxy *fullMessageProxy) encodedSignature() (*encodedSignature, error) {
|
||||
var encodedSig encodedSignature
|
||||
if err := json.Unmarshal(proxy.Protected, &encodedSig.Protected); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unmarshal 'protected' field`)
|
||||
}
|
||||
if err := json.Unmarshal(proxy.Signature, &encodedSig.Signature); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unmarshal 'signature' field`)
|
||||
}
|
||||
h := NewHeaders()
|
||||
if err := json.Unmarshal(proxy.Headers, h); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unmarshal 'header' field`)
|
||||
}
|
||||
|
||||
return &encodedSig, nil
|
||||
}
|
||||
|
||||
func parseJSON(src io.Reader) (result *Message, err error) {
|
||||
var proxy fullMessageProxy
|
||||
if err := json.NewDecoder(src).Decode(&proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to unmarshal jws message`)
|
||||
}
|
||||
|
||||
if len(proxy.Signature) > 0 {
|
||||
if len(proxy.Signatures) > 0 {
|
||||
return nil, errors.New("invalid message: mixed compact/full json serialization")
|
||||
}
|
||||
|
||||
encodedSig, err := proxy.encodedSignature()
|
||||
if err != nil {
|
||||
return nil, err // don't think we need to wrap this one
|
||||
}
|
||||
proxy.Signatures = append(proxy.Signatures, encodedSig)
|
||||
}
|
||||
|
||||
var plain Message
|
||||
plain.payload, err = base64.RawURLEncoding.DecodeString(proxy.Payload)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to decode payload`)
|
||||
}
|
||||
|
||||
for i, sig := range proxy.Signatures {
|
||||
var plainSig Signature
|
||||
|
||||
plainSig.headers = sig.Headers
|
||||
|
||||
if l := len(sig.Protected); l > 0 {
|
||||
plainSig.protected = NewHeaders()
|
||||
hdrbuf, err := base64.RawURLEncoding.DecodeString(sig.Protected)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to base64 decode protected header for signature #%d`, i+1)
|
||||
}
|
||||
if err := json.Unmarshal(hdrbuf, &plainSig.protected); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to unmarshal protected header for signature #%d`, i+1)
|
||||
}
|
||||
}
|
||||
|
||||
plainSig.signature, err = base64.RawURLEncoding.DecodeString(sig.Signature)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to decode signature #%d`, i)
|
||||
}
|
||||
|
||||
plain.signatures = append(plain.signatures, &plainSig)
|
||||
}
|
||||
|
||||
return &plain, nil
|
||||
}
|
||||
|
||||
// SplitCompact splits a JWT and returns its three parts
|
||||
// separately: protected headers, payload and signature.
|
||||
func SplitCompact(rdr io.Reader) ([]byte, []byte, []byte, error) {
|
||||
var protected []byte
|
||||
var payload []byte
|
||||
var signature []byte
|
||||
var periods int = 0
|
||||
var state int = 0
|
||||
|
||||
buf := make([]byte, 4096)
|
||||
var sofar []byte
|
||||
|
||||
for {
|
||||
// read next bytes
|
||||
n, err := rdr.Read(buf)
|
||||
// return on unexpected read error
|
||||
if err != nil && err != io.EOF {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
// append to current buffer
|
||||
sofar = append(sofar, buf[:n]...)
|
||||
// loop to capture multiple '.' in current buffer
|
||||
for loop := true; loop; {
|
||||
var i = bytes.IndexByte(sofar, '.')
|
||||
if i == -1 && err != io.EOF {
|
||||
// no '.' found -> exit and read next bytes (outer loop)
|
||||
loop = false
|
||||
continue
|
||||
} else if i == -1 && err == io.EOF {
|
||||
// no '.' found -> process rest and exit
|
||||
i = len(sofar)
|
||||
loop = false
|
||||
} else {
|
||||
// '.' found
|
||||
periods++
|
||||
}
|
||||
|
||||
// Reaching this point means we have found a '.' or EOF and process the rest of the buffer
|
||||
switch state {
|
||||
case 0:
|
||||
protected = sofar[:i]
|
||||
state++
|
||||
case 1:
|
||||
payload = sofar[:i]
|
||||
state++
|
||||
case 2:
|
||||
signature = sofar[:i]
|
||||
}
|
||||
// Shorten current buffer
|
||||
if len(sofar) > i {
|
||||
sofar = sofar[i+1:]
|
||||
}
|
||||
}
|
||||
// Exit on EOF
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
}
|
||||
if periods != 2 {
|
||||
return nil, nil, nil, errors.New(`invalid number of segments`)
|
||||
}
|
||||
|
||||
return protected, payload, signature, nil
|
||||
}
|
||||
|
||||
// parseCompact parses a JWS value serialized via compact serialization.
|
||||
func parseCompact(rdr io.Reader) (m *Message, err error) {
|
||||
protected, payload, signature, err := SplitCompact(rdr)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `invalid compact serialization format`)
|
||||
}
|
||||
|
||||
decodedHeader := make([]byte, base64.RawURLEncoding.DecodedLen(len(protected)))
|
||||
if _, err := base64.RawURLEncoding.Decode(decodedHeader, protected); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to decode headers`)
|
||||
}
|
||||
var hdr stdHeaders
|
||||
if err := json.Unmarshal(decodedHeader, &hdr); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to parse JOSE headers`)
|
||||
}
|
||||
|
||||
decodedPayload := make([]byte, base64.RawURLEncoding.DecodedLen(len(payload)))
|
||||
if _, err = base64.RawURLEncoding.Decode(decodedPayload, payload); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to decode payload`)
|
||||
}
|
||||
|
||||
decodedSignature := make([]byte, base64.RawURLEncoding.DecodedLen(len(signature)))
|
||||
if _, err := base64.RawURLEncoding.Decode(decodedSignature, signature); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to decode signature`)
|
||||
}
|
||||
|
||||
var msg Message
|
||||
msg.payload = decodedPayload
|
||||
msg.signatures = append(msg.signatures, &Signature{
|
||||
protected: &hdr,
|
||||
signature: decodedSignature,
|
||||
})
|
||||
return &msg, nil
|
||||
}
|
||||
+45
@@ -0,0 +1,45 @@
|
||||
package jws
|
||||
|
||||
func (s Signature) PublicHeaders() Headers {
|
||||
return s.headers
|
||||
}
|
||||
|
||||
func (s Signature) ProtectedHeaders() Headers {
|
||||
return s.protected
|
||||
}
|
||||
|
||||
func (s Signature) Signature() []byte {
|
||||
return s.signature
|
||||
}
|
||||
|
||||
func (m Message) Payload() []byte {
|
||||
return m.payload
|
||||
}
|
||||
|
||||
func (m Message) Signatures() []*Signature {
|
||||
return m.signatures
|
||||
}
|
||||
|
||||
// LookupSignature looks up a particular signature entry using
|
||||
// the `kid` value
|
||||
func (m Message) LookupSignature(kid string) []*Signature {
|
||||
var sigs []*Signature
|
||||
for _, sig := range m.signatures {
|
||||
if hdr := sig.PublicHeaders(); hdr != nil {
|
||||
hdrKeyID := hdr.KeyID()
|
||||
if hdrKeyID == kid {
|
||||
sigs = append(sigs, sig)
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if hdr := sig.ProtectedHeaders(); hdr != nil {
|
||||
hdrKeyID := hdr.KeyID()
|
||||
if hdrKeyID == kid {
|
||||
sigs = append(sigs, sig)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
return sigs
|
||||
}
|
||||
+26
@@ -0,0 +1,26 @@
|
||||
package jws
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/jwx/internal/option"
|
||||
"github.com/lestrrat-go/jwx/jws/sign"
|
||||
)
|
||||
|
||||
type Option = option.Interface
|
||||
|
||||
const (
|
||||
optkeyPayloadSigner = `payload-signer`
|
||||
optkeyHeaders = `headers`
|
||||
)
|
||||
|
||||
func WithSigner(signer sign.Signer, key interface{}, public, protected Headers) Option {
|
||||
return option.New(optkeyPayloadSigner, &payloadSigner{
|
||||
signer: signer,
|
||||
key: key,
|
||||
protected: protected,
|
||||
public: public,
|
||||
})
|
||||
}
|
||||
|
||||
func WithHeaders(h Headers) Option {
|
||||
return option.New(optkeyHeaders, h)
|
||||
}
|
||||
+88
@@ -0,0 +1,88 @@
|
||||
package sign
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/rand"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
var ecdsaSignFuncs = map[jwa.SignatureAlgorithm]ecdsaSignFunc{}
|
||||
|
||||
func init() {
|
||||
algs := map[jwa.SignatureAlgorithm]crypto.Hash{
|
||||
jwa.ES256: crypto.SHA256,
|
||||
jwa.ES384: crypto.SHA384,
|
||||
jwa.ES512: crypto.SHA512,
|
||||
}
|
||||
|
||||
for alg, h := range algs {
|
||||
ecdsaSignFuncs[alg] = makeECDSASignFunc(h)
|
||||
}
|
||||
}
|
||||
|
||||
func makeECDSASignFunc(hash crypto.Hash) ecdsaSignFunc {
|
||||
return func(payload []byte, key *ecdsa.PrivateKey) ([]byte, error) {
|
||||
curveBits := key.Curve.Params().BitSize
|
||||
keyBytes := curveBits / 8
|
||||
// Curve bits do not need to be a multiple of 8.
|
||||
if curveBits%8 > 0 {
|
||||
keyBytes++
|
||||
}
|
||||
h := hash.New()
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to write payload using ecdsa")
|
||||
}
|
||||
r, s, err := ecdsa.Sign(rand.Reader, key, h.Sum(nil))
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to sign payload using ecdsa")
|
||||
}
|
||||
|
||||
rBytes := r.Bytes()
|
||||
rBytesPadded := make([]byte, keyBytes)
|
||||
copy(rBytesPadded[keyBytes-len(rBytes):], rBytes)
|
||||
|
||||
sBytes := s.Bytes()
|
||||
sBytesPadded := make([]byte, keyBytes)
|
||||
copy(sBytesPadded[keyBytes-len(sBytes):], sBytes)
|
||||
|
||||
out := append(rBytesPadded, sBytesPadded...)
|
||||
return out, nil
|
||||
}
|
||||
}
|
||||
|
||||
func newECDSA(alg jwa.SignatureAlgorithm) (*ECDSASigner, error) {
|
||||
signfn, ok := ecdsaSignFuncs[alg]
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported algorithm while trying to create ECDSA signer: %s`, alg)
|
||||
}
|
||||
|
||||
return &ECDSASigner{
|
||||
alg: alg,
|
||||
sign: signfn,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s ECDSASigner) Algorithm() jwa.SignatureAlgorithm {
|
||||
return s.alg
|
||||
}
|
||||
|
||||
func (s ECDSASigner) Sign(payload []byte, key interface{}) ([]byte, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`missing private key while signing payload`)
|
||||
}
|
||||
|
||||
var pubkey *ecdsa.PrivateKey
|
||||
switch v := key.(type) {
|
||||
case ecdsa.PrivateKey:
|
||||
pubkey = &v
|
||||
case *ecdsa.PrivateKey:
|
||||
pubkey = v
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid key type %T. *ecdsa.PrivateKey is required`, key)
|
||||
}
|
||||
|
||||
return s.sign(payload, pubkey)
|
||||
}
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
package sign
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"crypto/sha512"
|
||||
"hash"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
var HMACSignFuncs = map[jwa.SignatureAlgorithm]hmacSignFunc{}
|
||||
|
||||
func init() {
|
||||
algs := map[jwa.SignatureAlgorithm]func() hash.Hash{
|
||||
jwa.HS256: sha256.New,
|
||||
jwa.HS384: sha512.New384,
|
||||
jwa.HS512: sha512.New,
|
||||
}
|
||||
|
||||
for alg, h := range algs {
|
||||
HMACSignFuncs[alg] = makeHMACSignFunc(h)
|
||||
}
|
||||
}
|
||||
|
||||
func newHMAC(alg jwa.SignatureAlgorithm) (*HMACSigner, error) {
|
||||
signer, ok := HMACSignFuncs[alg]
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported algorithm while trying to create HMAC signer: %s`, alg)
|
||||
}
|
||||
|
||||
return &HMACSigner{
|
||||
alg: alg,
|
||||
sign: signer,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func makeHMACSignFunc(hfunc func() hash.Hash) hmacSignFunc {
|
||||
return func(payload []byte, key []byte) ([]byte, error) {
|
||||
h := hmac.New(hfunc, key)
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to write payload using hmac")
|
||||
}
|
||||
return h.Sum(nil), nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s HMACSigner) Algorithm() jwa.SignatureAlgorithm {
|
||||
return s.alg
|
||||
}
|
||||
|
||||
func (s HMACSigner) Sign(payload []byte, key interface{}) ([]byte, error) {
|
||||
hmackey, ok := key.([]byte)
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`invalid key type %T. []byte is required`, key)
|
||||
}
|
||||
|
||||
if len(hmackey) == 0 {
|
||||
return nil, errors.New(`missing key while signing payload`)
|
||||
}
|
||||
|
||||
return s.sign(payload, hmackey)
|
||||
}
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
package sign
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/rsa"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
)
|
||||
|
||||
type Signer interface {
|
||||
// Sign creates a signature for the given `payload`.
|
||||
// `key` is the key used for signing the payload, and is usually
|
||||
// the private key type associated with the signature method. For example,
|
||||
// for `jwa.RSXXX` and `jwa.PSXXX` types, you need to pass the
|
||||
// `*"crypto/rsa".PrivateKey` type.
|
||||
// Check the documentation for each signer for details
|
||||
Sign(payload []byte, key interface{}) ([]byte, error)
|
||||
|
||||
Algorithm() jwa.SignatureAlgorithm
|
||||
}
|
||||
|
||||
type rsaSignFunc func([]byte, *rsa.PrivateKey) ([]byte, error)
|
||||
|
||||
// RSASigner uses crypto/rsa to sign the payloads.
|
||||
type RSASigner struct {
|
||||
alg jwa.SignatureAlgorithm
|
||||
sign rsaSignFunc
|
||||
}
|
||||
|
||||
type ecdsaSignFunc func([]byte, *ecdsa.PrivateKey) ([]byte, error)
|
||||
|
||||
// ECDSASigner uses crypto/ecdsa to sign the payloads.
|
||||
type ECDSASigner struct {
|
||||
alg jwa.SignatureAlgorithm
|
||||
sign ecdsaSignFunc
|
||||
}
|
||||
|
||||
type hmacSignFunc func([]byte, []byte) ([]byte, error)
|
||||
|
||||
// HMACSigner uses crypto/hmac to sign the payloads.
|
||||
type HMACSigner struct {
|
||||
alg jwa.SignatureAlgorithm
|
||||
sign hmacSignFunc
|
||||
}
|
||||
+105
@@ -0,0 +1,105 @@
|
||||
package sign
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
var rsaSignFuncs = map[jwa.SignatureAlgorithm]rsaSignFunc{}
|
||||
|
||||
func init() {
|
||||
algs := map[jwa.SignatureAlgorithm]struct {
|
||||
Hash crypto.Hash
|
||||
SignFunc func(crypto.Hash) rsaSignFunc
|
||||
}{
|
||||
jwa.RS256: {
|
||||
Hash: crypto.SHA256,
|
||||
SignFunc: makeSignPKCS1v15,
|
||||
},
|
||||
jwa.RS384: {
|
||||
Hash: crypto.SHA384,
|
||||
SignFunc: makeSignPKCS1v15,
|
||||
},
|
||||
jwa.RS512: {
|
||||
Hash: crypto.SHA512,
|
||||
SignFunc: makeSignPKCS1v15,
|
||||
},
|
||||
jwa.PS256: {
|
||||
Hash: crypto.SHA256,
|
||||
SignFunc: makeSignPSS,
|
||||
},
|
||||
jwa.PS384: {
|
||||
Hash: crypto.SHA384,
|
||||
SignFunc: makeSignPSS,
|
||||
},
|
||||
jwa.PS512: {
|
||||
Hash: crypto.SHA512,
|
||||
SignFunc: makeSignPSS,
|
||||
},
|
||||
}
|
||||
|
||||
for alg, item := range algs {
|
||||
rsaSignFuncs[alg] = item.SignFunc(item.Hash)
|
||||
}
|
||||
}
|
||||
|
||||
func makeSignPKCS1v15(hash crypto.Hash) rsaSignFunc {
|
||||
return func(payload []byte, key *rsa.PrivateKey) ([]byte, error) {
|
||||
h := hash.New()
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to write payload using SignPKCS1v15")
|
||||
}
|
||||
return rsa.SignPKCS1v15(rand.Reader, key, hash, h.Sum(nil))
|
||||
}
|
||||
}
|
||||
|
||||
func makeSignPSS(hash crypto.Hash) rsaSignFunc {
|
||||
return func(payload []byte, key *rsa.PrivateKey) ([]byte, error) {
|
||||
h := hash.New()
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to write payload using SignPSS")
|
||||
}
|
||||
return rsa.SignPSS(rand.Reader, key, hash, h.Sum(nil), &rsa.PSSOptions{
|
||||
SaltLength: rsa.PSSSaltLengthAuto,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func newRSA(alg jwa.SignatureAlgorithm) (*RSASigner, error) {
|
||||
signfn, ok := rsaSignFuncs[alg]
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported algorithm while trying to create RSA signer: %s`, alg)
|
||||
}
|
||||
return &RSASigner{
|
||||
alg: alg,
|
||||
sign: signfn,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s RSASigner) Algorithm() jwa.SignatureAlgorithm {
|
||||
return s.alg
|
||||
}
|
||||
|
||||
// Sign creates a signature using crypto/rsa. key must be a non-nil instance of
|
||||
// `*"crypto/rsa".PrivateKey`.
|
||||
func (s RSASigner) Sign(payload []byte, key interface{}) ([]byte, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`missing private key while signing payload`)
|
||||
}
|
||||
|
||||
var privkey *rsa.PrivateKey
|
||||
switch v := key.(type) {
|
||||
case rsa.PrivateKey:
|
||||
privkey = &v
|
||||
case *rsa.PrivateKey:
|
||||
privkey = v
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid key type %T. *rsa.PrivateKey is required`, key)
|
||||
}
|
||||
|
||||
return s.sign(payload, privkey)
|
||||
}
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
package sign
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// New creates a signer that signs payloads using the given signature algorithm.
|
||||
func New(alg jwa.SignatureAlgorithm) (Signer, error) {
|
||||
switch alg {
|
||||
case jwa.RS256, jwa.RS384, jwa.RS512, jwa.PS256, jwa.PS384, jwa.PS512:
|
||||
return newRSA(alg)
|
||||
case jwa.ES256, jwa.ES384, jwa.ES512:
|
||||
return newECDSA(alg)
|
||||
case jwa.HS256, jwa.HS384, jwa.HS512:
|
||||
return newHMAC(alg)
|
||||
default:
|
||||
return nil, errors.Errorf(`unsupported signature algorithm %s`, alg)
|
||||
}
|
||||
}
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
package jws
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type encodedSignatureProxy struct {
|
||||
Protected string `json:"protected,omitempty"`
|
||||
Headers json.RawMessage `json:"header,omitempty"`
|
||||
Signature string `json:"signature,omitempty"`
|
||||
}
|
||||
|
||||
func (sig *encodedSignature) UnmarshalJSON(buf []byte) error {
|
||||
var proxy encodedSignatureProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal into temporary struct`)
|
||||
}
|
||||
|
||||
var h Headers
|
||||
if len(proxy.Headers) > 0 {
|
||||
h = NewHeaders()
|
||||
if err := json.Unmarshal(proxy.Headers, h); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal headers`)
|
||||
}
|
||||
}
|
||||
|
||||
// XXX: sigh, dream of the day when we kill public fields
|
||||
sig.Protected = proxy.Protected
|
||||
sig.Signature = proxy.Signature
|
||||
sig.Headers = h
|
||||
|
||||
return nil
|
||||
}
|
||||
+1
@@ -0,0 +1 @@
|
||||
package jws
|
||||
+76
@@ -0,0 +1,76 @@
|
||||
package verify
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/pool"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
var ecdsaVerifyFuncs = map[jwa.SignatureAlgorithm]ecdsaVerifyFunc{}
|
||||
|
||||
func init() {
|
||||
algs := map[jwa.SignatureAlgorithm]crypto.Hash{
|
||||
jwa.ES256: crypto.SHA256,
|
||||
jwa.ES384: crypto.SHA384,
|
||||
jwa.ES512: crypto.SHA512,
|
||||
}
|
||||
|
||||
for alg, h := range algs {
|
||||
ecdsaVerifyFuncs[alg] = makeECDSAVerifyFunc(h)
|
||||
}
|
||||
}
|
||||
|
||||
func makeECDSAVerifyFunc(hash crypto.Hash) ecdsaVerifyFunc {
|
||||
return func(payload []byte, signature []byte, key *ecdsa.PublicKey) error {
|
||||
r := pool.GetBigInt()
|
||||
s := pool.GetBigInt()
|
||||
defer pool.ReleaseBigInt(r)
|
||||
defer pool.ReleaseBigInt(s)
|
||||
|
||||
n := len(signature) / 2
|
||||
r.SetBytes(signature[:n])
|
||||
s.SetBytes(signature[n:])
|
||||
|
||||
h := hash.New()
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return errors.Wrap(err, "failed to write payload using ecdsa")
|
||||
}
|
||||
|
||||
if !ecdsa.Verify(key, h.Sum(nil), r, s) {
|
||||
return errors.New(`failed to verify signature using ecdsa`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func newECDSA(alg jwa.SignatureAlgorithm) (*ECDSAVerifier, error) {
|
||||
verifyfn, ok := ecdsaVerifyFuncs[alg]
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported algorithm while trying to create ECDSA verifier: %s`, alg)
|
||||
}
|
||||
|
||||
return &ECDSAVerifier{
|
||||
verify: verifyfn,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (v ECDSAVerifier) Verify(payload []byte, signature []byte, key interface{}) error {
|
||||
if key == nil {
|
||||
return errors.New(`missing public key while verifying payload`)
|
||||
}
|
||||
|
||||
var pubkey *ecdsa.PublicKey
|
||||
switch v := key.(type) {
|
||||
case ecdsa.PublicKey:
|
||||
pubkey = &v
|
||||
case *ecdsa.PublicKey:
|
||||
pubkey = v
|
||||
default:
|
||||
return errors.Errorf(`invalid key type %T. *ecdsa.PublicKey is required`, key)
|
||||
}
|
||||
|
||||
return v.verify(payload, signature, pubkey)
|
||||
}
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
package verify
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jws/sign"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func newHMAC(alg jwa.SignatureAlgorithm) (*HMACVerifier, error) {
|
||||
_, ok := sign.HMACSignFuncs[alg]
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported algorithm while trying to create HMAC signer: %s`, alg)
|
||||
}
|
||||
s, err := sign.New(alg)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to generate HMAC signer`)
|
||||
}
|
||||
return &HMACVerifier{signer: s}, nil
|
||||
}
|
||||
|
||||
func (v HMACVerifier) Verify(payload, signature []byte, key interface{}) (err error) {
|
||||
expected, err := v.signer.Sign(payload, key)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to generated signature`)
|
||||
}
|
||||
|
||||
if !hmac.Equal(signature, expected) {
|
||||
return errors.New(`failed to match hmac signature`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
package verify
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/rsa"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jws/sign"
|
||||
)
|
||||
|
||||
type Verifier interface {
|
||||
// Verify checks whether the payload and signature are valid for
|
||||
// the given key.
|
||||
// `key` is the key used for verifying the payload, and is usually
|
||||
// the public key associated with the signature method. For example,
|
||||
// for `jwa.RSXXX` and `jwa.PSXXX` types, you need to pass the
|
||||
// `*"crypto/rsa".PublicKey` type.
|
||||
// Check the documentation for each verifier for details
|
||||
Verify(payload []byte, signature []byte, key interface{}) error
|
||||
}
|
||||
|
||||
type rsaVerifyFunc func([]byte, []byte, *rsa.PublicKey) error
|
||||
|
||||
type RSAVerifier struct {
|
||||
verify rsaVerifyFunc
|
||||
}
|
||||
|
||||
type ecdsaVerifyFunc func([]byte, []byte, *ecdsa.PublicKey) error
|
||||
|
||||
type ECDSAVerifier struct {
|
||||
verify ecdsaVerifyFunc
|
||||
}
|
||||
|
||||
type HMACVerifier struct {
|
||||
signer sign.Signer
|
||||
}
|
||||
+97
@@ -0,0 +1,97 @@
|
||||
package verify
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/rsa"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
var rsaVerifyFuncs = map[jwa.SignatureAlgorithm]rsaVerifyFunc{}
|
||||
|
||||
func init() {
|
||||
algs := map[jwa.SignatureAlgorithm]struct {
|
||||
Hash crypto.Hash
|
||||
VerifyFunc func(crypto.Hash) rsaVerifyFunc
|
||||
}{
|
||||
jwa.RS256: {
|
||||
Hash: crypto.SHA256,
|
||||
VerifyFunc: makeVerifyPKCS1v15,
|
||||
},
|
||||
jwa.RS384: {
|
||||
Hash: crypto.SHA384,
|
||||
VerifyFunc: makeVerifyPKCS1v15,
|
||||
},
|
||||
jwa.RS512: {
|
||||
Hash: crypto.SHA512,
|
||||
VerifyFunc: makeVerifyPKCS1v15,
|
||||
},
|
||||
jwa.PS256: {
|
||||
Hash: crypto.SHA256,
|
||||
VerifyFunc: makeVerifyPSS,
|
||||
},
|
||||
jwa.PS384: {
|
||||
Hash: crypto.SHA384,
|
||||
VerifyFunc: makeVerifyPSS,
|
||||
},
|
||||
jwa.PS512: {
|
||||
Hash: crypto.SHA512,
|
||||
VerifyFunc: makeVerifyPSS,
|
||||
},
|
||||
}
|
||||
|
||||
for alg, item := range algs {
|
||||
rsaVerifyFuncs[alg] = item.VerifyFunc(item.Hash)
|
||||
}
|
||||
}
|
||||
|
||||
func makeVerifyPKCS1v15(hash crypto.Hash) rsaVerifyFunc {
|
||||
return func(payload, signature []byte, key *rsa.PublicKey) error {
|
||||
h := hash.New()
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return errors.Wrap(err, "failed to write payload using PKCS1v15")
|
||||
}
|
||||
|
||||
return rsa.VerifyPKCS1v15(key, hash, h.Sum(nil), signature)
|
||||
}
|
||||
}
|
||||
|
||||
func makeVerifyPSS(hash crypto.Hash) rsaVerifyFunc {
|
||||
return func(payload, signature []byte, key *rsa.PublicKey) error {
|
||||
h := hash.New()
|
||||
if _, err := h.Write(payload); err != nil {
|
||||
return errors.Wrap(err, "failed to write payload using PSS")
|
||||
}
|
||||
return rsa.VerifyPSS(key, hash, h.Sum(nil), signature, nil)
|
||||
}
|
||||
}
|
||||
|
||||
func newRSA(alg jwa.SignatureAlgorithm) (*RSAVerifier, error) {
|
||||
verifyfn, ok := rsaVerifyFuncs[alg]
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported algorithm while trying to create RSA verifier: %s`, alg)
|
||||
}
|
||||
|
||||
return &RSAVerifier{
|
||||
verify: verifyfn,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (v RSAVerifier) Verify(payload, signature []byte, key interface{}) error {
|
||||
if key == nil {
|
||||
return errors.New(`missing public key while verifying payload`)
|
||||
}
|
||||
|
||||
var pubkey *rsa.PublicKey
|
||||
switch v := key.(type) {
|
||||
case rsa.PublicKey:
|
||||
pubkey = &v
|
||||
case *rsa.PublicKey:
|
||||
pubkey = v
|
||||
default:
|
||||
return errors.Errorf(`invalid key type %T. *rsa.PublicKey is required`, key)
|
||||
}
|
||||
|
||||
return v.verify(payload, signature, pubkey)
|
||||
}
|
||||
+21
@@ -0,0 +1,21 @@
|
||||
package verify
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// New creates a new JWS verifier using the specified algorithm
|
||||
// and the public key
|
||||
func New(alg jwa.SignatureAlgorithm) (Verifier, error) {
|
||||
switch alg {
|
||||
case jwa.RS256, jwa.RS384, jwa.RS512, jwa.PS256, jwa.PS384, jwa.PS512:
|
||||
return newRSA(alg)
|
||||
case jwa.ES256, jwa.ES384, jwa.ES512:
|
||||
return newECDSA(alg)
|
||||
case jwa.HS256, jwa.HS384, jwa.HS512:
|
||||
return newHMAC(alg)
|
||||
default:
|
||||
return nil, errors.Errorf(`unsupported signature algorithm: %s`, alg)
|
||||
}
|
||||
}
|
||||
+78
@@ -0,0 +1,78 @@
|
||||
# jwt
|
||||
|
||||
JWT tokens
|
||||
|
||||
# SYNOPSIS
|
||||
|
||||
```go
|
||||
package jwt_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jwt"
|
||||
)
|
||||
|
||||
func ExampleSignAndParse() {
|
||||
privKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
fmt.Printf("failed to generate private key: %s\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
var payload []byte
|
||||
{ // Create signed payload
|
||||
token := jwt.New()
|
||||
token.Set(`foo`, `bar`)
|
||||
payload, err = jwt.Sign(token, jwa.RS256, privKey)
|
||||
if err != nil {
|
||||
fmt.Printf("failed to generate signed payload: %s\n", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
{ // Parse signed payload
|
||||
// Use jwt.ParseVerify if you want to make absolutely sure that you
|
||||
// are going to verify the signatures every time
|
||||
token, err := jwt.Parse(bytes.NewReader(payload), jwt.WithVerify(jwa.RS256, &privKey.PublicKey))
|
||||
if err != nil {
|
||||
fmt.Printf("failed to parse JWT token: %s\n", err)
|
||||
return
|
||||
}
|
||||
buf, err := json.MarshalIndent(token, "", " ")
|
||||
if err != nil {
|
||||
fmt.Printf("failed to generate JSON: %s\n", err)
|
||||
return
|
||||
}
|
||||
fmt.Printf("%s\n", buf)
|
||||
}
|
||||
}
|
||||
|
||||
func ExampleToken() {
|
||||
t := jwt.New()
|
||||
t.Set(jwt.SubjectKey, `https://github.com/lestrrat-go/jwx/jwt`)
|
||||
t.Set(jwt.AudienceKey, `Golang Users`)
|
||||
t.Set(jwt.IssuedAtKey, time.Unix(aLongLongTimeAgo, 0))
|
||||
t.Set(`privateClaimKey`, `Hello, World!`)
|
||||
|
||||
buf, err := json.MarshalIndent(t, "", " ")
|
||||
if err != nil {
|
||||
fmt.Printf("failed to generate JSON: %s\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Printf("%s\n", buf)
|
||||
fmt.Printf("aud -> '%s'\n", t.Audience())
|
||||
fmt.Printf("iat -> '%s'\n", t.IssuedAt().Format(time.RFC3339))
|
||||
if v, ok := t.Get(`privateClaimKey`); ok {
|
||||
fmt.Printf("privateClaimKey -> '%s'\n", v)
|
||||
}
|
||||
fmt.Printf("sub -> '%s'\n", t.Subject())
|
||||
}
|
||||
```
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
package jwt
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
)
|
||||
|
||||
type ClaimPair = mapiter.Pair
|
||||
type Iterator = mapiter.Iterator
|
||||
type Visitor = iter.MapVisitor
|
||||
type VisitorFunc iter.MapVisitorFunc
|
||||
+97
@@ -0,0 +1,97 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// NumericDate represents the date format used in the 'nbf' claim
|
||||
type NumericDate struct {
|
||||
time.Time
|
||||
}
|
||||
|
||||
func (n *NumericDate) Get() time.Time {
|
||||
if n == nil {
|
||||
return (time.Time{}).UTC()
|
||||
}
|
||||
return n.Time
|
||||
}
|
||||
|
||||
func numericToTime(v interface{}, t *time.Time) bool {
|
||||
var n int64
|
||||
switch x := v.(type) {
|
||||
case int64:
|
||||
n = x
|
||||
case int32:
|
||||
n = int64(x)
|
||||
case int16:
|
||||
n = int64(x)
|
||||
case int8:
|
||||
n = int64(x)
|
||||
case int:
|
||||
n = int64(x)
|
||||
case float32:
|
||||
n = int64(x)
|
||||
case float64:
|
||||
n = int64(x)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
|
||||
*t = time.Unix(n, 0)
|
||||
return true
|
||||
}
|
||||
|
||||
func (n *NumericDate) Accept(v interface{}) error {
|
||||
var t time.Time
|
||||
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
i, err := strconv.ParseInt(x[:], 10, 64)
|
||||
if err != nil {
|
||||
return errors.Errorf(`invalid epoch value %#v`, x)
|
||||
}
|
||||
t = time.Unix(i, 0)
|
||||
|
||||
case json.Number:
|
||||
intval, err := x.Int64()
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to convert json value %#v to int64`, x)
|
||||
}
|
||||
t = time.Unix(intval, 0)
|
||||
case time.Time:
|
||||
t = x
|
||||
default:
|
||||
if !numericToTime(v, &t) {
|
||||
return errors.Errorf(`invalid type %T`, v)
|
||||
}
|
||||
}
|
||||
n.Time = t.UTC()
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalJSON translates from internal representation to JSON NumericDate
|
||||
// See https://tools.ietf.org/html/rfc7519#page-6
|
||||
func (n *NumericDate) MarshalJSON() ([]byte, error) {
|
||||
if n.IsZero() {
|
||||
return json.Marshal(nil)
|
||||
}
|
||||
return json.Marshal(n.Unix())
|
||||
}
|
||||
|
||||
func (n *NumericDate) UnmarshalJSON(data []byte) error {
|
||||
var v interface{}
|
||||
if err := json.Unmarshal(data, &v); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal date`)
|
||||
}
|
||||
|
||||
var n2 NumericDate
|
||||
if err := n2.Accept(v); err != nil {
|
||||
return errors.Wrap(err, `invalid value for NumericDate`)
|
||||
}
|
||||
*n = n2
|
||||
return nil
|
||||
}
|
||||
+43
@@ -0,0 +1,43 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type StringList []string
|
||||
|
||||
func (l StringList) Get() []string {
|
||||
return []string(l)
|
||||
}
|
||||
|
||||
func (l *StringList) Accept(v interface{}) error {
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
*l = StringList([]string{x})
|
||||
case []string:
|
||||
*l = StringList(x)
|
||||
case []interface{}:
|
||||
list := make(StringList, len(x))
|
||||
for i, e := range x {
|
||||
if s, ok := e.(string); ok {
|
||||
list[i] = s
|
||||
continue
|
||||
}
|
||||
return errors.Errorf(`invalid list element type %T`, e)
|
||||
}
|
||||
*l = list
|
||||
default:
|
||||
return errors.Errorf(`invalid type: %T`, v)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *StringList) UnmarshalJSON(data []byte) error {
|
||||
var v interface{}
|
||||
if err := json.Unmarshal(data, &v); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal data`)
|
||||
}
|
||||
return l.Accept(v)
|
||||
}
|
||||
+101
@@ -0,0 +1,101 @@
|
||||
//go:generate go run internal/cmd/gentoken/main.go
|
||||
|
||||
// Package jwt implements JSON Web Tokens as described in https://tools.ietf.org/html/rfc7519
|
||||
package jwt
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"io/ioutil"
|
||||
"strings"
|
||||
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jws"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// ParseString calls Parse with the given string
|
||||
func ParseString(s string, options ...Option) (Token, error) {
|
||||
return Parse(strings.NewReader(s), options...)
|
||||
}
|
||||
|
||||
// ParseString calls Parse with the given byte sequence
|
||||
func ParseBytes(s []byte, options ...Option) (Token, error) {
|
||||
return Parse(bytes.NewReader(s), options...)
|
||||
}
|
||||
|
||||
// Parse parses the JWT token payload and creates a new `jwt.Token` object.
|
||||
// The token must be encoded in either JSON format or compact format.
|
||||
//
|
||||
// If the token is signed and you want to verify the payload, you must
|
||||
// pass the jwt.WithVerify(alg, key) option. If you do not specify these
|
||||
// parameters, no verification will be performed.
|
||||
func Parse(src io.Reader, options ...Option) (Token, error) {
|
||||
var params VerifyParameters
|
||||
for _, o := range options {
|
||||
switch o.Name() {
|
||||
case optkeyVerify:
|
||||
params = o.Value().(VerifyParameters)
|
||||
}
|
||||
}
|
||||
|
||||
if params != nil {
|
||||
return ParseVerify(src, params.Algorithm(), params.Key())
|
||||
}
|
||||
|
||||
m, err := jws.Parse(src)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `invalid jws message`)
|
||||
}
|
||||
|
||||
token := New()
|
||||
if err := json.Unmarshal(m.Payload(), token); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to parse token`)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// ParseVerify is a function that is similar to Parse(), but does not
|
||||
// allow for parsing without signature verification parameters.
|
||||
func ParseVerify(src io.Reader, alg jwa.SignatureAlgorithm, key interface{}) (Token, error) {
|
||||
data, err := ioutil.ReadAll(src)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to read token from source`)
|
||||
}
|
||||
|
||||
v, err := jws.Verify(data, alg, key)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to verify jws signature`)
|
||||
}
|
||||
|
||||
t := New()
|
||||
if err := json.Unmarshal(v, t); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to parse token`)
|
||||
}
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// Sign is a convenience function to create a signed JWT token serialized in
|
||||
// compact form. `key` must match the key type required by the given
|
||||
// signature method `method`
|
||||
func Sign(t Token, method jwa.SignatureAlgorithm, key interface{}) ([]byte, error) {
|
||||
buf, err := json.Marshal(t)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to marshal token`)
|
||||
}
|
||||
|
||||
hdr := jws.NewHeaders()
|
||||
if hdr.Set(`alg`, method.String()) != nil {
|
||||
return nil, errors.Wrap(err, `failed to sign payload`)
|
||||
}
|
||||
if hdr.Set(`typ`, `JWT`) != nil {
|
||||
return nil, errors.Wrap(err, `failed to sign payload`)
|
||||
}
|
||||
sign, err := jws.Sign(buf, method, key, jws.WithHeaders(hdr))
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to sign payload`)
|
||||
}
|
||||
|
||||
return sign, nil
|
||||
}
|
||||
+221
@@ -0,0 +1,221 @@
|
||||
package openid
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
AddressFormattedKey = "formatted"
|
||||
AddressStreetAddressKey = "street_address"
|
||||
AddressLocalityKey = "locality"
|
||||
AddressRegionKey = "region"
|
||||
AddressPostalCodeKey = "postal_code"
|
||||
AddressCountryKey = "country"
|
||||
)
|
||||
|
||||
// AddressClaim is the address claim as described in https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
type AddressClaim struct {
|
||||
formatted *string // https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
streetAddress *string // https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
locality *string // https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
region *string // https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
postalCode *string // https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
country *string // https://openid.net/specs/openid-connect-core-1_0.html#AddressClaim
|
||||
}
|
||||
|
||||
type addressClaimMarshalProxy struct {
|
||||
Xformatted *string `json:"formatted,omitempty"`
|
||||
XstreetAddress *string `json:"street_address,omitempty"`
|
||||
Xlocality *string `json:"locality,omitempty"`
|
||||
Xregion *string `json:"region,omitempty"`
|
||||
XpostalCode *string `json:"postal_code,omitempty"`
|
||||
Xcountry *string `json:"country,omitempty"`
|
||||
}
|
||||
|
||||
func NewAddress() *AddressClaim {
|
||||
return &AddressClaim{}
|
||||
}
|
||||
|
||||
// Formatted is a convenience function to retrieve the corresponding value store in the token
|
||||
// if there is a problem retrieving the value, the zero value is returned. If you need to differentiate between existing/non-existing values, use `Get` instead
|
||||
func (t AddressClaim) Formatted() string {
|
||||
if t.formatted == nil {
|
||||
return ""
|
||||
}
|
||||
return *(t.formatted)
|
||||
}
|
||||
|
||||
// StreetAddress is a convenience function to retrieve the corresponding value store in the token
|
||||
// if there is a problem retrieving the value, the zero value is returned. If you need to differentiate between existing/non-existing values, use `Get` instead
|
||||
func (t AddressClaim) StreetAddress() string {
|
||||
if t.streetAddress == nil {
|
||||
return ""
|
||||
}
|
||||
return *(t.streetAddress)
|
||||
}
|
||||
|
||||
// Locality is a convenience function to retrieve the corresponding value store in the token
|
||||
// if there is a problem retrieving the value, the zero value is returned. If you need to differentiate between existing/non-existing values, use `Get` instead
|
||||
func (t AddressClaim) Locality() string {
|
||||
if t.locality == nil {
|
||||
return ""
|
||||
}
|
||||
return *(t.locality)
|
||||
}
|
||||
|
||||
// Region is a convenience function to retrieve the corresponding value store in the token
|
||||
// if there is a problem retrieving the value, the zero value is returned. If you need to differentiate between existing/non-existing values, use `Get` instead
|
||||
func (t AddressClaim) Region() string {
|
||||
if t.region == nil {
|
||||
return ""
|
||||
}
|
||||
return *(t.region)
|
||||
}
|
||||
|
||||
// PostalCode is a convenience function to retrieve the corresponding value store in the token
|
||||
// if there is a problem retrieving the value, the zero value is returned. If you need to differentiate between existing/non-existing values, use `Get` instead
|
||||
func (t AddressClaim) PostalCode() string {
|
||||
if t.postalCode == nil {
|
||||
return ""
|
||||
}
|
||||
return *(t.postalCode)
|
||||
}
|
||||
|
||||
// Country is a convenience function to retrieve the corresponding value store in the token
|
||||
// if there is a problem retrieving the value, the zero value is returned. If you need to differentiate between existing/non-existing values, use `Get` instead
|
||||
func (t AddressClaim) Country() string {
|
||||
if t.country == nil {
|
||||
return ""
|
||||
}
|
||||
return *(t.country)
|
||||
}
|
||||
|
||||
func (t *AddressClaim) Get(s string) (interface{}, bool) {
|
||||
switch s {
|
||||
case AddressFormattedKey:
|
||||
if t.formatted == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(t.formatted), true
|
||||
case AddressStreetAddressKey:
|
||||
if t.streetAddress == nil {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
return *(t.streetAddress), true
|
||||
case AddressLocalityKey:
|
||||
if t.locality == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(t.locality), true
|
||||
case AddressRegionKey:
|
||||
if t.region == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(t.region), true
|
||||
case AddressPostalCodeKey:
|
||||
if t.postalCode == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(t.postalCode), true
|
||||
case AddressCountryKey:
|
||||
if t.country == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *(t.country), true
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func (t *AddressClaim) Set(key string, value interface{}) error {
|
||||
switch key {
|
||||
case AddressFormattedKey:
|
||||
v, ok := value.(string)
|
||||
if ok {
|
||||
t.formatted = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid type for key 'formatted': %T`, value)
|
||||
case AddressStreetAddressKey:
|
||||
v, ok := value.(string)
|
||||
if ok {
|
||||
t.streetAddress = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid type for key 'streetAddress': %T`, value)
|
||||
case AddressLocalityKey:
|
||||
v, ok := value.(string)
|
||||
if ok {
|
||||
t.locality = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid type for key 'locality': %T`, value)
|
||||
case AddressRegionKey:
|
||||
v, ok := value.(string)
|
||||
if ok {
|
||||
t.region = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid type for key 'region': %T`, value)
|
||||
case AddressPostalCodeKey:
|
||||
v, ok := value.(string)
|
||||
if ok {
|
||||
t.postalCode = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid type for key 'postalCode': %T`, value)
|
||||
case AddressCountryKey:
|
||||
v, ok := value.(string)
|
||||
if ok {
|
||||
t.country = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid type for key 'country': %T`, value)
|
||||
default:
|
||||
return errors.Errorf(`invalid key for address claim: %s`, key)
|
||||
}
|
||||
}
|
||||
|
||||
func (t *AddressClaim) Accept(v interface{}) error {
|
||||
switch v := v.(type) {
|
||||
case map[string]interface{}:
|
||||
for key, value := range v {
|
||||
if err := t.Set(key, value); err != nil {
|
||||
return errors.Wrap(err, `failed to set header`)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return errors.Errorf(`invalid type for AddressClaim: %T`, v)
|
||||
}
|
||||
}
|
||||
|
||||
// MarshalJSON serializes the token in JSON format.
|
||||
func (t AddressClaim) MarshalJSON() ([]byte, error) {
|
||||
var proxy addressClaimMarshalProxy
|
||||
proxy.Xformatted = t.formatted
|
||||
proxy.XstreetAddress = t.streetAddress
|
||||
proxy.Xlocality = t.locality
|
||||
proxy.Xregion = t.region
|
||||
proxy.XpostalCode = t.postalCode
|
||||
proxy.Xcountry = t.country
|
||||
return json.Marshal(proxy)
|
||||
}
|
||||
|
||||
// UnmarshalJSON deserializes data from a JSON data buffer into a AddressClaim
|
||||
func (t *AddressClaim) UnmarshalJSON(data []byte) error {
|
||||
var proxy addressClaimMarshalProxy
|
||||
if err := json.Unmarshal(data, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshasl address claim`)
|
||||
}
|
||||
|
||||
t.formatted = proxy.Xformatted
|
||||
t.streetAddress = proxy.XstreetAddress
|
||||
t.locality = proxy.Xlocality
|
||||
t.region = proxy.Xregion
|
||||
t.postalCode = proxy.XpostalCode
|
||||
t.country = proxy.Xcountry
|
||||
return nil
|
||||
}
|
||||
+124
@@ -0,0 +1,124 @@
|
||||
package openid
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"regexp"
|
||||
"strconv"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// https://openid.net/specs/openid-connect-core-1_0.html
|
||||
//
|
||||
// End-User's birthday, represented as an ISO 8601:2004 [ISO8601‑2004] YYYY-MM-DD format.
|
||||
// The year MAY be 0000, indicating that it is omitted. To represent only the year, YYYY
|
||||
// format is allowed. Note that depending on the underlying platform's date related function,
|
||||
// providing just year can result in varying month and day, so the implementers need to
|
||||
// take this factor into account to correctly process the dates.
|
||||
|
||||
type BirthdateClaim struct {
|
||||
year *int
|
||||
month *int
|
||||
day *int
|
||||
}
|
||||
|
||||
func (b BirthdateClaim) Year() int {
|
||||
if b.year == nil {
|
||||
return 0
|
||||
}
|
||||
return *(b.year)
|
||||
}
|
||||
|
||||
func (b BirthdateClaim) Month() int {
|
||||
if b.month == nil {
|
||||
return 0
|
||||
}
|
||||
return *(b.month)
|
||||
}
|
||||
|
||||
func (b BirthdateClaim) Day() int {
|
||||
if b.day == nil {
|
||||
return 0
|
||||
}
|
||||
return *(b.day)
|
||||
}
|
||||
|
||||
func (b *BirthdateClaim) UnmarshalJSON(data []byte) error {
|
||||
var s string
|
||||
if err := json.Unmarshal(data, &s); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal JSON string for birthdate claim`)
|
||||
}
|
||||
|
||||
if err := b.Accept(s); err != nil {
|
||||
return errors.Wrap(err, `failed to accept JSON value for birthdate claim`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var birthdateRx = regexp.MustCompile(`^(\d{4})-(\d{2})-(\d{2})$`)
|
||||
|
||||
// Accepts a value read from JSON, and converts it to a BirthdateClaim.
|
||||
// This method DOES NOT verify the correctness of a date.
|
||||
// Consumers should check for validity of dates such as Apr 31 et al
|
||||
func (b *BirthdateClaim) Accept(v interface{}) error {
|
||||
switch v := v.(type) {
|
||||
case string:
|
||||
// yeah, yeah, regexp is slow. PR's welcome
|
||||
indices := birthdateRx.FindStringSubmatchIndex(v)
|
||||
if indices == nil {
|
||||
return errors.New(`invalid pattern for birthdate`)
|
||||
}
|
||||
var tmp BirthdateClaim
|
||||
|
||||
year, err := strconv.ParseInt(v[indices[2]:indices[3]], 10, 64)
|
||||
if err != nil {
|
||||
return errors.New(`failed to parse birthdate year`)
|
||||
}
|
||||
if year > 0 {
|
||||
var v int = int(year)
|
||||
tmp.year = &v
|
||||
}
|
||||
|
||||
month, err := strconv.ParseInt(v[indices[4]:indices[5]], 10, 64)
|
||||
if err != nil {
|
||||
return errors.New(`failed to parse birthdate month`)
|
||||
}
|
||||
if month > 0 {
|
||||
var v int = int(month)
|
||||
tmp.month = &v
|
||||
}
|
||||
|
||||
day, err := strconv.ParseInt(v[indices[6]:indices[7]], 10, 64)
|
||||
if err != nil {
|
||||
return errors.New(`failed to parse birthdate day`)
|
||||
}
|
||||
if day > 0 {
|
||||
var v int = int(day)
|
||||
tmp.day = &v
|
||||
}
|
||||
|
||||
*b = tmp
|
||||
return nil
|
||||
default:
|
||||
return errors.Errorf(`invalid type for birthdate: %T`, v)
|
||||
}
|
||||
}
|
||||
|
||||
func (b BirthdateClaim) encode(dst io.Writer) {
|
||||
fmt.Fprintf(dst, "%d-%02d-%02d", b.Year(), b.Month(), b.Day())
|
||||
}
|
||||
|
||||
func (b BirthdateClaim) String() string {
|
||||
var buf bytes.Buffer
|
||||
b.encode(&buf)
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func (b BirthdateClaim) MarshalText() ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
b.encode(&buf)
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
+8
@@ -0,0 +1,8 @@
|
||||
// Package openid provides a specialized token that provides utilities
|
||||
// to work with OpenID JWT tokens.
|
||||
//
|
||||
// In order to use OpenID claims, you specify the token to use in the
|
||||
// jwt.Parse method
|
||||
//
|
||||
// jwt.Parse(data, jwt.WithOpenIDClaims())
|
||||
package openid
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
package openid
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
)
|
||||
|
||||
type ClaimPair = mapiter.Pair
|
||||
type Iterator = mapiter.Iterator
|
||||
type Visitor = iter.MapVisitor
|
||||
type VisitorFunc = iter.MapVisitorFunc
|
||||
+948
@@ -0,0 +1,948 @@
|
||||
// This file is auto-generated by jwt/internal/cmd/gentoken/main.go. DO NOT EDIT
|
||||
package openid
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
"github.com/lestrrat-go/jwx/jwt/internal/types"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
AudienceKey = "aud"
|
||||
ExpirationKey = "exp"
|
||||
IssuedAtKey = "iat"
|
||||
IssuerKey = "iss"
|
||||
JwtIDKey = "jti"
|
||||
NotBeforeKey = "nbf"
|
||||
SubjectKey = "sub"
|
||||
NameKey = "name"
|
||||
GivenNameKey = "given_name"
|
||||
MiddleNameKey = "middle_name"
|
||||
FamilyNameKey = "family_name"
|
||||
NicknameKey = "nickname"
|
||||
PreferredUsernameKey = "preferred_username"
|
||||
ProfileKey = "profile"
|
||||
PictureKey = "picture"
|
||||
WebsiteKey = "website"
|
||||
EmailKey = "email"
|
||||
EmailVerifiedKey = "email_verified"
|
||||
GenderKey = "gender"
|
||||
BirthdateKey = "birthdate"
|
||||
ZoneinfoKey = "zoneinfo"
|
||||
LocaleKey = "locale"
|
||||
PhoneNumberKey = "phone_number"
|
||||
PhoneNumberVerifiedKey = "phone_number_verified"
|
||||
AddressKey = "address"
|
||||
UpdatedAtKey = "updated_at"
|
||||
)
|
||||
|
||||
type Token interface {
|
||||
Audience() []string
|
||||
Expiration() time.Time
|
||||
IssuedAt() time.Time
|
||||
Issuer() string
|
||||
JwtID() string
|
||||
NotBefore() time.Time
|
||||
Subject() string
|
||||
Name() string
|
||||
GivenName() string
|
||||
MiddleName() string
|
||||
FamilyName() string
|
||||
Nickname() string
|
||||
PreferredUsername() string
|
||||
Profile() string
|
||||
Picture() string
|
||||
Website() string
|
||||
Email() string
|
||||
EmailVerified() bool
|
||||
Gender() string
|
||||
Birthdate() *BirthdateClaim
|
||||
Zoneinfo() string
|
||||
Locale() string
|
||||
PhoneNumber() string
|
||||
PhoneNumberVerified() bool
|
||||
Address() *AddressClaim
|
||||
UpdatedAt() time.Time
|
||||
PrivateClaims() map[string]interface{}
|
||||
Get(string) (interface{}, bool)
|
||||
Set(string, interface{}) error
|
||||
Iterate(context.Context) Iterator
|
||||
Walk(context.Context, Visitor) error
|
||||
AsMap(context.Context) (map[string]interface{}, error)
|
||||
}
|
||||
type stdToken struct {
|
||||
audience types.StringList // https://tools.ietf.org/html/rfc7519#section-4.1.3
|
||||
expiration *types.NumericDate // https://tools.ietf.org/html/rfc7519#section-4.1.4
|
||||
issuedAt *types.NumericDate // https://tools.ietf.org/html/rfc7519#section-4.1.6
|
||||
issuer *string // https://tools.ietf.org/html/rfc7519#section-4.1.1
|
||||
jwtID *string // https://tools.ietf.org/html/rfc7519#section-4.1.7
|
||||
notBefore *types.NumericDate // https://tools.ietf.org/html/rfc7519#section-4.1.5
|
||||
subject *string // https://tools.ietf.org/html/rfc7519#section-4.1.2
|
||||
name *string //
|
||||
givenName *string //
|
||||
middleName *string //
|
||||
familyName *string //
|
||||
nickname *string //
|
||||
preferredUsername *string //
|
||||
profile *string //
|
||||
picture *string //
|
||||
website *string //
|
||||
email *string //
|
||||
emailVerified *bool //
|
||||
gender *string //
|
||||
birthdate *BirthdateClaim //
|
||||
zoneinfo *string //
|
||||
locale *string //
|
||||
phoneNumber *string //
|
||||
phoneNumberVerified *bool //
|
||||
address *AddressClaim //
|
||||
updatedAt *types.NumericDate //
|
||||
privateClaims map[string]interface{} `json:"-"`
|
||||
}
|
||||
|
||||
type openidTokenMarshalProxy struct {
|
||||
Xaudience types.StringList `json:"aud,omitempty"`
|
||||
Xexpiration *types.NumericDate `json:"exp,omitempty"`
|
||||
XissuedAt *types.NumericDate `json:"iat,omitempty"`
|
||||
Xissuer *string `json:"iss,omitempty"`
|
||||
XjwtID *string `json:"jti,omitempty"`
|
||||
XnotBefore *types.NumericDate `json:"nbf,omitempty"`
|
||||
Xsubject *string `json:"sub,omitempty"`
|
||||
Xname *string `json:"name,omitempty"`
|
||||
XgivenName *string `json:"given_name,omitempty"`
|
||||
XmiddleName *string `json:"middle_name,omitempty"`
|
||||
XfamilyName *string `json:"family_name,omitempty"`
|
||||
Xnickname *string `json:"nickname,omitempty"`
|
||||
XpreferredUsername *string `json:"preferred_username,omitempty"`
|
||||
Xprofile *string `json:"profile,omitempty"`
|
||||
Xpicture *string `json:"picture,omitempty"`
|
||||
Xwebsite *string `json:"website,omitempty"`
|
||||
Xemail *string `json:"email,omitempty"`
|
||||
XemailVerified *bool `json:"email_verified,omitempty"`
|
||||
Xgender *string `json:"gender,omitempty"`
|
||||
Xbirthdate *BirthdateClaim `json:"birthdate,omitempty"`
|
||||
Xzoneinfo *string `json:"zoneinfo,omitempty"`
|
||||
Xlocale *string `json:"locale,omitempty"`
|
||||
XphoneNumber *string `json:"phone_number,omitempty"`
|
||||
XphoneNumberVerified *bool `json:"phone_number_verified,omitempty"`
|
||||
Xaddress *AddressClaim `json:"address,omitempty"`
|
||||
XupdatedAt *types.NumericDate `json:"updated_at,omitempty"`
|
||||
}
|
||||
|
||||
// New creates a standard token, with minimal knowledge of
|
||||
// possible claims. Standard claims include"aud", "exp", "iat", "iss", "jti", "nbf", "sub", "name", "given_name", "middle_name", "family_name", "nickname", "preferred_username", "profile", "picture", "website", "email", "email_verified", "gender", "birthdate", "zoneinfo", "locale", "phone_number", "phone_number_verified", "address" and "updated_at".
|
||||
// Convenience accessors are provided for these standard claims
|
||||
func New() Token {
|
||||
return &stdToken{
|
||||
privateClaims: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Size returns the number of valid claims stored in this token
|
||||
func (t *stdToken) Size() int {
|
||||
var count int
|
||||
if len(t.audience) > 0 {
|
||||
count++
|
||||
}
|
||||
if t.birthdate != nil {
|
||||
count++
|
||||
}
|
||||
if t.address != nil {
|
||||
count++
|
||||
}
|
||||
count += len(t.privateClaims)
|
||||
return count
|
||||
}
|
||||
|
||||
func (t *stdToken) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case AudienceKey:
|
||||
if t.audience == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.audience.Get()
|
||||
return v, true
|
||||
case ExpirationKey:
|
||||
if t.expiration == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.expiration.Get()
|
||||
return v, true
|
||||
case IssuedAtKey:
|
||||
if t.issuedAt == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.issuedAt.Get()
|
||||
return v, true
|
||||
case IssuerKey:
|
||||
if t.issuer == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.issuer)
|
||||
return v, true
|
||||
case JwtIDKey:
|
||||
if t.jwtID == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.jwtID)
|
||||
return v, true
|
||||
case NotBeforeKey:
|
||||
if t.notBefore == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.notBefore.Get()
|
||||
return v, true
|
||||
case SubjectKey:
|
||||
if t.subject == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.subject)
|
||||
return v, true
|
||||
case NameKey:
|
||||
if t.name == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.name)
|
||||
return v, true
|
||||
case GivenNameKey:
|
||||
if t.givenName == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.givenName)
|
||||
return v, true
|
||||
case MiddleNameKey:
|
||||
if t.middleName == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.middleName)
|
||||
return v, true
|
||||
case FamilyNameKey:
|
||||
if t.familyName == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.familyName)
|
||||
return v, true
|
||||
case NicknameKey:
|
||||
if t.nickname == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.nickname)
|
||||
return v, true
|
||||
case PreferredUsernameKey:
|
||||
if t.preferredUsername == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.preferredUsername)
|
||||
return v, true
|
||||
case ProfileKey:
|
||||
if t.profile == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.profile)
|
||||
return v, true
|
||||
case PictureKey:
|
||||
if t.picture == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.picture)
|
||||
return v, true
|
||||
case WebsiteKey:
|
||||
if t.website == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.website)
|
||||
return v, true
|
||||
case EmailKey:
|
||||
if t.email == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.email)
|
||||
return v, true
|
||||
case EmailVerifiedKey:
|
||||
if t.emailVerified == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.emailVerified)
|
||||
return v, true
|
||||
case GenderKey:
|
||||
if t.gender == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.gender)
|
||||
return v, true
|
||||
case BirthdateKey:
|
||||
if t.birthdate == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.birthdate
|
||||
return v, true
|
||||
case ZoneinfoKey:
|
||||
if t.zoneinfo == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.zoneinfo)
|
||||
return v, true
|
||||
case LocaleKey:
|
||||
if t.locale == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.locale)
|
||||
return v, true
|
||||
case PhoneNumberKey:
|
||||
if t.phoneNumber == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.phoneNumber)
|
||||
return v, true
|
||||
case PhoneNumberVerifiedKey:
|
||||
if t.phoneNumberVerified == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.phoneNumberVerified)
|
||||
return v, true
|
||||
case AddressKey:
|
||||
if t.address == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.address
|
||||
return v, true
|
||||
case UpdatedAtKey:
|
||||
if t.updatedAt == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.updatedAt.Get()
|
||||
return v, true
|
||||
default:
|
||||
v, ok := t.privateClaims[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *stdToken) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case AudienceKey:
|
||||
var acceptor types.StringList
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, AudienceKey)
|
||||
}
|
||||
h.audience = acceptor
|
||||
return nil
|
||||
case ExpirationKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, ExpirationKey)
|
||||
}
|
||||
h.expiration = &acceptor
|
||||
return nil
|
||||
case IssuedAtKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, IssuedAtKey)
|
||||
}
|
||||
h.issuedAt = &acceptor
|
||||
return nil
|
||||
case IssuerKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.issuer = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, IssuerKey, value)
|
||||
case JwtIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.jwtID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, JwtIDKey, value)
|
||||
case NotBeforeKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, NotBeforeKey)
|
||||
}
|
||||
h.notBefore = &acceptor
|
||||
return nil
|
||||
case SubjectKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.subject = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, SubjectKey, value)
|
||||
case NameKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.name = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, NameKey, value)
|
||||
case GivenNameKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.givenName = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, GivenNameKey, value)
|
||||
case MiddleNameKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.middleName = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, MiddleNameKey, value)
|
||||
case FamilyNameKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.familyName = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, FamilyNameKey, value)
|
||||
case NicknameKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.nickname = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, NicknameKey, value)
|
||||
case PreferredUsernameKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.preferredUsername = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, PreferredUsernameKey, value)
|
||||
case ProfileKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.profile = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ProfileKey, value)
|
||||
case PictureKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.picture = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, PictureKey, value)
|
||||
case WebsiteKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.website = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, WebsiteKey, value)
|
||||
case EmailKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.email = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, EmailKey, value)
|
||||
case EmailVerifiedKey:
|
||||
if v, ok := value.(bool); ok {
|
||||
h.emailVerified = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, EmailVerifiedKey, value)
|
||||
case GenderKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.gender = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, GenderKey, value)
|
||||
case BirthdateKey:
|
||||
var acceptor BirthdateClaim
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, BirthdateKey)
|
||||
}
|
||||
h.birthdate = &acceptor
|
||||
return nil
|
||||
case ZoneinfoKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.zoneinfo = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, ZoneinfoKey, value)
|
||||
case LocaleKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.locale = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, LocaleKey, value)
|
||||
case PhoneNumberKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.phoneNumber = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, PhoneNumberKey, value)
|
||||
case PhoneNumberVerifiedKey:
|
||||
if v, ok := value.(bool); ok {
|
||||
h.phoneNumberVerified = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, PhoneNumberVerifiedKey, value)
|
||||
case AddressKey:
|
||||
var acceptor AddressClaim
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, AddressKey)
|
||||
}
|
||||
h.address = &acceptor
|
||||
return nil
|
||||
case UpdatedAtKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, UpdatedAtKey)
|
||||
}
|
||||
h.updatedAt = &acceptor
|
||||
return nil
|
||||
default:
|
||||
if h.privateClaims == nil {
|
||||
h.privateClaims = map[string]interface{}{}
|
||||
}
|
||||
h.privateClaims[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdToken) Audience() []string {
|
||||
if h.audience != nil {
|
||||
return h.audience.Get()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdToken) Expiration() time.Time {
|
||||
if h.expiration != nil {
|
||||
return h.expiration.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (h *stdToken) IssuedAt() time.Time {
|
||||
if h.issuedAt != nil {
|
||||
return h.issuedAt.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (h *stdToken) Issuer() string {
|
||||
if h.issuer != nil {
|
||||
return *(h.issuer)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) JwtID() string {
|
||||
if h.jwtID != nil {
|
||||
return *(h.jwtID)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) NotBefore() time.Time {
|
||||
if h.notBefore != nil {
|
||||
return h.notBefore.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (h *stdToken) Subject() string {
|
||||
if h.subject != nil {
|
||||
return *(h.subject)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Name() string {
|
||||
if h.name != nil {
|
||||
return *(h.name)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) GivenName() string {
|
||||
if h.givenName != nil {
|
||||
return *(h.givenName)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) MiddleName() string {
|
||||
if h.middleName != nil {
|
||||
return *(h.middleName)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) FamilyName() string {
|
||||
if h.familyName != nil {
|
||||
return *(h.familyName)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Nickname() string {
|
||||
if h.nickname != nil {
|
||||
return *(h.nickname)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) PreferredUsername() string {
|
||||
if h.preferredUsername != nil {
|
||||
return *(h.preferredUsername)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Profile() string {
|
||||
if h.profile != nil {
|
||||
return *(h.profile)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Picture() string {
|
||||
if h.picture != nil {
|
||||
return *(h.picture)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Website() string {
|
||||
if h.website != nil {
|
||||
return *(h.website)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Email() string {
|
||||
if h.email != nil {
|
||||
return *(h.email)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) EmailVerified() bool {
|
||||
if h.emailVerified != nil {
|
||||
return *(h.emailVerified)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *stdToken) Gender() string {
|
||||
if h.gender != nil {
|
||||
return *(h.gender)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Birthdate() *BirthdateClaim {
|
||||
return h.birthdate
|
||||
}
|
||||
|
||||
func (h *stdToken) Zoneinfo() string {
|
||||
if h.zoneinfo != nil {
|
||||
return *(h.zoneinfo)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) Locale() string {
|
||||
if h.locale != nil {
|
||||
return *(h.locale)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) PhoneNumber() string {
|
||||
if h.phoneNumber != nil {
|
||||
return *(h.phoneNumber)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) PhoneNumberVerified() bool {
|
||||
if h.phoneNumberVerified != nil {
|
||||
return *(h.phoneNumberVerified)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *stdToken) Address() *AddressClaim {
|
||||
return h.address
|
||||
}
|
||||
|
||||
func (h *stdToken) UpdatedAt() time.Time {
|
||||
if h.updatedAt != nil {
|
||||
return h.updatedAt.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (t *stdToken) PrivateClaims() map[string]interface{} {
|
||||
return t.privateClaims
|
||||
}
|
||||
|
||||
func (h *stdToken) iterate(ctx context.Context, ch chan *ClaimPair) {
|
||||
defer close(ch)
|
||||
|
||||
var pairs []*ClaimPair
|
||||
if h.audience != nil {
|
||||
v := h.audience.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: AudienceKey, Value: v})
|
||||
}
|
||||
if h.expiration != nil {
|
||||
v := h.expiration.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: ExpirationKey, Value: v})
|
||||
}
|
||||
if h.issuedAt != nil {
|
||||
v := h.issuedAt.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: IssuedAtKey, Value: v})
|
||||
}
|
||||
if h.issuer != nil {
|
||||
v := *(h.issuer)
|
||||
pairs = append(pairs, &ClaimPair{Key: IssuerKey, Value: v})
|
||||
}
|
||||
if h.jwtID != nil {
|
||||
v := *(h.jwtID)
|
||||
pairs = append(pairs, &ClaimPair{Key: JwtIDKey, Value: v})
|
||||
}
|
||||
if h.notBefore != nil {
|
||||
v := h.notBefore.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: NotBeforeKey, Value: v})
|
||||
}
|
||||
if h.subject != nil {
|
||||
v := *(h.subject)
|
||||
pairs = append(pairs, &ClaimPair{Key: SubjectKey, Value: v})
|
||||
}
|
||||
if h.name != nil {
|
||||
v := *(h.name)
|
||||
pairs = append(pairs, &ClaimPair{Key: NameKey, Value: v})
|
||||
}
|
||||
if h.givenName != nil {
|
||||
v := *(h.givenName)
|
||||
pairs = append(pairs, &ClaimPair{Key: GivenNameKey, Value: v})
|
||||
}
|
||||
if h.middleName != nil {
|
||||
v := *(h.middleName)
|
||||
pairs = append(pairs, &ClaimPair{Key: MiddleNameKey, Value: v})
|
||||
}
|
||||
if h.familyName != nil {
|
||||
v := *(h.familyName)
|
||||
pairs = append(pairs, &ClaimPair{Key: FamilyNameKey, Value: v})
|
||||
}
|
||||
if h.nickname != nil {
|
||||
v := *(h.nickname)
|
||||
pairs = append(pairs, &ClaimPair{Key: NicknameKey, Value: v})
|
||||
}
|
||||
if h.preferredUsername != nil {
|
||||
v := *(h.preferredUsername)
|
||||
pairs = append(pairs, &ClaimPair{Key: PreferredUsernameKey, Value: v})
|
||||
}
|
||||
if h.profile != nil {
|
||||
v := *(h.profile)
|
||||
pairs = append(pairs, &ClaimPair{Key: ProfileKey, Value: v})
|
||||
}
|
||||
if h.picture != nil {
|
||||
v := *(h.picture)
|
||||
pairs = append(pairs, &ClaimPair{Key: PictureKey, Value: v})
|
||||
}
|
||||
if h.website != nil {
|
||||
v := *(h.website)
|
||||
pairs = append(pairs, &ClaimPair{Key: WebsiteKey, Value: v})
|
||||
}
|
||||
if h.email != nil {
|
||||
v := *(h.email)
|
||||
pairs = append(pairs, &ClaimPair{Key: EmailKey, Value: v})
|
||||
}
|
||||
if h.emailVerified != nil {
|
||||
v := *(h.emailVerified)
|
||||
pairs = append(pairs, &ClaimPair{Key: EmailVerifiedKey, Value: v})
|
||||
}
|
||||
if h.gender != nil {
|
||||
v := *(h.gender)
|
||||
pairs = append(pairs, &ClaimPair{Key: GenderKey, Value: v})
|
||||
}
|
||||
if h.birthdate != nil {
|
||||
v := h.birthdate
|
||||
pairs = append(pairs, &ClaimPair{Key: BirthdateKey, Value: v})
|
||||
}
|
||||
if h.zoneinfo != nil {
|
||||
v := *(h.zoneinfo)
|
||||
pairs = append(pairs, &ClaimPair{Key: ZoneinfoKey, Value: v})
|
||||
}
|
||||
if h.locale != nil {
|
||||
v := *(h.locale)
|
||||
pairs = append(pairs, &ClaimPair{Key: LocaleKey, Value: v})
|
||||
}
|
||||
if h.phoneNumber != nil {
|
||||
v := *(h.phoneNumber)
|
||||
pairs = append(pairs, &ClaimPair{Key: PhoneNumberKey, Value: v})
|
||||
}
|
||||
if h.phoneNumberVerified != nil {
|
||||
v := *(h.phoneNumberVerified)
|
||||
pairs = append(pairs, &ClaimPair{Key: PhoneNumberVerifiedKey, Value: v})
|
||||
}
|
||||
if h.address != nil {
|
||||
v := h.address
|
||||
pairs = append(pairs, &ClaimPair{Key: AddressKey, Value: v})
|
||||
}
|
||||
if h.updatedAt != nil {
|
||||
v := h.updatedAt.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: UpdatedAtKey, Value: v})
|
||||
}
|
||||
for k, v := range h.privateClaims {
|
||||
pairs = append(pairs, &ClaimPair{Key: k, Value: v})
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// this is almost identical to json.Encoder.Encode(), but we use Marshal
|
||||
// to avoid having to remove the trailing newline for each successive
|
||||
// call to Encode()
|
||||
func writeJSON(buf *bytes.Buffer, v interface{}, keyName string) error {
|
||||
if enc, err := json.Marshal(v); err != nil {
|
||||
return errors.Wrapf(err, `failed to encode '%s'`, keyName)
|
||||
} else {
|
||||
buf.Write(enc)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdToken) UnmarshalJSON(buf []byte) error {
|
||||
var proxy openidTokenMarshalProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal stdToken`)
|
||||
}
|
||||
h.audience = proxy.Xaudience
|
||||
h.expiration = proxy.Xexpiration
|
||||
h.issuedAt = proxy.XissuedAt
|
||||
h.issuer = proxy.Xissuer
|
||||
h.jwtID = proxy.XjwtID
|
||||
h.notBefore = proxy.XnotBefore
|
||||
h.subject = proxy.Xsubject
|
||||
h.name = proxy.Xname
|
||||
h.givenName = proxy.XgivenName
|
||||
h.middleName = proxy.XmiddleName
|
||||
h.familyName = proxy.XfamilyName
|
||||
h.nickname = proxy.Xnickname
|
||||
h.preferredUsername = proxy.XpreferredUsername
|
||||
h.profile = proxy.Xprofile
|
||||
h.picture = proxy.Xpicture
|
||||
h.website = proxy.Xwebsite
|
||||
h.email = proxy.Xemail
|
||||
h.emailVerified = proxy.XemailVerified
|
||||
h.gender = proxy.Xgender
|
||||
h.birthdate = proxy.Xbirthdate
|
||||
h.zoneinfo = proxy.Xzoneinfo
|
||||
h.locale = proxy.Xlocale
|
||||
h.phoneNumber = proxy.XphoneNumber
|
||||
h.phoneNumberVerified = proxy.XphoneNumberVerified
|
||||
h.address = proxy.Xaddress
|
||||
h.updatedAt = proxy.XupdatedAt
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to parse privsate parameters`)
|
||||
}
|
||||
delete(m, AudienceKey)
|
||||
delete(m, ExpirationKey)
|
||||
delete(m, IssuedAtKey)
|
||||
delete(m, IssuerKey)
|
||||
delete(m, JwtIDKey)
|
||||
delete(m, NotBeforeKey)
|
||||
delete(m, SubjectKey)
|
||||
delete(m, NameKey)
|
||||
delete(m, GivenNameKey)
|
||||
delete(m, MiddleNameKey)
|
||||
delete(m, FamilyNameKey)
|
||||
delete(m, NicknameKey)
|
||||
delete(m, PreferredUsernameKey)
|
||||
delete(m, ProfileKey)
|
||||
delete(m, PictureKey)
|
||||
delete(m, WebsiteKey)
|
||||
delete(m, EmailKey)
|
||||
delete(m, EmailVerifiedKey)
|
||||
delete(m, GenderKey)
|
||||
delete(m, BirthdateKey)
|
||||
delete(m, ZoneinfoKey)
|
||||
delete(m, LocaleKey)
|
||||
delete(m, PhoneNumberKey)
|
||||
delete(m, PhoneNumberVerifiedKey)
|
||||
delete(m, AddressKey)
|
||||
delete(m, UpdatedAtKey)
|
||||
h.privateClaims = m
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h stdToken) MarshalJSON() ([]byte, error) {
|
||||
var proxy openidTokenMarshalProxy
|
||||
proxy.Xaudience = h.audience
|
||||
proxy.Xexpiration = h.expiration
|
||||
proxy.XissuedAt = h.issuedAt
|
||||
proxy.Xissuer = h.issuer
|
||||
proxy.XjwtID = h.jwtID
|
||||
proxy.XnotBefore = h.notBefore
|
||||
proxy.Xsubject = h.subject
|
||||
proxy.Xname = h.name
|
||||
proxy.XgivenName = h.givenName
|
||||
proxy.XmiddleName = h.middleName
|
||||
proxy.XfamilyName = h.familyName
|
||||
proxy.Xnickname = h.nickname
|
||||
proxy.XpreferredUsername = h.preferredUsername
|
||||
proxy.Xprofile = h.profile
|
||||
proxy.Xpicture = h.picture
|
||||
proxy.Xwebsite = h.website
|
||||
proxy.Xemail = h.email
|
||||
proxy.XemailVerified = h.emailVerified
|
||||
proxy.Xgender = h.gender
|
||||
proxy.Xbirthdate = h.birthdate
|
||||
proxy.Xzoneinfo = h.zoneinfo
|
||||
proxy.Xlocale = h.locale
|
||||
proxy.XphoneNumber = h.phoneNumber
|
||||
proxy.XphoneNumberVerified = h.phoneNumberVerified
|
||||
proxy.Xaddress = h.address
|
||||
proxy.XupdatedAt = h.updatedAt
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
if err := enc.Encode(proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to encode proxy to JSON`)
|
||||
}
|
||||
hasContent := buf.Len() > 3 // encoding/json always adds a newline, so "{}\n" is the empty hash
|
||||
if l := len(h.privateClaims); l > 0 {
|
||||
buf.Truncate(buf.Len() - 2)
|
||||
keys := make([]string, 0, l)
|
||||
for k := range h.privateClaims {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for i, k := range keys {
|
||||
if hasContent || i > 0 {
|
||||
fmt.Fprintf(&buf, `,`)
|
||||
}
|
||||
fmt.Fprintf(&buf, `%s:`, strconv.Quote(k))
|
||||
if err := enc.Encode(h.privateClaims[k]); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to encode private param %s`, k)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(&buf, `}`)
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func (h *stdToken) Iterate(ctx context.Context) Iterator {
|
||||
ch := make(chan *ClaimPair)
|
||||
go h.iterate(ctx, ch)
|
||||
return mapiter.New(ch)
|
||||
}
|
||||
|
||||
func (h *stdToken) Walk(ctx context.Context, visitor Visitor) error {
|
||||
return iter.WalkMap(ctx, h, visitor)
|
||||
}
|
||||
|
||||
func (h *stdToken) AsMap(ctx context.Context) (map[string]interface{}, error) {
|
||||
return iter.AsMap(ctx, h)
|
||||
}
|
||||
+54
@@ -0,0 +1,54 @@
|
||||
package jwt
|
||||
|
||||
import (
|
||||
"github.com/lestrrat-go/jwx/internal/option"
|
||||
"github.com/lestrrat-go/jwx/jwa"
|
||||
"github.com/lestrrat-go/jwx/jwt/openid"
|
||||
)
|
||||
|
||||
type Option = option.Interface
|
||||
|
||||
const (
|
||||
optkeyVerify = `verify`
|
||||
optkeyToken = `token`
|
||||
)
|
||||
|
||||
type VerifyParameters interface {
|
||||
Algorithm() jwa.SignatureAlgorithm
|
||||
Key() interface{}
|
||||
}
|
||||
|
||||
type verifyParams struct {
|
||||
alg jwa.SignatureAlgorithm
|
||||
key interface{}
|
||||
}
|
||||
|
||||
func (p *verifyParams) Algorithm() jwa.SignatureAlgorithm {
|
||||
return p.alg
|
||||
}
|
||||
|
||||
func (p *verifyParams) Key() interface{} {
|
||||
return p.key
|
||||
}
|
||||
|
||||
func WithVerify(alg jwa.SignatureAlgorithm, key interface{}) Option {
|
||||
return option.New(optkeyVerify, &verifyParams{
|
||||
alg: alg,
|
||||
key: key,
|
||||
})
|
||||
}
|
||||
|
||||
// WithToken specifies the token instance that is used when parsing
|
||||
// JWT tokens.
|
||||
func WithToken(t Token) Option {
|
||||
return option.New(optkeyToken, t)
|
||||
}
|
||||
|
||||
// WithOpenIDClaims is passed to the various JWT parsing functions, and
|
||||
// specifies that it should use an instance of `openid.Token` as the
|
||||
// destination to store the parsed results.
|
||||
//
|
||||
// This is exactly equivalent to specifying `jwt.WithToken(openid.New())`
|
||||
func WithOpenIDClaims() Option {
|
||||
return WithToken(openid.New())
|
||||
}
|
||||
+1
@@ -0,0 +1 @@
|
||||
package jwt
|
||||
+387
@@ -0,0 +1,387 @@
|
||||
// This file is auto-generated by jwt/internal/cmd/gentoken/main.go. DO NOT EDIT
|
||||
package jwt
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/lestrrat-go/iter/mapiter"
|
||||
"github.com/lestrrat-go/jwx/internal/iter"
|
||||
"github.com/lestrrat-go/jwx/jwt/internal/types"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
AudienceKey = "aud"
|
||||
ExpirationKey = "exp"
|
||||
IssuedAtKey = "iat"
|
||||
IssuerKey = "iss"
|
||||
JwtIDKey = "jti"
|
||||
NotBeforeKey = "nbf"
|
||||
SubjectKey = "sub"
|
||||
)
|
||||
|
||||
// Token represents a generic JWT token.
|
||||
// which are type-aware (to an extent). Other claims may be accessed via the `Get`/`Set`
|
||||
// methods but their types are not taken into consideration at all. If you have non-standard
|
||||
// claims that you must frequently access, consider creating accessors functions
|
||||
// like the following
|
||||
//
|
||||
// func SetFoo(tok jwt.Token) error
|
||||
// func GetFoo(tok jwt.Token) (*Customtyp, error)
|
||||
//
|
||||
// Embedding jwt.Token into another struct is not recommended, becase
|
||||
// jwt.Token needs to handle private claims, and this really does not
|
||||
// work well when it is embedded in other structure
|
||||
type Token interface {
|
||||
Audience() []string
|
||||
Expiration() time.Time
|
||||
IssuedAt() time.Time
|
||||
Issuer() string
|
||||
JwtID() string
|
||||
NotBefore() time.Time
|
||||
Subject() string
|
||||
PrivateClaims() map[string]interface{}
|
||||
Get(string) (interface{}, bool)
|
||||
Set(string, interface{}) error
|
||||
Iterate(context.Context) Iterator
|
||||
Walk(context.Context, Visitor) error
|
||||
AsMap(context.Context) (map[string]interface{}, error)
|
||||
}
|
||||
type stdToken struct {
|
||||
audience types.StringList // https://tools.ietf.org/html/rfc7519#section-4.1.3
|
||||
expiration *types.NumericDate // https://tools.ietf.org/html/rfc7519#section-4.1.4
|
||||
issuedAt *types.NumericDate // https://tools.ietf.org/html/rfc7519#section-4.1.6
|
||||
issuer *string // https://tools.ietf.org/html/rfc7519#section-4.1.1
|
||||
jwtID *string // https://tools.ietf.org/html/rfc7519#section-4.1.7
|
||||
notBefore *types.NumericDate // https://tools.ietf.org/html/rfc7519#section-4.1.5
|
||||
subject *string // https://tools.ietf.org/html/rfc7519#section-4.1.2
|
||||
privateClaims map[string]interface{} `json:"-"`
|
||||
}
|
||||
|
||||
type stdTokenMarshalProxy struct {
|
||||
Xaudience types.StringList `json:"aud,omitempty"`
|
||||
Xexpiration *types.NumericDate `json:"exp,omitempty"`
|
||||
XissuedAt *types.NumericDate `json:"iat,omitempty"`
|
||||
Xissuer *string `json:"iss,omitempty"`
|
||||
XjwtID *string `json:"jti,omitempty"`
|
||||
XnotBefore *types.NumericDate `json:"nbf,omitempty"`
|
||||
Xsubject *string `json:"sub,omitempty"`
|
||||
}
|
||||
|
||||
// New creates a standard token, with minimal knowledge of
|
||||
// possible claims. Standard claims include"aud", "exp", "iat", "iss", "jti", "nbf" and "sub".
|
||||
// Convenience accessors are provided for these standard claims
|
||||
func New() Token {
|
||||
return &stdToken{
|
||||
privateClaims: make(map[string]interface{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Size returns the number of valid claims stored in this token
|
||||
func (t *stdToken) Size() int {
|
||||
var count int
|
||||
if len(t.audience) > 0 {
|
||||
count++
|
||||
}
|
||||
count += len(t.privateClaims)
|
||||
return count
|
||||
}
|
||||
|
||||
func (t *stdToken) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case AudienceKey:
|
||||
if t.audience == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.audience.Get()
|
||||
return v, true
|
||||
case ExpirationKey:
|
||||
if t.expiration == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.expiration.Get()
|
||||
return v, true
|
||||
case IssuedAtKey:
|
||||
if t.issuedAt == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.issuedAt.Get()
|
||||
return v, true
|
||||
case IssuerKey:
|
||||
if t.issuer == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.issuer)
|
||||
return v, true
|
||||
case JwtIDKey:
|
||||
if t.jwtID == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.jwtID)
|
||||
return v, true
|
||||
case NotBeforeKey:
|
||||
if t.notBefore == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := t.notBefore.Get()
|
||||
return v, true
|
||||
case SubjectKey:
|
||||
if t.subject == nil {
|
||||
return nil, false
|
||||
}
|
||||
v := *(t.subject)
|
||||
return v, true
|
||||
default:
|
||||
v, ok := t.privateClaims[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *stdToken) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case AudienceKey:
|
||||
var acceptor types.StringList
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, AudienceKey)
|
||||
}
|
||||
h.audience = acceptor
|
||||
return nil
|
||||
case ExpirationKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, ExpirationKey)
|
||||
}
|
||||
h.expiration = &acceptor
|
||||
return nil
|
||||
case IssuedAtKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, IssuedAtKey)
|
||||
}
|
||||
h.issuedAt = &acceptor
|
||||
return nil
|
||||
case IssuerKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.issuer = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, IssuerKey, value)
|
||||
case JwtIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.jwtID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, JwtIDKey, value)
|
||||
case NotBeforeKey:
|
||||
var acceptor types.NumericDate
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, NotBeforeKey)
|
||||
}
|
||||
h.notBefore = &acceptor
|
||||
return nil
|
||||
case SubjectKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.subject = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, SubjectKey, value)
|
||||
default:
|
||||
if h.privateClaims == nil {
|
||||
h.privateClaims = map[string]interface{}{}
|
||||
}
|
||||
h.privateClaims[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdToken) Audience() []string {
|
||||
if h.audience != nil {
|
||||
return h.audience.Get()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdToken) Expiration() time.Time {
|
||||
if h.expiration != nil {
|
||||
return h.expiration.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (h *stdToken) IssuedAt() time.Time {
|
||||
if h.issuedAt != nil {
|
||||
return h.issuedAt.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (h *stdToken) Issuer() string {
|
||||
if h.issuer != nil {
|
||||
return *(h.issuer)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) JwtID() string {
|
||||
if h.jwtID != nil {
|
||||
return *(h.jwtID)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *stdToken) NotBefore() time.Time {
|
||||
if h.notBefore != nil {
|
||||
return h.notBefore.Get()
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
func (h *stdToken) Subject() string {
|
||||
if h.subject != nil {
|
||||
return *(h.subject)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (t *stdToken) PrivateClaims() map[string]interface{} {
|
||||
return t.privateClaims
|
||||
}
|
||||
|
||||
func (h *stdToken) iterate(ctx context.Context, ch chan *ClaimPair) {
|
||||
defer close(ch)
|
||||
|
||||
var pairs []*ClaimPair
|
||||
if h.audience != nil {
|
||||
v := h.audience.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: AudienceKey, Value: v})
|
||||
}
|
||||
if h.expiration != nil {
|
||||
v := h.expiration.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: ExpirationKey, Value: v})
|
||||
}
|
||||
if h.issuedAt != nil {
|
||||
v := h.issuedAt.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: IssuedAtKey, Value: v})
|
||||
}
|
||||
if h.issuer != nil {
|
||||
v := *(h.issuer)
|
||||
pairs = append(pairs, &ClaimPair{Key: IssuerKey, Value: v})
|
||||
}
|
||||
if h.jwtID != nil {
|
||||
v := *(h.jwtID)
|
||||
pairs = append(pairs, &ClaimPair{Key: JwtIDKey, Value: v})
|
||||
}
|
||||
if h.notBefore != nil {
|
||||
v := h.notBefore.Get()
|
||||
pairs = append(pairs, &ClaimPair{Key: NotBeforeKey, Value: v})
|
||||
}
|
||||
if h.subject != nil {
|
||||
v := *(h.subject)
|
||||
pairs = append(pairs, &ClaimPair{Key: SubjectKey, Value: v})
|
||||
}
|
||||
for k, v := range h.privateClaims {
|
||||
pairs = append(pairs, &ClaimPair{Key: k, Value: v})
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case ch <- pair:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// this is almost identical to json.Encoder.Encode(), but we use Marshal
|
||||
// to avoid having to remove the trailing newline for each successive
|
||||
// call to Encode()
|
||||
func writeJSON(buf *bytes.Buffer, v interface{}, keyName string) error {
|
||||
if enc, err := json.Marshal(v); err != nil {
|
||||
return errors.Wrapf(err, `failed to encode '%s'`, keyName)
|
||||
} else {
|
||||
buf.Write(enc)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *stdToken) UnmarshalJSON(buf []byte) error {
|
||||
var proxy stdTokenMarshalProxy
|
||||
if err := json.Unmarshal(buf, &proxy); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal stdToken`)
|
||||
}
|
||||
h.audience = proxy.Xaudience
|
||||
h.expiration = proxy.Xexpiration
|
||||
h.issuedAt = proxy.XissuedAt
|
||||
h.issuer = proxy.Xissuer
|
||||
h.jwtID = proxy.XjwtID
|
||||
h.notBefore = proxy.XnotBefore
|
||||
h.subject = proxy.Xsubject
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to parse privsate parameters`)
|
||||
}
|
||||
delete(m, AudienceKey)
|
||||
delete(m, ExpirationKey)
|
||||
delete(m, IssuedAtKey)
|
||||
delete(m, IssuerKey)
|
||||
delete(m, JwtIDKey)
|
||||
delete(m, NotBeforeKey)
|
||||
delete(m, SubjectKey)
|
||||
h.privateClaims = m
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h stdToken) MarshalJSON() ([]byte, error) {
|
||||
var proxy stdTokenMarshalProxy
|
||||
proxy.Xaudience = h.audience
|
||||
proxy.Xexpiration = h.expiration
|
||||
proxy.XissuedAt = h.issuedAt
|
||||
proxy.Xissuer = h.issuer
|
||||
proxy.XjwtID = h.jwtID
|
||||
proxy.XnotBefore = h.notBefore
|
||||
proxy.Xsubject = h.subject
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
if err := enc.Encode(proxy); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to encode proxy to JSON`)
|
||||
}
|
||||
hasContent := buf.Len() > 3 // encoding/json always adds a newline, so "{}\n" is the empty hash
|
||||
if l := len(h.privateClaims); l > 0 {
|
||||
buf.Truncate(buf.Len() - 2)
|
||||
keys := make([]string, 0, l)
|
||||
for k := range h.privateClaims {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for i, k := range keys {
|
||||
if hasContent || i > 0 {
|
||||
fmt.Fprintf(&buf, `,`)
|
||||
}
|
||||
fmt.Fprintf(&buf, `%s:`, strconv.Quote(k))
|
||||
if err := enc.Encode(h.privateClaims[k]); err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to encode private param %s`, k)
|
||||
}
|
||||
}
|
||||
fmt.Fprintf(&buf, `}`)
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
func (h *stdToken) Iterate(ctx context.Context) Iterator {
|
||||
ch := make(chan *ClaimPair)
|
||||
go h.iterate(ctx, ch)
|
||||
return mapiter.New(ch)
|
||||
}
|
||||
|
||||
func (h *stdToken) Walk(ctx context.Context, visitor Visitor) error {
|
||||
return iter.WalkMap(ctx, h, visitor)
|
||||
}
|
||||
|
||||
func (h *stdToken) AsMap(ctx context.Context) (map[string]interface{}, error) {
|
||||
return iter.AsMap(ctx, h)
|
||||
}
|
||||
+173
@@ -0,0 +1,173 @@
|
||||
package jwt
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/lestrrat-go/jwx/internal/option"
|
||||
)
|
||||
|
||||
const (
|
||||
optkeyAcceptableSkew = "acceptableSkew"
|
||||
optkeyClock = "clock"
|
||||
optkeyIssuer = "issuer"
|
||||
optkeySubject = "subject"
|
||||
optkeyAudience = "audience"
|
||||
optkeyJwtid = "jwtid"
|
||||
)
|
||||
|
||||
type Clock interface {
|
||||
Now() time.Time
|
||||
}
|
||||
type ClockFunc func() time.Time
|
||||
|
||||
func (f ClockFunc) Now() time.Time {
|
||||
return f()
|
||||
}
|
||||
|
||||
// WithClock specifies the `Clock` to be used when verifying
|
||||
// claims exp and nbf.
|
||||
func WithClock(c Clock) Option {
|
||||
return option.New(optkeyClock, c)
|
||||
}
|
||||
|
||||
// WithAcceptableSkew specifies the duration in which exp and nbf
|
||||
// claims may differ by. This value should be positive
|
||||
func WithAcceptableSkew(dur time.Duration) Option {
|
||||
return option.New(optkeyAcceptableSkew, dur)
|
||||
}
|
||||
|
||||
// WithIssuer specifies that expected issuer value. If not specified,
|
||||
// the value of issuer is not verified at all.
|
||||
func WithIssuer(s string) Option {
|
||||
return option.New(optkeyIssuer, s)
|
||||
}
|
||||
|
||||
// WithSubject specifies that expected subject value. If not specified,
|
||||
// the value of subject is not verified at all.
|
||||
func WithSubject(s string) Option {
|
||||
return option.New(optkeySubject, s)
|
||||
}
|
||||
|
||||
// WithJwtID specifies that expected jti value. If not specified,
|
||||
// the value of jti is not verified at all.
|
||||
func WithJwtID(s string) Option {
|
||||
return option.New(optkeyJwtid, s)
|
||||
}
|
||||
|
||||
// WithAudience specifies that expected audience value.
|
||||
// Verify will return true if one of the values in the `aud` element
|
||||
// matches this value. If not specified, the value of issuer is not
|
||||
// verified at all.
|
||||
func WithAudience(s string) Option {
|
||||
return option.New(optkeyAudience, s)
|
||||
}
|
||||
|
||||
// WithClaimValue specifies that expected any claim value.
|
||||
func WithClaimValue(name string, v interface{}) Option {
|
||||
return option.New(name, v)
|
||||
}
|
||||
|
||||
// Verify makes sure that the essential claims stand.
|
||||
//
|
||||
// See the various `WithXXX` functions for optional parameters
|
||||
// that can control the behavior of this method.
|
||||
func Verify(t Token, options ...Option) error {
|
||||
var issuer string
|
||||
var subject string
|
||||
var audience string
|
||||
var jwtid string
|
||||
var clock Clock = ClockFunc(time.Now)
|
||||
var skew time.Duration
|
||||
claimValues := make(map[string]interface{})
|
||||
for _, o := range options {
|
||||
switch o.Name() {
|
||||
case optkeyClock:
|
||||
clock = o.Value().(Clock)
|
||||
case optkeyAcceptableSkew:
|
||||
skew = o.Value().(time.Duration)
|
||||
case optkeyIssuer:
|
||||
issuer = o.Value().(string)
|
||||
case optkeySubject:
|
||||
subject = o.Value().(string)
|
||||
case optkeyAudience:
|
||||
audience = o.Value().(string)
|
||||
case optkeyJwtid:
|
||||
jwtid = o.Value().(string)
|
||||
default:
|
||||
claimValues[o.Name()] = o.Value()
|
||||
}
|
||||
}
|
||||
|
||||
// check for iss
|
||||
if len(issuer) > 0 {
|
||||
if v := t.Issuer(); v != "" && v != issuer {
|
||||
return errors.New(`iss not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
// check for jti
|
||||
if len(jwtid) > 0 {
|
||||
if v := t.JwtID(); v != "" && v != jwtid {
|
||||
return errors.New(`jti not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
// check for sub
|
||||
if len(subject) > 0 {
|
||||
if v := t.Subject(); v != "" && v != subject {
|
||||
return errors.New(`sub not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
// check for aud
|
||||
if len(audience) > 0 {
|
||||
var found bool
|
||||
for _, v := range t.Audience() {
|
||||
if v == audience {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return errors.New(`aud not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
// check for exp
|
||||
if tv := t.Expiration(); !tv.IsZero() {
|
||||
now := clock.Now().Truncate(time.Second)
|
||||
ttv := tv.Truncate(time.Second)
|
||||
if !now.Before(ttv.Add(skew)) {
|
||||
return errors.New(`exp not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
// check for iat
|
||||
if tv := t.IssuedAt(); !tv.IsZero() {
|
||||
now := clock.Now().Truncate(time.Second)
|
||||
ttv := tv.Truncate(time.Second)
|
||||
if now.Before(ttv.Add(-1 * skew)) {
|
||||
return errors.New(`iat not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
// check for nbf
|
||||
if tv := t.NotBefore(); !tv.IsZero() {
|
||||
now := clock.Now().Truncate(time.Second)
|
||||
ttv := tv.Truncate(time.Second)
|
||||
// now cannot be before t, so we check for now > t - skew
|
||||
if !now.After(ttv.Add(-1 * skew)) {
|
||||
return errors.New(`nbf not satisfied`)
|
||||
}
|
||||
}
|
||||
|
||||
for name, expectedValue := range claimValues {
|
||||
if v, ok := t.Get(name); !ok || v != expectedValue {
|
||||
return fmt.Errorf(`%v not satisfied`, name)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
The MIT License (MIT)
|
||||
|
||||
Copyright (c) 2015 lestrrat
|
||||
|
||||
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.
|
||||
|
||||
+28
@@ -0,0 +1,28 @@
|
||||
package base64
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
)
|
||||
|
||||
func EncodeToString(src []byte) string {
|
||||
return base64.RawURLEncoding.EncodeToString(src)
|
||||
}
|
||||
|
||||
func EncodeUint64ToString(v uint64) string {
|
||||
data := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(data, v)
|
||||
|
||||
i := 0
|
||||
for ; i < len(data); i++ {
|
||||
if data[i] != 0x0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return EncodeToString(data[i:])
|
||||
}
|
||||
|
||||
func DecodeString(src string) ([]byte, error) {
|
||||
return base64.RawURLEncoding.DecodeString(src)
|
||||
}
|
||||
+43
@@ -0,0 +1,43 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// CompressionAlgorithm represents the compression algorithms as described in https://tools.ietf.org/html/rfc7518#section-7.3
|
||||
type CompressionAlgorithm string
|
||||
|
||||
// Supported values for CompressionAlgorithm
|
||||
const (
|
||||
Deflate CompressionAlgorithm = "DEF" // DEFLATE (RFC 1951)
|
||||
NoCompress CompressionAlgorithm = "" // No compression
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *CompressionAlgorithm) Accept(value interface{}) error {
|
||||
var tmp CompressionAlgorithm
|
||||
switch x := value.(type) {
|
||||
case string:
|
||||
tmp = CompressionAlgorithm(x)
|
||||
case CompressionAlgorithm:
|
||||
tmp = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.CompressionAlgorithm: %T`, value)
|
||||
}
|
||||
switch tmp {
|
||||
case Deflate, NoCompress:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.CompressionAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a CompressionAlgorithm
|
||||
func (v CompressionAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// ContentEncryptionAlgorithm represents the various encryption algorithms as described in https://tools.ietf.org/html/rfc7518#section-5
|
||||
type ContentEncryptionAlgorithm string
|
||||
|
||||
// Supported values for ContentEncryptionAlgorithm
|
||||
const (
|
||||
A128CBC_HS256 ContentEncryptionAlgorithm = "A128CBC-HS256" // AES-CBC + HMAC-SHA256 (128)
|
||||
A128GCM ContentEncryptionAlgorithm = "A128GCM" // AES-GCM (128)
|
||||
A192CBC_HS384 ContentEncryptionAlgorithm = "A192CBC-HS384" // AES-CBC + HMAC-SHA384 (192)
|
||||
A192GCM ContentEncryptionAlgorithm = "A192GCM" // AES-GCM (192)
|
||||
A256CBC_HS512 ContentEncryptionAlgorithm = "A256CBC-HS512" // AES-CBC + HMAC-SHA512 (256)
|
||||
A256GCM ContentEncryptionAlgorithm = "A256GCM" // AES-GCM (256)
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *ContentEncryptionAlgorithm) Accept(value interface{}) error {
|
||||
var tmp ContentEncryptionAlgorithm
|
||||
switch x := value.(type) {
|
||||
case string:
|
||||
tmp = ContentEncryptionAlgorithm(x)
|
||||
case ContentEncryptionAlgorithm:
|
||||
tmp = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.ContentEncryptionAlgorithm: %T`, value)
|
||||
}
|
||||
switch tmp {
|
||||
case A128CBC_HS256, A128GCM, A192CBC_HS384, A192GCM, A256CBC_HS512, A256GCM:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.ContentEncryptionAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a ContentEncryptionAlgorithm
|
||||
func (v ContentEncryptionAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// EllipticCurveAlgorithm represents the algorithms used for EC keys
|
||||
type EllipticCurveAlgorithm string
|
||||
|
||||
// Supported values for EllipticCurveAlgorithm
|
||||
const (
|
||||
P256 EllipticCurveAlgorithm = "P-256"
|
||||
P384 EllipticCurveAlgorithm = "P-384"
|
||||
P521 EllipticCurveAlgorithm = "P-521"
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *EllipticCurveAlgorithm) Accept(value interface{}) error {
|
||||
var tmp EllipticCurveAlgorithm
|
||||
switch x := value.(type) {
|
||||
case string:
|
||||
tmp = EllipticCurveAlgorithm(x)
|
||||
case EllipticCurveAlgorithm:
|
||||
tmp = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.EllipticCurveAlgorithm: %T`, value)
|
||||
}
|
||||
switch tmp {
|
||||
case P256, P384, P521:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.EllipticCurveAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a EllipticCurveAlgorithm
|
||||
func (v EllipticCurveAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+17
@@ -0,0 +1,17 @@
|
||||
//go:generate go run internal/cmd/gentypes/main.go
|
||||
|
||||
// Package jwa defines the various algorithm described in https://tools.ietf.org/html/rfc7518
|
||||
package jwa
|
||||
|
||||
// Size returns the size of the EllipticCurveAlgorithm
|
||||
func (crv EllipticCurveAlgorithm) Size() int {
|
||||
switch crv {
|
||||
case P256:
|
||||
return 32
|
||||
case P384:
|
||||
return 48
|
||||
case P521:
|
||||
return 66
|
||||
}
|
||||
return 0
|
||||
}
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// KeyEncryptionAlgorithm represents the various encryption algorithms as described in https://tools.ietf.org/html/rfc7518#section-4.1
|
||||
type KeyEncryptionAlgorithm string
|
||||
|
||||
// Supported values for KeyEncryptionAlgorithm
|
||||
const (
|
||||
A128GCMKW KeyEncryptionAlgorithm = "A128GCMKW" // AES-GCM key wrap (128)
|
||||
A128KW KeyEncryptionAlgorithm = "A128KW" // AES key wrap (128)
|
||||
A192GCMKW KeyEncryptionAlgorithm = "A192GCMKW" // AES-GCM key wrap (192)
|
||||
A192KW KeyEncryptionAlgorithm = "A192KW" // AES key wrap (192)
|
||||
A256GCMKW KeyEncryptionAlgorithm = "A256GCMKW" // AES-GCM key wrap (256)
|
||||
A256KW KeyEncryptionAlgorithm = "A256KW" // AES key wrap (256)
|
||||
DIRECT KeyEncryptionAlgorithm = "dir" // Direct encryption
|
||||
ECDH_ES KeyEncryptionAlgorithm = "ECDH-ES" // ECDH-ES
|
||||
ECDH_ES_A128KW KeyEncryptionAlgorithm = "ECDH-ES+A128KW" // ECDH-ES + AES key wrap (128)
|
||||
ECDH_ES_A192KW KeyEncryptionAlgorithm = "ECDH-ES+A192KW" // ECDH-ES + AES key wrap (192)
|
||||
ECDH_ES_A256KW KeyEncryptionAlgorithm = "ECDH-ES+A256KW" // ECDH-ES + AES key wrap (256)
|
||||
PBES2_HS256_A128KW KeyEncryptionAlgorithm = "PBES2-HS256+A128KW" // PBES2 + HMAC-SHA256 + AES key wrap (128)
|
||||
PBES2_HS384_A192KW KeyEncryptionAlgorithm = "PBES2-HS384+A192KW" // PBES2 + HMAC-SHA384 + AES key wrap (192)
|
||||
PBES2_HS512_A256KW KeyEncryptionAlgorithm = "PBES2-HS512+A256KW" // PBES2 + HMAC-SHA512 + AES key wrap (256)
|
||||
RSA1_5 KeyEncryptionAlgorithm = "RSA1_5" // RSA-PKCS1v1.5
|
||||
RSA_OAEP KeyEncryptionAlgorithm = "RSA-OAEP" // RSA-OAEP-SHA1
|
||||
RSA_OAEP_256 KeyEncryptionAlgorithm = "RSA-OAEP-256" // RSA-OAEP-SHA256
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *KeyEncryptionAlgorithm) Accept(value interface{}) error {
|
||||
var tmp KeyEncryptionAlgorithm
|
||||
switch x := value.(type) {
|
||||
case string:
|
||||
tmp = KeyEncryptionAlgorithm(x)
|
||||
case KeyEncryptionAlgorithm:
|
||||
tmp = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.KeyEncryptionAlgorithm: %T`, value)
|
||||
}
|
||||
switch tmp {
|
||||
case A128GCMKW, A128KW, A192GCMKW, A192KW, A256GCMKW, A256KW, DIRECT, ECDH_ES, ECDH_ES_A128KW, ECDH_ES_A192KW, ECDH_ES_A256KW, PBES2_HS256_A128KW, PBES2_HS384_A192KW, PBES2_HS512_A256KW, RSA1_5, RSA_OAEP, RSA_OAEP_256:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.KeyEncryptionAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a KeyEncryptionAlgorithm
|
||||
func (v KeyEncryptionAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+45
@@ -0,0 +1,45 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// KeyType represents the key type ("kty") that are supported
|
||||
type KeyType string
|
||||
|
||||
// Supported values for KeyType
|
||||
const (
|
||||
EC KeyType = "EC" // Elliptic Curve
|
||||
InvalidKeyType KeyType = "" // Invalid KeyType
|
||||
OctetSeq KeyType = "oct" // Octet sequence (used to represent symmetric keys)
|
||||
RSA KeyType = "RSA" // RSA
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *KeyType) Accept(value interface{}) error {
|
||||
var tmp KeyType
|
||||
switch x := value.(type) {
|
||||
case string:
|
||||
tmp = KeyType(x)
|
||||
case KeyType:
|
||||
tmp = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.KeyType: %T`, value)
|
||||
}
|
||||
switch tmp {
|
||||
case EC, OctetSeq, RSA:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.KeyType value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a KeyType
|
||||
func (v KeyType) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+54
@@ -0,0 +1,54 @@
|
||||
// this file was auto-generated by internal/cmd/gentypes/main.go: DO NOT EDIT
|
||||
|
||||
package jwa
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// SignatureAlgorithm represents the various signature algorithms as described in https://tools.ietf.org/html/rfc7518#section-3.1
|
||||
type SignatureAlgorithm string
|
||||
|
||||
// Supported values for SignatureAlgorithm
|
||||
const (
|
||||
ES256 SignatureAlgorithm = "ES256" // ECDSA using P-256 and SHA-256
|
||||
ES384 SignatureAlgorithm = "ES384" // ECDSA using P-384 and SHA-384
|
||||
ES512 SignatureAlgorithm = "ES512" // ECDSA using P-521 and SHA-512
|
||||
HS256 SignatureAlgorithm = "HS256" // HMAC using SHA-256
|
||||
HS384 SignatureAlgorithm = "HS384" // HMAC using SHA-384
|
||||
HS512 SignatureAlgorithm = "HS512" // HMAC using SHA-512
|
||||
NoSignature SignatureAlgorithm = "none"
|
||||
PS256 SignatureAlgorithm = "PS256" // RSASSA-PSS using SHA256 and MGF1-SHA256
|
||||
PS384 SignatureAlgorithm = "PS384" // RSASSA-PSS using SHA384 and MGF1-SHA384
|
||||
PS512 SignatureAlgorithm = "PS512" // RSASSA-PSS using SHA512 and MGF1-SHA512
|
||||
RS256 SignatureAlgorithm = "RS256" // RSASSA-PKCS-v1.5 using SHA-256
|
||||
RS384 SignatureAlgorithm = "RS384" // RSASSA-PKCS-v1.5 using SHA-384
|
||||
RS512 SignatureAlgorithm = "RS512" // RSASSA-PKCS-v1.5 using SHA-512
|
||||
)
|
||||
|
||||
// Accept is used when conversion from values given by
|
||||
// outside sources (such as JSON payloads) is required
|
||||
func (v *SignatureAlgorithm) Accept(value interface{}) error {
|
||||
var tmp SignatureAlgorithm
|
||||
switch x := value.(type) {
|
||||
case string:
|
||||
tmp = SignatureAlgorithm(x)
|
||||
case SignatureAlgorithm:
|
||||
tmp = x
|
||||
default:
|
||||
return errors.Errorf(`invalid type for jwa.SignatureAlgorithm: %T`, value)
|
||||
}
|
||||
switch tmp {
|
||||
case ES256, ES384, ES512, HS256, HS384, HS512, NoSignature, PS256, PS384, PS512, RS256, RS384, RS512:
|
||||
default:
|
||||
return errors.Errorf(`invalid jwa.SignatureAlgorithm value`)
|
||||
}
|
||||
|
||||
*v = tmp
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns the string representation of a SignatureAlgorithm
|
||||
func (v SignatureAlgorithm) String() string {
|
||||
return string(v)
|
||||
}
|
||||
+49
@@ -0,0 +1,49 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/base64"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func (c CertificateChain) Get() []*x509.Certificate {
|
||||
return c.certs
|
||||
}
|
||||
|
||||
func (c *CertificateChain) Accept(v interface{}) error {
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
return c.Accept([]string{x})
|
||||
case []interface{}:
|
||||
l := make([]string, len(x))
|
||||
for i, e := range x {
|
||||
if es, ok := e.(string); ok {
|
||||
l[i] = es
|
||||
} else {
|
||||
return errors.Errorf(`invalid list element type: expected string, got %T`, v)
|
||||
}
|
||||
}
|
||||
return c.Accept(l)
|
||||
case []string:
|
||||
certs := make([]*x509.Certificate, len(x))
|
||||
for i, e := range x {
|
||||
buf, err := base64.StdEncoding.DecodeString(e)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to base64 decode list element`)
|
||||
}
|
||||
cert, err := x509.ParseCertificate(buf)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to parse certificate`)
|
||||
}
|
||||
certs[i] = cert
|
||||
}
|
||||
|
||||
*c = CertificateChain{
|
||||
certs: certs,
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return errors.Errorf(`invalid value %T`, v)
|
||||
}
|
||||
}
|
||||
+307
@@ -0,0 +1,307 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math/big"
|
||||
|
||||
"github.com/lestrrat/go-jwx/internal/base64"
|
||||
"github.com/lestrrat/go-jwx/jwa"
|
||||
pdebug "github.com/lestrrat/go-pdebug"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func newECDSAPublicKey(key *ecdsa.PublicKey) (*ECDSAPublicKey, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`non-nil ecdsa.PublicKey required`)
|
||||
}
|
||||
|
||||
var hdr StandardHeaders
|
||||
hdr.Set(KeyTypeKey, jwa.EC)
|
||||
return &ECDSAPublicKey{
|
||||
headers: &hdr,
|
||||
key: key,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newECDSAPrivateKey(key *ecdsa.PrivateKey) (*ECDSAPrivateKey, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`non-nil ecdsa.PrivateKey required`)
|
||||
}
|
||||
|
||||
var hdr StandardHeaders
|
||||
hdr.Set(KeyTypeKey, jwa.EC)
|
||||
return &ECDSAPrivateKey{
|
||||
headers: &hdr,
|
||||
key: key,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (k ECDSAPrivateKey) PublicKey() (*ECDSAPublicKey, error) {
|
||||
return newECDSAPublicKey(&k.key.PublicKey)
|
||||
}
|
||||
|
||||
// Materialize returns the EC-DSA public key represented by this JWK
|
||||
func (k ECDSAPublicKey) Materialize() (interface{}, error) {
|
||||
return k.key, nil
|
||||
}
|
||||
|
||||
func (k ECDSAPublicKey) Curve() jwa.EllipticCurveAlgorithm {
|
||||
return jwa.EllipticCurveAlgorithm(k.key.Curve.Params().Name)
|
||||
}
|
||||
|
||||
func (k ECDSAPrivateKey) Curve() jwa.EllipticCurveAlgorithm {
|
||||
return jwa.EllipticCurveAlgorithm(k.key.PublicKey.Curve.Params().Name)
|
||||
}
|
||||
|
||||
func ecdsaThumbprint(hash crypto.Hash, crv, x, y string) []byte {
|
||||
h := hash.New()
|
||||
fmt.Fprintf(h, `{"crv":"`)
|
||||
fmt.Fprintf(h, crv)
|
||||
fmt.Fprintf(h, `","kty":"EC","x":"`)
|
||||
fmt.Fprintf(h, x)
|
||||
fmt.Fprintf(h, `","y":"`)
|
||||
fmt.Fprintf(h, y)
|
||||
fmt.Fprintf(h, `"}`)
|
||||
return h.Sum(nil)
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (k ECDSAPublicKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
return ecdsaThumbprint(
|
||||
hash,
|
||||
k.key.Curve.Params().Name,
|
||||
base64.EncodeToString(k.key.X.Bytes()),
|
||||
base64.EncodeToString(k.key.Y.Bytes()),
|
||||
), nil
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (k ECDSAPrivateKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
return ecdsaThumbprint(
|
||||
hash,
|
||||
k.key.Curve.Params().Name,
|
||||
base64.EncodeToString(k.key.X.Bytes()),
|
||||
base64.EncodeToString(k.key.Y.Bytes()),
|
||||
), nil
|
||||
}
|
||||
|
||||
// Materialize returns the EC-DSA private key represented by this JWK
|
||||
func (k ECDSAPrivateKey) Materialize() (interface{}, error) {
|
||||
return k.key, nil
|
||||
}
|
||||
|
||||
func (k ECDSAPublicKey) MarshalJSON() (buf []byte, err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPublicKey.MarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := make(map[string]interface{})
|
||||
if err := k.PopulateMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to populate pulibc key values`)
|
||||
}
|
||||
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
func (k ECDSAPublicKey) PopulateMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPublicKey.PopulateJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
if err := k.headers.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate header values`)
|
||||
}
|
||||
|
||||
const (
|
||||
xKey = `x`
|
||||
yKey = `y`
|
||||
crvKey = `crv`
|
||||
)
|
||||
m[xKey] = base64.EncodeToString(k.key.X.Bytes())
|
||||
m[yKey] = base64.EncodeToString(k.key.Y.Bytes())
|
||||
m[crvKey] = k.key.Curve.Params().Name
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k ECDSAPrivateKey) MarshalJSON() (buf []byte, err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPrivateKey.MarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := make(map[string]interface{})
|
||||
if err := k.PopulateMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to populate pulibc key values`)
|
||||
}
|
||||
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
func (k ECDSAPrivateKey) PopulateMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPrivateKey.PopulateJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
if err := k.headers.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate header values`)
|
||||
}
|
||||
|
||||
pubkey, err := newECDSAPublicKey(&k.key.PublicKey)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to construct public key from private key`)
|
||||
}
|
||||
|
||||
if err := pubkey.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate public key values`)
|
||||
}
|
||||
|
||||
m[`d`] = base64.EncodeToString(k.key.D.Bytes())
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *ECDSAPublicKey) UnmarshalJSON(data []byte) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPublicKey.UnmarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := map[string]interface{}{}
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal public key`)
|
||||
}
|
||||
|
||||
if err := k.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract data from map`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *ECDSAPublicKey) ExtractMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPublicKey.ExtractMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
const (
|
||||
xKey = `x`
|
||||
yKey = `y`
|
||||
crvKey = `crv`
|
||||
)
|
||||
|
||||
crvname, ok := m[crvKey]
|
||||
if !ok {
|
||||
return errors.Errorf(`failed to get required key crv`)
|
||||
}
|
||||
delete(m, crvKey)
|
||||
|
||||
var crv jwa.EllipticCurveAlgorithm
|
||||
if err := crv.Accept(crvname); err != nil {
|
||||
return errors.Wrap(err, `failed to accept value for crv key`)
|
||||
}
|
||||
|
||||
var curve elliptic.Curve
|
||||
switch crv {
|
||||
case jwa.P256:
|
||||
curve = elliptic.P256()
|
||||
case jwa.P384:
|
||||
curve = elliptic.P384()
|
||||
case jwa.P521:
|
||||
curve = elliptic.P521()
|
||||
default:
|
||||
return errors.Errorf(`invalid curve name %s`, crv)
|
||||
}
|
||||
|
||||
xbuf, err := getRequiredKey(m, xKey)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to get required key %s`, xKey)
|
||||
}
|
||||
delete(m, xKey)
|
||||
|
||||
ybuf, err := getRequiredKey(m, yKey)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to get required key %s`, yKey)
|
||||
}
|
||||
delete(m, yKey)
|
||||
|
||||
var x, y big.Int
|
||||
x.SetBytes(xbuf)
|
||||
y.SetBytes(ybuf)
|
||||
|
||||
var hdrs StandardHeaders
|
||||
if err := hdrs.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract header values`)
|
||||
}
|
||||
|
||||
*k = ECDSAPublicKey{
|
||||
headers: &hdrs,
|
||||
key: &ecdsa.PublicKey{
|
||||
Curve: curve,
|
||||
X: &x,
|
||||
Y: &y,
|
||||
},
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *ECDSAPrivateKey) UnmarshalJSON(data []byte) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPrivateKey.UnmarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := map[string]interface{}{}
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal public key`)
|
||||
}
|
||||
|
||||
if err := k.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract data from map`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *ECDSAPrivateKey) ExtractMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.ECDSAPrivateKey.ExtractMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
const (
|
||||
dKey = `d`
|
||||
)
|
||||
|
||||
dbuf, err := getRequiredKey(m, dKey)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to get required key %s`, dKey)
|
||||
}
|
||||
delete(m, dKey)
|
||||
|
||||
var pubkey ECDSAPublicKey
|
||||
if err := pubkey.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract public key values`)
|
||||
}
|
||||
|
||||
var d big.Int
|
||||
d.SetBytes(dbuf)
|
||||
|
||||
*k = ECDSAPrivateKey{
|
||||
headers: pubkey.headers,
|
||||
key: &ecdsa.PrivateKey{
|
||||
PublicKey: *(pubkey.key),
|
||||
D: &d,
|
||||
},
|
||||
}
|
||||
pubkey.headers = nil
|
||||
return nil
|
||||
}
|
||||
+382
@@ -0,0 +1,382 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/lestrrat/go-jwx/jwa"
|
||||
"github.com/lestrrat/go-pdebug"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
AlgorithmKey = "alg"
|
||||
KeyIDKey = "kid"
|
||||
KeyTypeKey = "kty"
|
||||
KeyUsageKey = "use"
|
||||
KeyOpsKey = "key_ops"
|
||||
X509CertChainKey = "x5c"
|
||||
X509CertThumbprintKey = "x5t"
|
||||
X509CertThumbprintS256Key = "x5t#S256"
|
||||
X509URLKey = "x5u"
|
||||
)
|
||||
|
||||
type Headers interface {
|
||||
Remove(string)
|
||||
Get(string) (interface{}, bool)
|
||||
Set(string, interface{}) error
|
||||
PopulateMap(map[string]interface{}) error
|
||||
ExtractMap(map[string]interface{}) error
|
||||
Walk(func(string, interface{}) error) error
|
||||
Algorithm() string
|
||||
KeyID() string
|
||||
KeyType() jwa.KeyType
|
||||
KeyUsage() string
|
||||
KeyOps() []KeyOperation
|
||||
X509CertChain() []*x509.Certificate
|
||||
X509CertThumbprint() string
|
||||
X509CertThumbprintS256() string
|
||||
X509URL() string
|
||||
}
|
||||
|
||||
type StandardHeaders struct {
|
||||
algorithm *string // https://tools.ietf.org/html/rfc7517#section-4.4
|
||||
keyID *string // https://tools.ietf.org/html/rfc7515#section-4.1.4
|
||||
keyType *jwa.KeyType // https://tools.ietf.org/html/rfc7517#section-4.1
|
||||
keyUsage *string // https://tools.ietf.org/html/rfc7517#section-4.2
|
||||
keyops []KeyOperation // https://tools.ietf.org/html/rfc7517#section-4.3
|
||||
x509CertChain *CertificateChain // https://tools.ietf.org/html/rfc7515#section-4.1.6
|
||||
x509CertThumbprint *string // https://tools.ietf.org/html/rfc7515#section-4.1.7
|
||||
x509CertThumbprintS256 *string // https://tools.ietf.org/html/rfc7515#section-4.1.8
|
||||
x509URL *string // https://tools.ietf.org/html/rfc7515#section-4.1.5
|
||||
privateParams map[string]interface{}
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) Remove(s string) {
|
||||
delete(h.privateParams, s)
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) Algorithm() string {
|
||||
if v := h.algorithm; v != nil {
|
||||
return *v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) KeyID() string {
|
||||
if v := h.keyID; v != nil {
|
||||
return *v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) KeyType() jwa.KeyType {
|
||||
if v := h.keyType; v != nil {
|
||||
return *v
|
||||
}
|
||||
return jwa.InvalidKeyType
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) KeyUsage() string {
|
||||
if v := h.keyUsage; v != nil {
|
||||
return *v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) KeyOps() []KeyOperation {
|
||||
return h.keyops
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) X509CertChain() []*x509.Certificate {
|
||||
return h.x509CertChain.Get()
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) X509CertThumbprint() string {
|
||||
if v := h.x509CertThumbprint; v != nil {
|
||||
return *v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) X509CertThumbprintS256() string {
|
||||
if v := h.x509CertThumbprintS256; v != nil {
|
||||
return *v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) X509URL() string {
|
||||
if v := h.x509URL; v != nil {
|
||||
return *v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) Get(name string) (interface{}, bool) {
|
||||
switch name {
|
||||
case AlgorithmKey:
|
||||
v := h.algorithm
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
case KeyIDKey:
|
||||
v := h.keyID
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
case KeyTypeKey:
|
||||
v := h.keyType
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
case KeyUsageKey:
|
||||
v := h.keyUsage
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
case KeyOpsKey:
|
||||
v := h.keyops
|
||||
if len(v) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
return v, true
|
||||
case X509CertChainKey:
|
||||
v := h.x509CertChain
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return v.Get(), true
|
||||
case X509CertThumbprintKey:
|
||||
v := h.x509CertThumbprint
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
case X509CertThumbprintS256Key:
|
||||
v := h.x509CertThumbprintS256
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
case X509URLKey:
|
||||
v := h.x509URL
|
||||
if v == nil {
|
||||
return nil, false
|
||||
}
|
||||
return *v, true
|
||||
default:
|
||||
v, ok := h.privateParams[name]
|
||||
return v, ok
|
||||
}
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) Set(name string, value interface{}) error {
|
||||
switch name {
|
||||
case AlgorithmKey:
|
||||
switch v := value.(type) {
|
||||
case string:
|
||||
h.algorithm = &v
|
||||
return nil
|
||||
case fmt.Stringer:
|
||||
s := v.String()
|
||||
h.algorithm = &s
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, AlgorithmKey, value)
|
||||
case KeyIDKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyID = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyIDKey, value)
|
||||
case KeyTypeKey:
|
||||
var acceptor jwa.KeyType
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, KeyTypeKey)
|
||||
}
|
||||
h.keyType = &acceptor
|
||||
return nil
|
||||
case KeyUsageKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.keyUsage = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyUsageKey, value)
|
||||
case KeyOpsKey:
|
||||
if v, ok := value.([]KeyOperation); ok {
|
||||
h.keyops = v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, KeyOpsKey, value)
|
||||
case X509CertChainKey:
|
||||
var acceptor CertificateChain
|
||||
if err := acceptor.Accept(value); err != nil {
|
||||
return errors.Wrapf(err, `invalid value for %s key`, X509CertChainKey)
|
||||
}
|
||||
h.x509CertChain = &acceptor
|
||||
return nil
|
||||
case X509CertThumbprintKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprint = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintKey, value)
|
||||
case X509CertThumbprintS256Key:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509CertThumbprintS256 = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509CertThumbprintS256Key, value)
|
||||
case X509URLKey:
|
||||
if v, ok := value.(string); ok {
|
||||
h.x509URL = &v
|
||||
return nil
|
||||
}
|
||||
return errors.Errorf(`invalid value for %s key: %T`, X509URLKey, value)
|
||||
default:
|
||||
if h.privateParams == nil {
|
||||
h.privateParams = map[string]interface{}{}
|
||||
}
|
||||
h.privateParams[name] = value
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h StandardHeaders) MarshalJSON() ([]byte, error) {
|
||||
m := map[string]interface{}{}
|
||||
if err := h.PopulateMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to populate map for serialization`)
|
||||
}
|
||||
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
// PopulateMap populates a map with appropriate values that represent
|
||||
// the headers as a JSON object. This exists primarily because JWKs are
|
||||
// represented as flat objects instead of differentiating the different
|
||||
// parts of the message in separate sub objects.
|
||||
func (h StandardHeaders) PopulateMap(m map[string]interface{}) error {
|
||||
for k, v := range h.privateParams {
|
||||
m[k] = v
|
||||
}
|
||||
if v, ok := h.Get(AlgorithmKey); ok {
|
||||
m[AlgorithmKey] = v
|
||||
}
|
||||
if v, ok := h.Get(KeyIDKey); ok {
|
||||
m[KeyIDKey] = v
|
||||
}
|
||||
if v, ok := h.Get(KeyTypeKey); ok {
|
||||
m[KeyTypeKey] = v
|
||||
}
|
||||
if v, ok := h.Get(KeyUsageKey); ok {
|
||||
m[KeyUsageKey] = v
|
||||
}
|
||||
if v, ok := h.Get(KeyOpsKey); ok {
|
||||
m[KeyOpsKey] = v
|
||||
}
|
||||
if v, ok := h.Get(X509CertChainKey); ok {
|
||||
m[X509CertChainKey] = v
|
||||
}
|
||||
if v, ok := h.Get(X509CertThumbprintKey); ok {
|
||||
m[X509CertThumbprintKey] = v
|
||||
}
|
||||
if v, ok := h.Get(X509CertThumbprintS256Key); ok {
|
||||
m[X509CertThumbprintS256Key] = v
|
||||
}
|
||||
if v, ok := h.Get(X509URLKey); ok {
|
||||
m[X509URLKey] = v
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExtractMap populates the appropriate values from a map that represent
|
||||
// the headers as a JSON object. This exists primarily because JWKs are
|
||||
// represented as flat objects instead of differentiating the different
|
||||
// parts of the message in separate sub objects.
|
||||
func (h *StandardHeaders) ExtractMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker(`jwk.StandardHeaders.ExtractMap`).BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
if v, ok := m[AlgorithmKey]; ok {
|
||||
if err := h.Set(AlgorithmKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, AlgorithmKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[KeyIDKey]; ok {
|
||||
if err := h.Set(KeyIDKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, KeyIDKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[KeyTypeKey]; ok {
|
||||
if err := h.Set(KeyTypeKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, KeyTypeKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[KeyUsageKey]; ok {
|
||||
if err := h.Set(KeyUsageKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, KeyUsageKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[KeyOpsKey]; ok {
|
||||
if err := h.Set(KeyOpsKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, KeyOpsKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[X509CertChainKey]; ok {
|
||||
if err := h.Set(X509CertChainKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, X509CertChainKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[X509CertThumbprintKey]; ok {
|
||||
if err := h.Set(X509CertThumbprintKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, X509CertThumbprintKey)
|
||||
}
|
||||
}
|
||||
if v, ok := m[X509CertThumbprintS256Key]; ok {
|
||||
if err := h.Set(X509CertThumbprintS256Key, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, X509CertThumbprintS256Key)
|
||||
}
|
||||
}
|
||||
if v, ok := m[X509URLKey]; ok {
|
||||
if err := h.Set(X509URLKey, v); err != nil {
|
||||
return errors.Wrapf(err, `failed to set value for key %s`, X509URLKey)
|
||||
}
|
||||
}
|
||||
h.privateParams = m
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *StandardHeaders) UnmarshalJSON(buf []byte) error {
|
||||
var m map[string]interface{}
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal headers`)
|
||||
}
|
||||
|
||||
return h.ExtractMap(m)
|
||||
}
|
||||
|
||||
func (h StandardHeaders) Walk(f func(string, interface{}) error) error {
|
||||
for _, key := range []string{AlgorithmKey, KeyIDKey, KeyTypeKey, KeyUsageKey, KeyOpsKey, X509CertChainKey, X509CertThumbprintKey, X509CertThumbprintS256Key, X509URLKey} {
|
||||
if v, ok := h.Get(key); ok {
|
||||
if err := f(key, v); err != nil {
|
||||
return errors.Wrapf(err, `walk function returned error for %s`, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for k, v := range h.privateParams {
|
||||
if err := f(k, v); err != nil {
|
||||
return errors.Wrapf(err, `walk function returned error for %s`, k)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+104
@@ -0,0 +1,104 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
)
|
||||
|
||||
// KeyUsageType is used to denote what this key should be used for
|
||||
type KeyUsageType string
|
||||
|
||||
const (
|
||||
// ForSignature is the value used in the headers to indicate that
|
||||
// this key should be used for signatures
|
||||
ForSignature KeyUsageType = "sig"
|
||||
// ForEncryption is the value used in the headers to indicate that
|
||||
// this key should be used for encryptiong
|
||||
ForEncryption KeyUsageType = "enc"
|
||||
)
|
||||
|
||||
type CertificateChain struct {
|
||||
certs []*x509.Certificate
|
||||
}
|
||||
|
||||
// Errors related to JWK
|
||||
var (
|
||||
ErrInvalidHeaderName = errors.New("invalid header name")
|
||||
ErrInvalidHeaderValue = errors.New("invalid value for header key")
|
||||
ErrUnsupportedKty = errors.New("unsupported kty")
|
||||
ErrUnsupportedCurve = errors.New("unsupported curve")
|
||||
)
|
||||
|
||||
type KeyOperation string
|
||||
|
||||
const (
|
||||
KeyOpSign KeyOperation = "sign" // (compute digital signature or MAC)
|
||||
KeyOpVerify = "verify" // (verify digital signature or MAC)
|
||||
KeyOpEncrypt = "encrypt" // (encrypt content)
|
||||
KeyOpDecrypt = "decrypt" // (decrypt content and validate decryption, if applicable)
|
||||
KeyOpWrapKey = "wrapKey" // (encrypt key)
|
||||
KeyOpUnwrapKey = "unwrapKey" // (decrypt key and validate decryption, if applicable)
|
||||
KeyOpDeriveKey = "deriveKey" // (derive key)
|
||||
KeyOpDeriveBits = "deriveBits" // (derive bits not to be used as a key)
|
||||
)
|
||||
|
||||
// Set is a convenience struct to allow generating and parsing
|
||||
// JWK sets as opposed to single JWKs
|
||||
type Set struct {
|
||||
Keys []Key `json:"keys"`
|
||||
}
|
||||
|
||||
// Key defines the minimal interface for each of the
|
||||
// key types. Their use and implementation differ significantly
|
||||
// between each key types, so you should use type assertions
|
||||
// to perform more specific tasks with each key
|
||||
type Key interface {
|
||||
Headers
|
||||
|
||||
// Materialize creates the corresponding key. For example,
|
||||
// RSA types would create *rsa.PublicKey or *rsa.PrivateKey,
|
||||
// EC types would create *ecdsa.PublicKey or *ecdsa.PrivateKey,
|
||||
// and OctetSeq types create a []byte key.
|
||||
Materialize() (interface{}, error)
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
Thumbprint(crypto.Hash) ([]byte, error)
|
||||
}
|
||||
|
||||
type headers interface {
|
||||
Headers
|
||||
}
|
||||
|
||||
// RSAPublicKey is a type of JWK generated from RSA public keys
|
||||
type RSAPublicKey struct {
|
||||
headers
|
||||
key *rsa.PublicKey
|
||||
}
|
||||
|
||||
// RSAPrivateKey is a type of JWK generated from RSA private keys
|
||||
type RSAPrivateKey struct {
|
||||
headers
|
||||
key *rsa.PrivateKey
|
||||
}
|
||||
|
||||
// SymmetricKey is a type of JWK generated from symmetric keys
|
||||
type SymmetricKey struct {
|
||||
headers
|
||||
key []byte
|
||||
}
|
||||
|
||||
// ECDSAPublicKey is a type of JWK generated from ECDSA public keys
|
||||
type ECDSAPublicKey struct {
|
||||
headers
|
||||
key *ecdsa.PublicKey
|
||||
}
|
||||
|
||||
// ECDSAPrivateKey is a type of JWK generated from ECDH-ES private keys
|
||||
type ECDSAPrivateKey struct {
|
||||
headers
|
||||
key *ecdsa.PrivateKey
|
||||
}
|
||||
+240
@@ -0,0 +1,240 @@
|
||||
//go:generate go run internal/cmd/genheader/main.go
|
||||
|
||||
// Package jwk implements JWK as described in https://tools.ietf.org/html/rfc7517
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto/ecdsa"
|
||||
"crypto/rsa"
|
||||
"encoding/json"
|
||||
"io/ioutil"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
|
||||
"github.com/lestrrat/go-jwx/internal/base64"
|
||||
"github.com/lestrrat/go-jwx/jwa"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// New creates a jwk.Key from the given key.
|
||||
func New(key interface{}) (Key, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`jwk.New requires a non-nil key`)
|
||||
}
|
||||
|
||||
switch v := key.(type) {
|
||||
case *rsa.PrivateKey:
|
||||
return newRSAPrivateKey(v)
|
||||
case *rsa.PublicKey:
|
||||
return newRSAPublicKey(v)
|
||||
case *ecdsa.PrivateKey:
|
||||
return newECDSAPrivateKey(v)
|
||||
case *ecdsa.PublicKey:
|
||||
return newECDSAPublicKey(v)
|
||||
case []byte:
|
||||
return newSymmetricKey(v)
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid key type %T`, key)
|
||||
}
|
||||
}
|
||||
|
||||
// Fetch fetches a JWK resource specified by a URL
|
||||
func Fetch(urlstring string) (*Set, error) {
|
||||
u, err := url.Parse(urlstring)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to parse url`)
|
||||
}
|
||||
|
||||
var src []byte
|
||||
switch u.Scheme {
|
||||
case "http", "https":
|
||||
res, err := http.Get(u.String())
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to fetch remote JWK")
|
||||
}
|
||||
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return nil, errors.New("failed to fetch remote JWK (status != 200)")
|
||||
}
|
||||
|
||||
// XXX Check for maximum length to read?
|
||||
buf, err := ioutil.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to read JWK HTTP response body")
|
||||
}
|
||||
defer res.Body.Close()
|
||||
src = buf
|
||||
case "file":
|
||||
f, err := os.Open(u.Path)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to open jwk file`)
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
buf, err := ioutil.ReadAll(f)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed read content from jwk file`)
|
||||
}
|
||||
src = buf
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid url scheme %s`, u.Scheme)
|
||||
}
|
||||
|
||||
return Parse(src)
|
||||
}
|
||||
|
||||
// FetchHTTP fetches the remote JWK and parses its contents
|
||||
func FetchHTTP(jwkurl string) (*Set, error) {
|
||||
res, err := http.Get(jwkurl)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to fetch remote JWK")
|
||||
}
|
||||
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return nil, errors.New("failed to fetch remote JWK (status != 200)")
|
||||
}
|
||||
|
||||
// XXX Check for maximum length to read?
|
||||
buf, err := ioutil.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to read JWK HTTP response body")
|
||||
}
|
||||
defer res.Body.Close()
|
||||
|
||||
return Parse(buf)
|
||||
}
|
||||
|
||||
// Parse parses JWK from the incoming byte buffer.
|
||||
func Parse(buf []byte) (*Set, error) {
|
||||
m := make(map[string]interface{})
|
||||
if err := json.Unmarshal(buf, &m); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to unmarshal JWK")
|
||||
}
|
||||
|
||||
// We must change what the underlying structure that gets decoded
|
||||
// out of this JSON is based on parameters within the already parsed
|
||||
// JSON (m). In order to do this, we have to go through the tedious
|
||||
// task of parsing the contents of this map :/
|
||||
if _, ok := m["keys"]; ok {
|
||||
var set Set
|
||||
if err := set.ExtractMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to extract from map`)
|
||||
}
|
||||
return &set, nil
|
||||
}
|
||||
|
||||
k, err := constructKey(m)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, `failed to construct key from keys`)
|
||||
}
|
||||
return &Set{Keys: []Key{k}}, nil
|
||||
}
|
||||
|
||||
// ParseString parses JWK from the incoming string.
|
||||
func ParseString(s string) (*Set, error) {
|
||||
return Parse([]byte(s))
|
||||
}
|
||||
|
||||
// LookupKeyID looks for keys matching the given key id. Note that the
|
||||
// Set *may* contain multiple keys with the same key id
|
||||
func (s Set) LookupKeyID(kid string) []Key {
|
||||
var keys []Key
|
||||
for _, key := range s.Keys {
|
||||
if key.KeyID() == kid {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func (s *Set) ExtractMap(m map[string]interface{}) error {
|
||||
raw, ok := m["keys"]
|
||||
if !ok {
|
||||
return errors.New("missing 'keys' parameter")
|
||||
}
|
||||
|
||||
v, ok := raw.([]interface{})
|
||||
if !ok {
|
||||
return errors.New("invalid 'keys' parameter")
|
||||
}
|
||||
|
||||
var ks Set
|
||||
for _, c := range v {
|
||||
conf, ok := c.(map[string]interface{})
|
||||
if !ok {
|
||||
return errors.New("invalid element in 'keys'")
|
||||
}
|
||||
|
||||
k, err := constructKey(conf)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to construct key from map`)
|
||||
}
|
||||
ks.Keys = append(ks.Keys, k)
|
||||
}
|
||||
|
||||
*s = ks
|
||||
return nil
|
||||
}
|
||||
|
||||
func constructKey(m map[string]interface{}) (Key, error) {
|
||||
kty, ok := m["kty"].(string)
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`unsupported kty type %T`, m[KeyTypeKey])
|
||||
}
|
||||
|
||||
var key Key
|
||||
switch jwa.KeyType(kty) {
|
||||
case jwa.RSA:
|
||||
if _, ok := m["d"]; ok {
|
||||
key = &RSAPrivateKey{}
|
||||
} else {
|
||||
key = &RSAPublicKey{}
|
||||
}
|
||||
case jwa.EC:
|
||||
if _, ok := m["d"]; ok {
|
||||
key = &ECDSAPrivateKey{}
|
||||
} else {
|
||||
key = &ECDSAPublicKey{}
|
||||
}
|
||||
case jwa.OctetSeq:
|
||||
key = &SymmetricKey{}
|
||||
default:
|
||||
return nil, errors.Errorf(`invalid kty %s`, kty)
|
||||
}
|
||||
|
||||
if err := key.ExtractMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to extract key from map`)
|
||||
}
|
||||
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func getRequiredKey(m map[string]interface{}, key string) ([]byte, error) {
|
||||
return getKey(m, key, true)
|
||||
}
|
||||
|
||||
func getOptionalKey(m map[string]interface{}, key string) ([]byte, error) {
|
||||
return getKey(m, key, false)
|
||||
}
|
||||
|
||||
func getKey(m map[string]interface{}, key string, required bool) ([]byte, error) {
|
||||
v, ok := m[key]
|
||||
if !ok {
|
||||
if !required {
|
||||
return nil, errors.Errorf(`missing parameter '%s'`, key)
|
||||
}
|
||||
return nil, errors.Errorf(`missing required parameter '%s'`, key)
|
||||
}
|
||||
|
||||
vs, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, errors.Errorf(`invalid type for parameter '%s': %T`, key, v)
|
||||
}
|
||||
|
||||
buf, err := base64.DecodeString(vs)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, `failed to base64 decode key %s`, key)
|
||||
}
|
||||
return buf, nil
|
||||
}
|
||||
+346
@@ -0,0 +1,346 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto"
|
||||
"crypto/rsa"
|
||||
"encoding/json"
|
||||
"math/big"
|
||||
|
||||
"github.com/lestrrat/go-jwx/internal/base64"
|
||||
"github.com/lestrrat/go-jwx/jwa"
|
||||
pdebug "github.com/lestrrat/go-pdebug"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func newRSAPublicKey(key *rsa.PublicKey) (*RSAPublicKey, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`non-nil rsa.PublicKey required`)
|
||||
}
|
||||
|
||||
var hdr StandardHeaders
|
||||
hdr.Set(KeyTypeKey, jwa.RSA)
|
||||
return &RSAPublicKey{
|
||||
headers: &hdr,
|
||||
key: key,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newRSAPrivateKey(key *rsa.PrivateKey) (*RSAPrivateKey, error) {
|
||||
if key == nil {
|
||||
return nil, errors.New(`non-nil rsa.PrivateKey required`)
|
||||
}
|
||||
|
||||
if len(key.Primes) < 2 {
|
||||
return nil, errors.New("two primes required for RSA private key")
|
||||
}
|
||||
|
||||
var hdr StandardHeaders
|
||||
hdr.Set(KeyTypeKey, jwa.RSA)
|
||||
return &RSAPrivateKey{
|
||||
headers: &hdr,
|
||||
key: key,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (k RSAPrivateKey) PublicKey() (*RSAPublicKey, error) {
|
||||
return newRSAPublicKey(&k.key.PublicKey)
|
||||
}
|
||||
|
||||
func (k *RSAPublicKey) Materialize() (interface{}, error) {
|
||||
if k.key == nil {
|
||||
return nil, errors.New(`key has no rsa.PublicKey associated with it`)
|
||||
}
|
||||
return k.key, nil
|
||||
}
|
||||
|
||||
func (k *RSAPrivateKey) Materialize() (interface{}, error) {
|
||||
if k.key == nil {
|
||||
return nil, errors.New(`key has no rsa.PrivateKey associated with it`)
|
||||
}
|
||||
return k.key, nil
|
||||
}
|
||||
|
||||
func (k RSAPublicKey) MarshalJSON() (buf []byte, err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPublicKey.MarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := map[string]interface{}{}
|
||||
if err := k.PopulateMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to populate pulibc key values`)
|
||||
}
|
||||
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
func (k RSAPublicKey) PopulateMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPublicKey.PopulateJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
if err := k.headers.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate header values`)
|
||||
}
|
||||
|
||||
m[`n`] = base64.EncodeToString(k.key.N.Bytes())
|
||||
m[`e`] = base64.EncodeUint64ToString(uint64(k.key.E))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *RSAPublicKey) UnmarshalJSON(data []byte) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPublicKey.UnmarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := map[string]interface{}{}
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal public key`)
|
||||
}
|
||||
|
||||
if err := k.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract data from map`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *RSAPublicKey) ExtractMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPublicKey.ExtractMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
const (
|
||||
eKey = `e`
|
||||
nKey = `n`
|
||||
)
|
||||
|
||||
nbuf, err := getRequiredKey(m, nKey)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to get required key %s`, nKey)
|
||||
}
|
||||
delete(m, nKey)
|
||||
|
||||
ebuf, err := getRequiredKey(m, eKey)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to get required key %s`, eKey)
|
||||
}
|
||||
delete(m, eKey)
|
||||
|
||||
var n, e big.Int
|
||||
n.SetBytes(nbuf)
|
||||
e.SetBytes(ebuf)
|
||||
|
||||
var hdrs StandardHeaders
|
||||
if err := hdrs.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract header values`)
|
||||
}
|
||||
|
||||
*k = RSAPublicKey{
|
||||
headers: &hdrs,
|
||||
key: &rsa.PublicKey{E: int(e.Int64()), N: &n},
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k RSAPrivateKey) MarshalJSON() (buf []byte, err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPrivateKey.MarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := make(map[string]interface{})
|
||||
if err := k.PopulateMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to populate private key values`)
|
||||
}
|
||||
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
func (k RSAPrivateKey) PopulateMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPrivateKey.PopulateMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
const (
|
||||
dKey = `d`
|
||||
pKey = `p`
|
||||
qKey = `q`
|
||||
dpKey = `dp`
|
||||
dqKey = `dq`
|
||||
qiKey = `qi`
|
||||
)
|
||||
|
||||
if err := k.headers.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate header values`)
|
||||
}
|
||||
|
||||
pubkey, _ := newRSAPublicKey(&k.key.PublicKey)
|
||||
if err := pubkey.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate public key values`)
|
||||
}
|
||||
|
||||
if err := k.headers.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate header values`)
|
||||
}
|
||||
m[dKey] = base64.EncodeToString(k.key.D.Bytes())
|
||||
m[pKey] = base64.EncodeToString(k.key.Primes[0].Bytes())
|
||||
m[qKey] = base64.EncodeToString(k.key.Primes[1].Bytes())
|
||||
if v := k.key.Precomputed.Dp; v != nil {
|
||||
m[dpKey] = base64.EncodeToString(v.Bytes())
|
||||
}
|
||||
if v := k.key.Precomputed.Dq; v != nil {
|
||||
m[dqKey] = base64.EncodeToString(v.Bytes())
|
||||
}
|
||||
if v := k.key.Precomputed.Qinv; v != nil {
|
||||
m[qiKey] = base64.EncodeToString(v.Bytes())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *RSAPrivateKey) UnmarshalJSON(data []byte) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPrivateKey.UnmarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
|
||||
pdebug.Printf("data --> %s", data)
|
||||
}
|
||||
|
||||
m := map[string]interface{}{}
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal public key`)
|
||||
}
|
||||
|
||||
var key RSAPrivateKey
|
||||
if err := key.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract data from map`)
|
||||
}
|
||||
*k = key
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (k *RSAPrivateKey) ExtractMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.RSAPrivateKey.ExractMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
const (
|
||||
dKey = `d`
|
||||
pKey = `p`
|
||||
qKey = `q`
|
||||
dpKey = `dp`
|
||||
dqKey = `dq`
|
||||
qiKey = `qi`
|
||||
)
|
||||
|
||||
dbuf, err := getRequiredKey(m, dKey)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to get required key`)
|
||||
}
|
||||
delete(m, dKey)
|
||||
|
||||
pbuf, err := getRequiredKey(m, pKey)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to get required key`)
|
||||
}
|
||||
delete(m, pKey)
|
||||
|
||||
qbuf, err := getRequiredKey(m, qKey)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to get required key`)
|
||||
}
|
||||
delete(m, qKey)
|
||||
|
||||
var d, q, p big.Int
|
||||
d.SetBytes(dbuf)
|
||||
q.SetBytes(qbuf)
|
||||
p.SetBytes(pbuf)
|
||||
|
||||
var dp, dq, qi *big.Int
|
||||
|
||||
dpbuf, err := getOptionalKey(m, dpKey)
|
||||
if err == nil {
|
||||
delete(m, dpKey)
|
||||
|
||||
dp = &big.Int{}
|
||||
dp.SetBytes(dpbuf)
|
||||
}
|
||||
|
||||
dqbuf, err := getOptionalKey(m, dqKey)
|
||||
if err == nil {
|
||||
delete(m, dqKey)
|
||||
|
||||
dq = &big.Int{}
|
||||
dq.SetBytes(dqbuf)
|
||||
}
|
||||
|
||||
qibuf, err := getOptionalKey(m, qiKey)
|
||||
if err == nil {
|
||||
delete(m, qiKey)
|
||||
|
||||
qi = &big.Int{}
|
||||
qi.SetBytes(qibuf)
|
||||
}
|
||||
|
||||
var pubkey RSAPublicKey
|
||||
if err := pubkey.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract fields for public key`)
|
||||
}
|
||||
|
||||
materialized, err := pubkey.Materialize()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, `failed to materialize RSA public key`)
|
||||
}
|
||||
rsaPubkey := materialized.(*rsa.PublicKey)
|
||||
|
||||
var key rsa.PrivateKey
|
||||
key.PublicKey = *rsaPubkey
|
||||
key.D = &d
|
||||
key.Primes = []*big.Int{&p, &q}
|
||||
|
||||
if dp != nil {
|
||||
key.Precomputed.Dp = dp
|
||||
}
|
||||
if dq != nil {
|
||||
key.Precomputed.Dq = dq
|
||||
}
|
||||
if qi != nil {
|
||||
key.Precomputed.Qinv = qi
|
||||
}
|
||||
|
||||
*k = RSAPrivateKey{
|
||||
headers: pubkey.headers,
|
||||
key: &key,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (k RSAPrivateKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
return rsaThumbprint(hash, &k.key.PublicKey)
|
||||
}
|
||||
|
||||
func (k RSAPublicKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
return rsaThumbprint(hash, k.key)
|
||||
}
|
||||
|
||||
func rsaThumbprint(hash crypto.Hash, key *rsa.PublicKey) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
buf.WriteString(`{"e":"`)
|
||||
buf.WriteString(base64.EncodeUint64ToString(uint64(key.E)))
|
||||
buf.WriteString(`","kty":"RSA","n":"`)
|
||||
buf.WriteString(base64.EncodeToString(key.N.Bytes()))
|
||||
buf.WriteString(`"}`)
|
||||
|
||||
h := hash.New()
|
||||
buf.WriteTo(h)
|
||||
return h.Sum(nil), nil
|
||||
}
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
package jwk
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/lestrrat/go-jwx/internal/base64"
|
||||
"github.com/lestrrat/go-jwx/jwa"
|
||||
pdebug "github.com/lestrrat/go-pdebug"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func newSymmetricKey(key []byte) (*SymmetricKey, error) {
|
||||
if len(key) == 0 {
|
||||
return nil, errors.New(`non-empty []byte key required`)
|
||||
}
|
||||
|
||||
var hdr StandardHeaders
|
||||
hdr.Set(KeyTypeKey, jwa.OctetSeq)
|
||||
return &SymmetricKey{
|
||||
headers: &hdr,
|
||||
key: key,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Materialize returns the octets for this symmetric key.
|
||||
// Since this is a symmetric key, this just calls Octets
|
||||
func (s SymmetricKey) Materialize() (interface{}, error) {
|
||||
return s.Octets(), nil
|
||||
}
|
||||
|
||||
// Octets returns the octets in the key
|
||||
func (s SymmetricKey) Octets() []byte {
|
||||
return s.key
|
||||
}
|
||||
|
||||
// Thumbprint returns the JWK thumbprint using the indicated
|
||||
// hashing algorithm, according to RFC 7638
|
||||
func (s SymmetricKey) Thumbprint(hash crypto.Hash) ([]byte, error) {
|
||||
h := hash.New()
|
||||
fmt.Fprintf(h, `{"k":"`)
|
||||
fmt.Fprintf(h, base64.EncodeToString(s.key))
|
||||
fmt.Fprintf(h, `","kty":"oct"}`)
|
||||
return h.Sum(nil), nil
|
||||
}
|
||||
|
||||
func (k *SymmetricKey) UnmarshalJSON(data []byte) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.SymmetricKey.UnmarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := map[string]interface{}{}
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return errors.Wrap(err, `failed to unmarshal public key`)
|
||||
}
|
||||
|
||||
if err := k.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract data from map`)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SymmetricKey) ExtractMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.SymmetricKey.ExtractMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
const kKey = `k`
|
||||
|
||||
kbuf, err := getRequiredKey(m, kKey)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, `failed to get required key '%s'`, kKey)
|
||||
}
|
||||
delete(m, kKey)
|
||||
|
||||
var hdrs StandardHeaders
|
||||
if err := hdrs.ExtractMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to extract header values`)
|
||||
}
|
||||
|
||||
*s = SymmetricKey{
|
||||
headers: &hdrs,
|
||||
key: kbuf,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SymmetricKey) MarshalJSON() (buf []byte, err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.SymmetricKey.MarshalJSON").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
m := make(map[string]interface{})
|
||||
if err := s.PopulateMap(m); err != nil {
|
||||
return nil, errors.Wrap(err, `failed to populate symmetric key values`)
|
||||
}
|
||||
|
||||
return json.Marshal(m)
|
||||
}
|
||||
|
||||
func (s SymmetricKey) PopulateMap(m map[string]interface{}) (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("jwk.SymmetricKey.PopulateMap").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
if err := s.headers.PopulateMap(m); err != nil {
|
||||
return errors.Wrap(err, `failed to populate header values`)
|
||||
}
|
||||
|
||||
const kKey = `k`
|
||||
m[kKey] = base64.EncodeToString(s.key)
|
||||
return nil
|
||||
}
|
||||
+24
@@ -0,0 +1,24 @@
|
||||
# Compiled Object files, Static and Dynamic libs (Shared Objects)
|
||||
*.o
|
||||
*.a
|
||||
*.so
|
||||
|
||||
# Folders
|
||||
_obj
|
||||
_test
|
||||
|
||||
# Architecture specific extensions/prefixes
|
||||
*.[568vq]
|
||||
[568vq].out
|
||||
|
||||
*.cgo1.go
|
||||
*.cgo2.c
|
||||
_cgo_defun.c
|
||||
_cgo_gotypes.go
|
||||
_cgo_export.*
|
||||
|
||||
_testmain.go
|
||||
|
||||
*.exe
|
||||
*.test
|
||||
*.prof
|
||||
+14
@@ -0,0 +1,14 @@
|
||||
language: go
|
||||
sudo: false
|
||||
go:
|
||||
- 1.6
|
||||
- 1.7
|
||||
- tip
|
||||
install:
|
||||
- go get -t -v ./...
|
||||
- go get -t -tags debug0 -v ./...
|
||||
script:
|
||||
- go test -v ./...
|
||||
- go test -tags debug ./...
|
||||
- PDEBUG_TRACE=1 go test -tags debug ./...
|
||||
- go test -tags debug0 ./...
|
||||
+21
@@ -0,0 +1,21 @@
|
||||
The MIT License (MIT)
|
||||
|
||||
Copyright (c) 2016 lestrrat
|
||||
|
||||
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.
|
||||
+99
@@ -0,0 +1,99 @@
|
||||
# go-pdebug
|
||||
|
||||
[](https://travis-ci.org/lestrrat/go-pdebug)
|
||||
|
||||
[](https://godoc.org/github.com/lestrrat/go-pdebug)
|
||||
|
||||
Utilities for my print debugging fun. YMMV
|
||||
|
||||
# WARNING
|
||||
|
||||
This repository has been moved to [github.com/lestrrat-go/pdebug](https://github.com/lestrrat-go/pdebug). This repository exists so that libraries pointing to this URL will keep functioning, but this repository will NOT be updated in the future. Please use the new import path.
|
||||
|
||||
# Synopsis
|
||||
|
||||

|
||||
|
||||
# Description
|
||||
|
||||
Building with `pdebug` declares a constant, `pdebug.Enabled` which you
|
||||
can use to easily compile in/out depending on the presence of a build tag.
|
||||
|
||||
```go
|
||||
func Foo() {
|
||||
// will only be available if you compile with `-tags debug`
|
||||
if pdebug.Enabled {
|
||||
pdebug.Printf("Starting Foo()!
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Note that using `github.com/lestrrat/go-pdebug` and `-tags debug` only
|
||||
compiles in the code. In order to actually show the debug trace, you need
|
||||
to specify an environment variable:
|
||||
|
||||
```shell
|
||||
# For example, to show debug code during testing:
|
||||
PDEBUG_TRACE=1 go test -tags debug
|
||||
```
|
||||
|
||||
If you want to forcefully show the trace (which is handy when you're
|
||||
debugging/testing), you can use the `debug0` tag instead:
|
||||
|
||||
```shell
|
||||
go test -tags debug0
|
||||
```
|
||||
|
||||
# Markers
|
||||
|
||||
When you want to print debug a chain of function calls, you can use the
|
||||
`Marker` functions:
|
||||
|
||||
```go
|
||||
func Foo() {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("Foo")
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
pdebug.Printf("Inside Foo()!")
|
||||
}
|
||||
```
|
||||
|
||||
This will cause all of the `Printf` calls to automatically indent
|
||||
the output so it's visually easier to see where a certain trace log
|
||||
is being generated.
|
||||
|
||||
By default it will print something like:
|
||||
|
||||
```
|
||||
|DEBUG| START Foo
|
||||
|DEBUG| Inside Foo()!
|
||||
|DEBUG| END Foo (1.23μs)
|
||||
```
|
||||
|
||||
If you want to automatically show the error value you are returning
|
||||
(but only if there is an error), you can use the `BindError` method:
|
||||
|
||||
```go
|
||||
func Foo() (err error) {
|
||||
if pdebug.Enabled {
|
||||
g := pdebug.Marker("Foo").BindError(&err)
|
||||
defer g.End()
|
||||
}
|
||||
|
||||
pdebug.Printf("Inside Foo()!")
|
||||
|
||||
return errors.New("boo")
|
||||
}
|
||||
```
|
||||
|
||||
This will print something like:
|
||||
|
||||
|
||||
```
|
||||
|DEBUG| START Foo
|
||||
|DEBUG| Inside Foo()!
|
||||
|DEBUG| END Foo (1.23μs): ERROR boo
|
||||
```
|
||||
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
// +build debug
|
||||
|
||||
package pdebug
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
var Trace = false
|
||||
func init() {
|
||||
if b, err := strconv.ParseBool(os.Getenv("PDEBUG_TRACE")); err == nil && b {
|
||||
Trace = true
|
||||
}
|
||||
}
|
||||
+6
@@ -0,0 +1,6 @@
|
||||
// +build debug0
|
||||
|
||||
package pdebug
|
||||
|
||||
var Trace = true
|
||||
|
||||
+43
@@ -0,0 +1,43 @@
|
||||
package pdebug
|
||||
|
||||
import (
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type pdctx struct {
|
||||
mutex sync.Mutex
|
||||
indentL int
|
||||
LogTime bool
|
||||
Prefix string
|
||||
Writer io.Writer
|
||||
}
|
||||
|
||||
var emptyMarkerGuard = &markerg{}
|
||||
|
||||
type markerg struct {
|
||||
indentg guard
|
||||
ctx *pdctx
|
||||
f string
|
||||
args []interface{}
|
||||
start time.Time
|
||||
errptr *error
|
||||
}
|
||||
|
||||
var DefaultCtx = &pdctx{
|
||||
LogTime: true,
|
||||
Prefix: "|DEBUG| ",
|
||||
Writer: os.Stdout,
|
||||
}
|
||||
|
||||
type guard struct {
|
||||
cb func()
|
||||
}
|
||||
|
||||
func (g *guard) End() {
|
||||
if cb := g.cb; cb != nil {
|
||||
cb()
|
||||
}
|
||||
}
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
//+build !debug,!debug0
|
||||
|
||||
package pdebug
|
||||
|
||||
// Enabled is true if `-tags debug` or `-tags debug0` is used
|
||||
// during compilation. Use this to "ifdef-out" debug blocks.
|
||||
const Enabled = false
|
||||
|
||||
// Trace is true if `-tags debug` is used AND the environment
|
||||
// variable `PDEBUG_TRACE` is set to a `true` value (i.e.,
|
||||
// 1, true, etc), or `-tags debug0` is used. This allows you to
|
||||
// compile-in the trace logs, but only show them when you
|
||||
// set the environment variable
|
||||
const Trace = false
|
||||
|
||||
// IRelease is deprecated. Use Marker()/End() instead
|
||||
func (g guard) IRelease(f string, args ...interface{}) {}
|
||||
|
||||
// IPrintf is deprecated. Use Marker()/End() instead
|
||||
func IPrintf(f string, args ...interface{}) guard { return guard{} }
|
||||
|
||||
// Printf prints to standard out, just like a normal fmt.Printf,
|
||||
// but respects the indentation level set by IPrintf/IRelease.
|
||||
// Printf is no op unless you compile with the `debug` tag.
|
||||
func Printf(f string, args ...interface{}) {}
|
||||
|
||||
// Dump dumps the objects using go-spew.
|
||||
// Dump is a no op unless you compile with the `debug` tag.
|
||||
func Dump(v ...interface{}) {}
|
||||
|
||||
// Marker marks the beginning of an indented block. The message
|
||||
// you specify in the arguments is prefixed witha "START", and
|
||||
// subsequent calls to Printf will be indented one level more.
|
||||
//
|
||||
// To reset this, you must call End() on the guard object that
|
||||
// gets returned by Marker().
|
||||
func Marker(f string, args ...interface{}) *markerg { return emptyMarkerGuard }
|
||||
func (g *markerg) BindError(_ *error) *markerg { return g }
|
||||
func (g *markerg) End() {}
|
||||
+170
@@ -0,0 +1,170 @@
|
||||
// +build debug OR debug0
|
||||
|
||||
package pdebug
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/davecgh/go-spew/spew"
|
||||
)
|
||||
|
||||
const Enabled = true
|
||||
|
||||
type Guard interface {
|
||||
End()
|
||||
}
|
||||
|
||||
var emptyGuard = &guard{}
|
||||
|
||||
func (ctx *pdctx) Unindent() {
|
||||
ctx.mutex.Lock()
|
||||
defer ctx.mutex.Unlock()
|
||||
ctx.indentL--
|
||||
}
|
||||
|
||||
func (ctx *pdctx) Indent() guard {
|
||||
ctx.mutex.Lock()
|
||||
ctx.indentL++
|
||||
ctx.mutex.Unlock()
|
||||
|
||||
return guard{cb: ctx.Unindent}
|
||||
}
|
||||
|
||||
func (ctx *pdctx) preamble(buf *bytes.Buffer) {
|
||||
if p := ctx.Prefix; len(p) > 0 {
|
||||
buf.WriteString(p)
|
||||
}
|
||||
if ctx.LogTime {
|
||||
fmt.Fprintf(buf, "%0.5f ", float64(time.Now().UnixNano()) / 1000000.0)
|
||||
}
|
||||
|
||||
for i := 0; i < ctx.indentL; i++ {
|
||||
buf.WriteString(" ")
|
||||
}
|
||||
}
|
||||
|
||||
func (ctx *pdctx) Printf(f string, args ...interface{}) {
|
||||
if !strings.HasSuffix(f, "\n") {
|
||||
f = f + "\n"
|
||||
}
|
||||
buf := bytes.Buffer{}
|
||||
ctx.preamble(&buf)
|
||||
fmt.Fprintf(&buf, f, args...)
|
||||
buf.WriteTo(ctx.Writer)
|
||||
}
|
||||
|
||||
func Marker(f string, args ...interface{}) *markerg {
|
||||
return DefaultCtx.Marker(f, args...)
|
||||
}
|
||||
|
||||
func (ctx *pdctx) Marker(f string, args ...interface{}) *markerg {
|
||||
if !Trace {
|
||||
return emptyMarkerGuard
|
||||
}
|
||||
|
||||
buf := &bytes.Buffer{}
|
||||
ctx.preamble(buf)
|
||||
buf.WriteString("START ")
|
||||
fmt.Fprintf(buf, f, args...)
|
||||
if buf.Len() > 0 {
|
||||
if b := buf.Bytes(); b[buf.Len()-1] != '\n' {
|
||||
buf.WriteRune('\n')
|
||||
}
|
||||
}
|
||||
|
||||
buf.WriteTo(ctx.Writer)
|
||||
|
||||
g := ctx.Indent()
|
||||
return &markerg{
|
||||
indentg: g,
|
||||
ctx: ctx,
|
||||
f: f,
|
||||
args: args,
|
||||
start: time.Now(),
|
||||
errptr: nil,
|
||||
}
|
||||
}
|
||||
|
||||
func (g *markerg) BindError(errptr *error) *markerg {
|
||||
if g.ctx == nil {
|
||||
return g
|
||||
}
|
||||
g.ctx.mutex.Lock()
|
||||
defer g.ctx.mutex.Unlock()
|
||||
|
||||
g.errptr = errptr
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *markerg) End() {
|
||||
if g.ctx == nil {
|
||||
return
|
||||
}
|
||||
|
||||
g.indentg.End() // unindent
|
||||
buf := &bytes.Buffer{}
|
||||
g.ctx.preamble(buf)
|
||||
fmt.Fprint(buf, "END ")
|
||||
fmt.Fprintf(buf, g.f, g.args...)
|
||||
fmt.Fprintf(buf, " (%s)", time.Since(g.start))
|
||||
if errptr := g.errptr; errptr != nil && *errptr != nil {
|
||||
fmt.Fprintf(buf, ": ERROR: %s", *errptr)
|
||||
}
|
||||
|
||||
if buf.Len() > 0 {
|
||||
if b := buf.Bytes(); b[buf.Len()-1] != '\n' {
|
||||
buf.WriteRune('\n')
|
||||
}
|
||||
}
|
||||
|
||||
buf.WriteTo(g.ctx.Writer)
|
||||
}
|
||||
|
||||
type legacyg struct {
|
||||
guard
|
||||
start time.Time
|
||||
}
|
||||
|
||||
var emptylegacyg = legacyg{}
|
||||
|
||||
func (g legacyg) IRelease(f string, args ...interface{}) {
|
||||
if !Trace {
|
||||
return
|
||||
}
|
||||
g.End()
|
||||
dur := time.Since(g.start)
|
||||
Printf("%s (%s)", fmt.Sprintf(f, args...), dur)
|
||||
}
|
||||
|
||||
// IPrintf indents and then prints debug messages. Execute the callback
|
||||
// to undo the indent
|
||||
func IPrintf(f string, args ...interface{}) legacyg {
|
||||
if !Trace {
|
||||
return emptylegacyg
|
||||
}
|
||||
|
||||
DefaultCtx.Printf(f, args...)
|
||||
g := legacyg{
|
||||
guard: DefaultCtx.Indent(),
|
||||
start: time.Now(),
|
||||
}
|
||||
return g
|
||||
}
|
||||
|
||||
// Printf prints debug messages. Only available if compiled with "debug" tag
|
||||
func Printf(f string, args ...interface{}) {
|
||||
if !Trace {
|
||||
return
|
||||
}
|
||||
DefaultCtx.Printf(f, args...)
|
||||
}
|
||||
|
||||
func Dump(v ...interface{}) {
|
||||
if !Trace {
|
||||
return
|
||||
}
|
||||
spew.Dump(v...)
|
||||
}
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
// Package pdebug provides tools to produce debug logs the way the author
|
||||
// (Daisuke Maki a.k.a. lestrrat) likes. All of the functions are no-ops
|
||||
// unless you compile with the `-tags debug` option.
|
||||
//
|
||||
// When you compile your program with `-tags debug`, no trace is displayed,
|
||||
// but the code enclosed within `if pdebug.Enabled { ... }` is compiled in.
|
||||
// To show the debug trace, set the PDEBUG_TRACE environment variable to
|
||||
// true (or 1, or whatever `strconv.ParseBool` parses to true)
|
||||
//
|
||||
// If you want to show the debug trace regardless of an environment variable,
|
||||
// for example, perhaps while you are debugging or running tests, use the
|
||||
// `-tags debug0` build tag instead. This will enable the debug trace
|
||||
// forcefully
|
||||
package pdebug
|
||||
|
||||
+67
-11
@@ -32,7 +32,8 @@ func Containsf(t TestingT, s interface{}, contains interface{}, msg string, args
|
||||
return Contains(t, s, contains, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func DirExistsf(t TestingT, path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -160,7 +161,8 @@ func Falsef(t TestingT, value bool, msg string, args ...interface{}) bool {
|
||||
return False(t, value, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func FileExistsf(t TestingT, path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -267,7 +269,7 @@ func Implementsf(t TestingT, interfaceObject interface{}, object interface{}, ms
|
||||
|
||||
// InDeltaf asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// assert.InDeltaf(t, math.Pi, (22 / 7.0, "error message %s", "formatted"), 0.01)
|
||||
// assert.InDeltaf(t, math.Pi, 22/7.0, 0.01, "error message %s", "formatted")
|
||||
func InDeltaf(t TestingT, expected interface{}, actual interface{}, delta float64, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -325,14 +327,6 @@ func JSONEqf(t TestingT, expected string, actual string, msg string, args ...int
|
||||
return JSONEq(t, expected, actual, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func YAMLEqf(t TestingT, expected string, actual string, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return YAMLEq(t, expected, actual, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// Lenf asserts that the specified object has specific length.
|
||||
// Lenf also fails if the object has a type that len() not accept.
|
||||
//
|
||||
@@ -369,6 +363,17 @@ func LessOrEqualf(t TestingT, e1 interface{}, e2 interface{}, msg string, args .
|
||||
return LessOrEqual(t, e1, e2, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// Neverf asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// assert.Neverf(t, func() bool { return false; }, time.Second, 10*time.Millisecond, "error message %s", "formatted")
|
||||
func Neverf(t TestingT, condition func() bool, waitFor time.Duration, tick time.Duration, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return Never(t, condition, waitFor, tick, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// Nilf asserts that the specified object is nil.
|
||||
//
|
||||
// assert.Nilf(t, err, "error message %s", "formatted")
|
||||
@@ -379,6 +384,15 @@ func Nilf(t TestingT, object interface{}, msg string, args ...interface{}) bool
|
||||
return Nil(t, object, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// NoDirExistsf checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func NoDirExistsf(t TestingT, path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NoDirExists(t, path, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// NoErrorf asserts that a function returned no error (i.e. `nil`).
|
||||
//
|
||||
// actualObj, err := SomeFunction()
|
||||
@@ -392,6 +406,15 @@ func NoErrorf(t TestingT, err error, msg string, args ...interface{}) bool {
|
||||
return NoError(t, err, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// NoFileExistsf checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func NoFileExistsf(t TestingT, path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NoFileExists(t, path, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// NotContainsf asserts that the specified string, list(array, slice...) or map does NOT contain the
|
||||
// specified substring or element.
|
||||
//
|
||||
@@ -462,6 +485,19 @@ func NotRegexpf(t TestingT, rx interface{}, str interface{}, msg string, args ..
|
||||
return NotRegexp(t, rx, str, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// NotSamef asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// assert.NotSamef(t, ptr1, ptr2, "error message %s", "formatted")
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func NotSamef(t TestingT, expected interface{}, actual interface{}, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NotSame(t, expected, actual, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// NotSubsetf asserts that the specified list(array, slice...) contains not all
|
||||
// elements given in the specified subset(array, slice...).
|
||||
//
|
||||
@@ -491,6 +527,18 @@ func Panicsf(t TestingT, f PanicTestFunc, msg string, args ...interface{}) bool
|
||||
return Panics(t, f, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// PanicsWithErrorf asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// assert.PanicsWithErrorf(t, "crazy error", func(){ GoCrazy() }, "error message %s", "formatted")
|
||||
func PanicsWithErrorf(t TestingT, errString string, f PanicTestFunc, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return PanicsWithError(t, errString, f, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// PanicsWithValuef asserts that the code inside the specified PanicTestFunc panics, and that
|
||||
// the recovered panic value equals the expected panic value.
|
||||
//
|
||||
@@ -557,6 +605,14 @@ func WithinDurationf(t TestingT, expected time.Time, actual time.Time, delta tim
|
||||
return WithinDuration(t, expected, actual, delta, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func YAMLEqf(t TestingT, expected string, actual string, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return YAMLEq(t, expected, actual, append([]interface{}{msg}, args...)...)
|
||||
}
|
||||
|
||||
// Zerof asserts that i is the zero value for its type.
|
||||
func Zerof(t TestingT, i interface{}, msg string, args ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
|
||||
+134
-22
@@ -53,7 +53,8 @@ func (a *Assertions) Containsf(s interface{}, contains interface{}, msg string,
|
||||
return Containsf(a.t, s, contains, msg, args...)
|
||||
}
|
||||
|
||||
// DirExists checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExists checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func (a *Assertions) DirExists(path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -61,7 +62,8 @@ func (a *Assertions) DirExists(path string, msgAndArgs ...interface{}) bool {
|
||||
return DirExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func (a *Assertions) DirExistsf(path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -309,7 +311,8 @@ func (a *Assertions) Falsef(value bool, msg string, args ...interface{}) bool {
|
||||
return Falsef(a.t, value, msg, args...)
|
||||
}
|
||||
|
||||
// FileExists checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExists checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func (a *Assertions) FileExists(path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -317,7 +320,8 @@ func (a *Assertions) FileExists(path string, msgAndArgs ...interface{}) bool {
|
||||
return FileExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func (a *Assertions) FileExistsf(path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -521,7 +525,7 @@ func (a *Assertions) Implementsf(interfaceObject interface{}, object interface{}
|
||||
|
||||
// InDelta asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// a.InDelta(math.Pi, (22 / 7.0), 0.01)
|
||||
// a.InDelta(math.Pi, 22/7.0, 0.01)
|
||||
func (a *Assertions) InDelta(expected interface{}, actual interface{}, delta float64, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -563,7 +567,7 @@ func (a *Assertions) InDeltaSlicef(expected interface{}, actual interface{}, del
|
||||
|
||||
// InDeltaf asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// a.InDeltaf(math.Pi, (22 / 7.0, "error message %s", "formatted"), 0.01)
|
||||
// a.InDeltaf(math.Pi, 22/7.0, 0.01, "error message %s", "formatted")
|
||||
func (a *Assertions) InDeltaf(expected interface{}, actual interface{}, delta float64, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -639,22 +643,6 @@ func (a *Assertions) JSONEqf(expected string, actual string, msg string, args ..
|
||||
return JSONEqf(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// YAMLEq asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEq(expected string, actual string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return YAMLEq(a.t, expected, actual, msgAndArgs...)
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEqf(expected string, actual string, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return YAMLEqf(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// Len asserts that the specified object has specific length.
|
||||
// Len also fails if the object has a type that len() not accept.
|
||||
//
|
||||
@@ -727,6 +715,28 @@ func (a *Assertions) Lessf(e1 interface{}, e2 interface{}, msg string, args ...i
|
||||
return Lessf(a.t, e1, e2, msg, args...)
|
||||
}
|
||||
|
||||
// Never asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// a.Never(func() bool { return false; }, time.Second, 10*time.Millisecond)
|
||||
func (a *Assertions) Never(condition func() bool, waitFor time.Duration, tick time.Duration, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return Never(a.t, condition, waitFor, tick, msgAndArgs...)
|
||||
}
|
||||
|
||||
// Neverf asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// a.Neverf(func() bool { return false; }, time.Second, 10*time.Millisecond, "error message %s", "formatted")
|
||||
func (a *Assertions) Neverf(condition func() bool, waitFor time.Duration, tick time.Duration, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return Neverf(a.t, condition, waitFor, tick, msg, args...)
|
||||
}
|
||||
|
||||
// Nil asserts that the specified object is nil.
|
||||
//
|
||||
// a.Nil(err)
|
||||
@@ -747,6 +757,24 @@ func (a *Assertions) Nilf(object interface{}, msg string, args ...interface{}) b
|
||||
return Nilf(a.t, object, msg, args...)
|
||||
}
|
||||
|
||||
// NoDirExists checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func (a *Assertions) NoDirExists(path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NoDirExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// NoDirExistsf checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func (a *Assertions) NoDirExistsf(path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NoDirExistsf(a.t, path, msg, args...)
|
||||
}
|
||||
|
||||
// NoError asserts that a function returned no error (i.e. `nil`).
|
||||
//
|
||||
// actualObj, err := SomeFunction()
|
||||
@@ -773,6 +801,24 @@ func (a *Assertions) NoErrorf(err error, msg string, args ...interface{}) bool {
|
||||
return NoErrorf(a.t, err, msg, args...)
|
||||
}
|
||||
|
||||
// NoFileExists checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func (a *Assertions) NoFileExists(path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NoFileExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// NoFileExistsf checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func (a *Assertions) NoFileExistsf(path string, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NoFileExistsf(a.t, path, msg, args...)
|
||||
}
|
||||
|
||||
// NotContains asserts that the specified string, list(array, slice...) or map does NOT contain the
|
||||
// specified substring or element.
|
||||
//
|
||||
@@ -913,6 +959,32 @@ func (a *Assertions) NotRegexpf(rx interface{}, str interface{}, msg string, arg
|
||||
return NotRegexpf(a.t, rx, str, msg, args...)
|
||||
}
|
||||
|
||||
// NotSame asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// a.NotSame(ptr1, ptr2)
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func (a *Assertions) NotSame(expected interface{}, actual interface{}, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NotSame(a.t, expected, actual, msgAndArgs...)
|
||||
}
|
||||
|
||||
// NotSamef asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// a.NotSamef(ptr1, ptr2, "error message %s", "formatted")
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func (a *Assertions) NotSamef(expected interface{}, actual interface{}, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return NotSamef(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// NotSubset asserts that the specified list(array, slice...) contains not all
|
||||
// elements given in the specified subset(array, slice...).
|
||||
//
|
||||
@@ -961,6 +1033,30 @@ func (a *Assertions) Panics(f PanicTestFunc, msgAndArgs ...interface{}) bool {
|
||||
return Panics(a.t, f, msgAndArgs...)
|
||||
}
|
||||
|
||||
// PanicsWithError asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// a.PanicsWithError("crazy error", func(){ GoCrazy() })
|
||||
func (a *Assertions) PanicsWithError(errString string, f PanicTestFunc, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return PanicsWithError(a.t, errString, f, msgAndArgs...)
|
||||
}
|
||||
|
||||
// PanicsWithErrorf asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// a.PanicsWithErrorf("crazy error", func(){ GoCrazy() }, "error message %s", "formatted")
|
||||
func (a *Assertions) PanicsWithErrorf(errString string, f PanicTestFunc, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return PanicsWithErrorf(a.t, errString, f, msg, args...)
|
||||
}
|
||||
|
||||
// PanicsWithValue asserts that the code inside the specified PanicTestFunc panics, and that
|
||||
// the recovered panic value equals the expected panic value.
|
||||
//
|
||||
@@ -1103,6 +1199,22 @@ func (a *Assertions) WithinDurationf(expected time.Time, actual time.Time, delta
|
||||
return WithinDurationf(a.t, expected, actual, delta, msg, args...)
|
||||
}
|
||||
|
||||
// YAMLEq asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEq(expected string, actual string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return YAMLEq(a.t, expected, actual, msgAndArgs...)
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEqf(expected string, actual string, msg string, args ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
return YAMLEqf(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// Zero asserts that i is the zero value for its type.
|
||||
func (a *Assertions) Zero(i interface{}, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
|
||||
+173
-45
@@ -11,6 +11,7 @@ import (
|
||||
"reflect"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
@@ -21,7 +22,7 @@ import (
|
||||
yaml "gopkg.in/yaml.v2"
|
||||
)
|
||||
|
||||
//go:generate go run ../_codegen/main.go -output-package=assert -template=assertion_format.go.tmpl
|
||||
//go:generate sh -c "cd ../_codegen && go build && cd - && ../_codegen/_codegen -output-package=assert -template=assertion_format.go.tmpl"
|
||||
|
||||
// TestingT is an interface wrapper around *testing.T
|
||||
type TestingT interface {
|
||||
@@ -351,6 +352,19 @@ func Equal(t TestingT, expected, actual interface{}, msgAndArgs ...interface{})
|
||||
|
||||
}
|
||||
|
||||
// validateEqualArgs checks whether provided arguments can be safely used in the
|
||||
// Equal/NotEqual functions.
|
||||
func validateEqualArgs(expected, actual interface{}) error {
|
||||
if expected == nil && actual == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if isFunction(expected) || isFunction(actual) {
|
||||
return errors.New("cannot take func type as argument")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Same asserts that two pointers reference the same object.
|
||||
//
|
||||
// assert.Same(t, ptr1, ptr2)
|
||||
@@ -362,18 +376,7 @@ func Same(t TestingT, expected, actual interface{}, msgAndArgs ...interface{}) b
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
expectedPtr, actualPtr := reflect.ValueOf(expected), reflect.ValueOf(actual)
|
||||
if expectedPtr.Kind() != reflect.Ptr || actualPtr.Kind() != reflect.Ptr {
|
||||
return Fail(t, "Invalid operation: both arguments must be pointers", msgAndArgs...)
|
||||
}
|
||||
|
||||
expectedType, actualType := reflect.TypeOf(expected), reflect.TypeOf(actual)
|
||||
if expectedType != actualType {
|
||||
return Fail(t, fmt.Sprintf("Pointer expected to be of type %v, but was %v",
|
||||
expectedType, actualType), msgAndArgs...)
|
||||
}
|
||||
|
||||
if expected != actual {
|
||||
if !samePointers(expected, actual) {
|
||||
return Fail(t, fmt.Sprintf("Not same: \n"+
|
||||
"expected: %p %#v\n"+
|
||||
"actual : %p %#v", expected, expected, actual, actual), msgAndArgs...)
|
||||
@@ -382,6 +385,42 @@ func Same(t TestingT, expected, actual interface{}, msgAndArgs ...interface{}) b
|
||||
return true
|
||||
}
|
||||
|
||||
// NotSame asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// assert.NotSame(t, ptr1, ptr2)
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func NotSame(t TestingT, expected, actual interface{}, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
if samePointers(expected, actual) {
|
||||
return Fail(t, fmt.Sprintf(
|
||||
"Expected and actual point to the same object: %p %#v",
|
||||
expected, expected), msgAndArgs...)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// samePointers compares two generic interface objects and returns whether
|
||||
// they point to the same object
|
||||
func samePointers(first, second interface{}) bool {
|
||||
firstPtr, secondPtr := reflect.ValueOf(first), reflect.ValueOf(second)
|
||||
if firstPtr.Kind() != reflect.Ptr || secondPtr.Kind() != reflect.Ptr {
|
||||
return false
|
||||
}
|
||||
|
||||
firstType, secondType := reflect.TypeOf(first), reflect.TypeOf(second)
|
||||
if firstType != secondType {
|
||||
return false
|
||||
}
|
||||
|
||||
// compare pointer addresses
|
||||
return first == second
|
||||
}
|
||||
|
||||
// formatUnequalValues takes two values of arbitrary types and returns string
|
||||
// representations appropriate to be presented to the user.
|
||||
//
|
||||
@@ -393,9 +432,11 @@ func formatUnequalValues(expected, actual interface{}) (e string, a string) {
|
||||
return fmt.Sprintf("%T(%#v)", expected, expected),
|
||||
fmt.Sprintf("%T(%#v)", actual, actual)
|
||||
}
|
||||
|
||||
return fmt.Sprintf("%#v", expected),
|
||||
fmt.Sprintf("%#v", actual)
|
||||
switch expected.(type) {
|
||||
case time.Duration:
|
||||
return fmt.Sprintf("%v", expected), fmt.Sprintf("%v", actual)
|
||||
}
|
||||
return fmt.Sprintf("%#v", expected), fmt.Sprintf("%#v", actual)
|
||||
}
|
||||
|
||||
// EqualValues asserts that two objects are equal or convertable to the same types
|
||||
@@ -901,15 +942,17 @@ func Condition(t TestingT, comp Comparison, msgAndArgs ...interface{}) bool {
|
||||
type PanicTestFunc func()
|
||||
|
||||
// didPanic returns true if the function passed to it panics. Otherwise, it returns false.
|
||||
func didPanic(f PanicTestFunc) (bool, interface{}) {
|
||||
func didPanic(f PanicTestFunc) (bool, interface{}, string) {
|
||||
|
||||
didPanic := false
|
||||
var message interface{}
|
||||
var stack string
|
||||
func() {
|
||||
|
||||
defer func() {
|
||||
if message = recover(); message != nil {
|
||||
didPanic = true
|
||||
stack = string(debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -918,7 +961,7 @@ func didPanic(f PanicTestFunc) (bool, interface{}) {
|
||||
|
||||
}()
|
||||
|
||||
return didPanic, message
|
||||
return didPanic, message, stack
|
||||
|
||||
}
|
||||
|
||||
@@ -930,7 +973,7 @@ func Panics(t TestingT, f PanicTestFunc, msgAndArgs ...interface{}) bool {
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
if funcDidPanic, panicValue := didPanic(f); !funcDidPanic {
|
||||
if funcDidPanic, panicValue, _ := didPanic(f); !funcDidPanic {
|
||||
return Fail(t, fmt.Sprintf("func %#v should panic\n\tPanic value:\t%#v", f, panicValue), msgAndArgs...)
|
||||
}
|
||||
|
||||
@@ -946,12 +989,34 @@ func PanicsWithValue(t TestingT, expected interface{}, f PanicTestFunc, msgAndAr
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
funcDidPanic, panicValue := didPanic(f)
|
||||
funcDidPanic, panicValue, panickedStack := didPanic(f)
|
||||
if !funcDidPanic {
|
||||
return Fail(t, fmt.Sprintf("func %#v should panic\n\tPanic value:\t%#v", f, panicValue), msgAndArgs...)
|
||||
}
|
||||
if panicValue != expected {
|
||||
return Fail(t, fmt.Sprintf("func %#v should panic with value:\t%#v\n\tPanic value:\t%#v", f, expected, panicValue), msgAndArgs...)
|
||||
return Fail(t, fmt.Sprintf("func %#v should panic with value:\t%#v\n\tPanic value:\t%#v\n\tPanic stack:\t%s", f, expected, panicValue, panickedStack), msgAndArgs...)
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// PanicsWithError asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// assert.PanicsWithError(t, "crazy error", func(){ GoCrazy() })
|
||||
func PanicsWithError(t TestingT, errString string, f PanicTestFunc, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
funcDidPanic, panicValue, panickedStack := didPanic(f)
|
||||
if !funcDidPanic {
|
||||
return Fail(t, fmt.Sprintf("func %#v should panic\n\tPanic value:\t%#v", f, panicValue), msgAndArgs...)
|
||||
}
|
||||
panicErr, ok := panicValue.(error)
|
||||
if !ok || panicErr.Error() != errString {
|
||||
return Fail(t, fmt.Sprintf("func %#v should panic with error message:\t%#v\n\tPanic value:\t%#v\n\tPanic stack:\t%s", f, errString, panicValue, panickedStack), msgAndArgs...)
|
||||
}
|
||||
|
||||
return true
|
||||
@@ -965,8 +1030,8 @@ func NotPanics(t TestingT, f PanicTestFunc, msgAndArgs ...interface{}) bool {
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
if funcDidPanic, panicValue := didPanic(f); funcDidPanic {
|
||||
return Fail(t, fmt.Sprintf("func %#v should not panic\n\tPanic value:\t%v", f, panicValue), msgAndArgs...)
|
||||
if funcDidPanic, panicValue, panickedStack := didPanic(f); funcDidPanic {
|
||||
return Fail(t, fmt.Sprintf("func %#v should not panic\n\tPanic value:\t%v\n\tPanic stack:\t%s", f, panicValue, panickedStack), msgAndArgs...)
|
||||
}
|
||||
|
||||
return true
|
||||
@@ -1026,7 +1091,7 @@ func toFloat(x interface{}) (float64, bool) {
|
||||
|
||||
// InDelta asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// assert.InDelta(t, math.Pi, (22 / 7.0), 0.01)
|
||||
// assert.InDelta(t, math.Pi, 22/7.0, 0.01)
|
||||
func InDelta(t TestingT, expected, actual interface{}, delta float64, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -1314,7 +1379,8 @@ func NotZero(t TestingT, i interface{}, msgAndArgs ...interface{}) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// FileExists checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExists checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func FileExists(t TestingT, path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -1332,7 +1398,24 @@ func FileExists(t TestingT, path string, msgAndArgs ...interface{}) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// DirExists checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// NoFileExists checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func NoFileExists(t TestingT, path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
if info.IsDir() {
|
||||
return true
|
||||
}
|
||||
return Fail(t, fmt.Sprintf("file %q exists", path), msgAndArgs...)
|
||||
}
|
||||
|
||||
// DirExists checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func DirExists(t TestingT, path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -1350,6 +1433,25 @@ func DirExists(t TestingT, path string, msgAndArgs ...interface{}) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// NoDirExists checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func NoDirExists(t TestingT, path string, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return true
|
||||
}
|
||||
return true
|
||||
}
|
||||
if !info.IsDir() {
|
||||
return true
|
||||
}
|
||||
return Fail(t, fmt.Sprintf("directory %q exists", path), msgAndArgs...)
|
||||
}
|
||||
|
||||
// JSONEq asserts that two JSON strings are equivalent.
|
||||
//
|
||||
// assert.JSONEq(t, `{"hello": "world", "foo": "bar"}`, `{"foo": "bar", "hello": "world"}`)
|
||||
@@ -1439,15 +1541,6 @@ func diff(expected interface{}, actual interface{}) string {
|
||||
return "\n\nDiff:\n" + diff
|
||||
}
|
||||
|
||||
// validateEqualArgs checks whether provided arguments can be safely used in the
|
||||
// Equal/NotEqual functions.
|
||||
func validateEqualArgs(expected, actual interface{}) error {
|
||||
if isFunction(expected) || isFunction(actual) {
|
||||
return errors.New("cannot take func type as argument")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isFunction(arg interface{}) bool {
|
||||
if arg == nil {
|
||||
return false
|
||||
@@ -1475,24 +1568,59 @@ func Eventually(t TestingT, condition func() bool, waitFor time.Duration, tick t
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
ch := make(chan bool, 1)
|
||||
|
||||
timer := time.NewTimer(waitFor)
|
||||
ticker := time.NewTicker(tick)
|
||||
checkPassed := make(chan bool)
|
||||
defer timer.Stop()
|
||||
|
||||
ticker := time.NewTicker(tick)
|
||||
defer ticker.Stop()
|
||||
defer close(checkPassed)
|
||||
for {
|
||||
|
||||
for tick := ticker.C; ; {
|
||||
select {
|
||||
case <-timer.C:
|
||||
return Fail(t, "Condition never satisfied", msgAndArgs...)
|
||||
case result := <-checkPassed:
|
||||
if result {
|
||||
case <-tick:
|
||||
tick = nil
|
||||
go func() { ch <- condition() }()
|
||||
case v := <-ch:
|
||||
if v {
|
||||
return true
|
||||
}
|
||||
case <-ticker.C:
|
||||
go func() {
|
||||
checkPassed <- condition()
|
||||
}()
|
||||
tick = ticker.C
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Never asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// assert.Never(t, func() bool { return false; }, time.Second, 10*time.Millisecond)
|
||||
func Never(t TestingT, condition func() bool, waitFor time.Duration, tick time.Duration, msgAndArgs ...interface{}) bool {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
|
||||
ch := make(chan bool, 1)
|
||||
|
||||
timer := time.NewTimer(waitFor)
|
||||
defer timer.Stop()
|
||||
|
||||
ticker := time.NewTicker(tick)
|
||||
defer ticker.Stop()
|
||||
|
||||
for tick := ticker.C; ; {
|
||||
select {
|
||||
case <-timer.C:
|
||||
return true
|
||||
case <-tick:
|
||||
tick = nil
|
||||
go func() { ch <- condition() }()
|
||||
case v := <-ch:
|
||||
if v {
|
||||
return Fail(t, "Condition satisfied", msgAndArgs...)
|
||||
}
|
||||
tick = ticker.C
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -13,4 +13,4 @@ func New(t TestingT) *Assertions {
|
||||
}
|
||||
}
|
||||
|
||||
//go:generate go run ../_codegen/main.go -output-package=assert -template=assertion_forward.go.tmpl -include-format-funcs
|
||||
//go:generate sh -c "cd ../_codegen && go build && cd - && ../_codegen/_codegen -output-package=assert -template=assertion_forward.go.tmpl -include-format-funcs"
|
||||
|
||||
+24
-1
@@ -147,7 +147,7 @@ func (c *Call) After(d time.Duration) *Call {
|
||||
}
|
||||
|
||||
// Run sets a handler to be called before returning. It can be used when
|
||||
// mocking a method such as unmarshalers that takes a pointer to a struct and
|
||||
// mocking a method (such as an unmarshaler) that takes a pointer to a struct and
|
||||
// sets properties in such struct
|
||||
//
|
||||
// Mock.On("Unmarshal", AnythingOfType("*map[string]interface{}").Return().Run(func(args Arguments) {
|
||||
@@ -578,6 +578,23 @@ func AnythingOfType(t string) AnythingOfTypeArgument {
|
||||
return AnythingOfTypeArgument(t)
|
||||
}
|
||||
|
||||
// IsTypeArgument is a struct that contains the type of an argument
|
||||
// for use when type checking. This is an alternative to AnythingOfType.
|
||||
// Used in Diff and Assert.
|
||||
type IsTypeArgument struct {
|
||||
t interface{}
|
||||
}
|
||||
|
||||
// IsType returns an IsTypeArgument object containing the type to check for.
|
||||
// You can provide a zero-value of the type to check. This is an
|
||||
// alternative to AnythingOfType. Used in Diff and Assert.
|
||||
//
|
||||
// For example:
|
||||
// Assert(t, IsType(""), IsType(0))
|
||||
func IsType(t interface{}) *IsTypeArgument {
|
||||
return &IsTypeArgument{t: t}
|
||||
}
|
||||
|
||||
// argumentMatcher performs custom argument matching, returning whether or
|
||||
// not the argument is matched by the expectation fixture function.
|
||||
type argumentMatcher struct {
|
||||
@@ -711,6 +728,12 @@ func (args Arguments) Diff(objects []interface{}) (string, int) {
|
||||
output = fmt.Sprintf("%s\t%d: FAIL: type %s != type %s - %s\n", output, i, expected, reflect.TypeOf(actual).Name(), actualFmt)
|
||||
}
|
||||
|
||||
} else if reflect.TypeOf(expected) == reflect.TypeOf((*IsTypeArgument)(nil)) {
|
||||
t := expected.(*IsTypeArgument).t
|
||||
if reflect.TypeOf(t) != reflect.TypeOf(actual) {
|
||||
differences++
|
||||
output = fmt.Sprintf("%s\t%d: FAIL: type %s != type %s - %s\n", output, i, reflect.TypeOf(t).Name(), reflect.TypeOf(actual).Name(), actualFmt)
|
||||
}
|
||||
} else {
|
||||
|
||||
// normal checking
|
||||
|
||||
+1
-1
@@ -13,4 +13,4 @@ func New(t TestingT) *Assertions {
|
||||
}
|
||||
}
|
||||
|
||||
//go:generate go run ../_codegen/main.go -output-package=require -template=require_forward.go.tmpl -include-format-funcs
|
||||
//go:generate sh -c "cd ../_codegen && go build && cd - && ../_codegen/_codegen -output-package=require -template=require_forward.go.tmpl -include-format-funcs"
|
||||
|
||||
+176
-34
@@ -66,7 +66,8 @@ func Containsf(t TestingT, s interface{}, contains interface{}, msg string, args
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// DirExists checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExists checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func DirExists(t TestingT, path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -77,7 +78,8 @@ func DirExists(t TestingT, path string, msgAndArgs ...interface{}) {
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func DirExistsf(t TestingT, path string, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -275,12 +277,12 @@ func Errorf(t TestingT, err error, msg string, args ...interface{}) {
|
||||
//
|
||||
// assert.Eventually(t, func() bool { return true; }, time.Second, 10*time.Millisecond)
|
||||
func Eventually(t TestingT, condition func() bool, waitFor time.Duration, tick time.Duration, msgAndArgs ...interface{}) {
|
||||
if assert.Eventually(t, condition, waitFor, tick, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.Eventually(t, condition, waitFor, tick, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
@@ -289,12 +291,12 @@ func Eventually(t TestingT, condition func() bool, waitFor time.Duration, tick t
|
||||
//
|
||||
// assert.Eventuallyf(t, func() bool { return true; }, time.Second, 10*time.Millisecond, "error message %s", "formatted")
|
||||
func Eventuallyf(t TestingT, condition func() bool, waitFor time.Duration, tick time.Duration, msg string, args ...interface{}) {
|
||||
if assert.Eventuallyf(t, condition, waitFor, tick, msg, args...) {
|
||||
return
|
||||
}
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.Eventuallyf(t, condition, waitFor, tick, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
@@ -394,7 +396,8 @@ func Falsef(t TestingT, value bool, msg string, args ...interface{}) {
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// FileExists checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExists checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func FileExists(t TestingT, path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -405,7 +408,8 @@ func FileExists(t TestingT, path string, msgAndArgs ...interface{}) {
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func FileExistsf(t TestingT, path string, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -660,7 +664,7 @@ func Implementsf(t TestingT, interfaceObject interface{}, object interface{}, ms
|
||||
|
||||
// InDelta asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// assert.InDelta(t, math.Pi, (22 / 7.0), 0.01)
|
||||
// assert.InDelta(t, math.Pi, 22/7.0, 0.01)
|
||||
func InDelta(t TestingT, expected interface{}, actual interface{}, delta float64, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -717,7 +721,7 @@ func InDeltaSlicef(t TestingT, expected interface{}, actual interface{}, delta f
|
||||
|
||||
// InDeltaf asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// assert.InDeltaf(t, math.Pi, (22 / 7.0, "error message %s", "formatted"), 0.01)
|
||||
// assert.InDeltaf(t, math.Pi, 22/7.0, 0.01, "error message %s", "formatted")
|
||||
func InDeltaf(t TestingT, expected interface{}, actual interface{}, delta float64, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -820,28 +824,6 @@ func JSONEqf(t TestingT, expected string, actual string, msg string, args ...int
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// YAMLEq asserts that two YAML strings are equivalent.
|
||||
func YAMLEq(t TestingT, expected string, actual string, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.YAMLEq(t, expected, actual, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func YAMLEqf(t TestingT, expected string, actual string, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.YAMLEqf(t, expected, actual, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// Len asserts that the specified object has specific length.
|
||||
// Len also fails if the object has a type that len() not accept.
|
||||
//
|
||||
@@ -932,6 +914,34 @@ func Lessf(t TestingT, e1 interface{}, e2 interface{}, msg string, args ...inter
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// Never asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// assert.Never(t, func() bool { return false; }, time.Second, 10*time.Millisecond)
|
||||
func Never(t TestingT, condition func() bool, waitFor time.Duration, tick time.Duration, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.Never(t, condition, waitFor, tick, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// Neverf asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// assert.Neverf(t, func() bool { return false; }, time.Second, 10*time.Millisecond, "error message %s", "formatted")
|
||||
func Neverf(t TestingT, condition func() bool, waitFor time.Duration, tick time.Duration, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.Neverf(t, condition, waitFor, tick, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// Nil asserts that the specified object is nil.
|
||||
//
|
||||
// assert.Nil(t, err)
|
||||
@@ -958,6 +968,30 @@ func Nilf(t TestingT, object interface{}, msg string, args ...interface{}) {
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NoDirExists checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func NoDirExists(t TestingT, path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.NoDirExists(t, path, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NoDirExistsf checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func NoDirExistsf(t TestingT, path string, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.NoDirExistsf(t, path, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NoError asserts that a function returned no error (i.e. `nil`).
|
||||
//
|
||||
// actualObj, err := SomeFunction()
|
||||
@@ -990,6 +1024,30 @@ func NoErrorf(t TestingT, err error, msg string, args ...interface{}) {
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NoFileExists checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func NoFileExists(t TestingT, path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.NoFileExists(t, path, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NoFileExistsf checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func NoFileExistsf(t TestingT, path string, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.NoFileExistsf(t, path, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NotContains asserts that the specified string, list(array, slice...) or map does NOT contain the
|
||||
// specified substring or element.
|
||||
//
|
||||
@@ -1166,6 +1224,38 @@ func NotRegexpf(t TestingT, rx interface{}, str interface{}, msg string, args ..
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NotSame asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// assert.NotSame(t, ptr1, ptr2)
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func NotSame(t TestingT, expected interface{}, actual interface{}, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.NotSame(t, expected, actual, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NotSamef asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// assert.NotSamef(t, ptr1, ptr2, "error message %s", "formatted")
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func NotSamef(t TestingT, expected interface{}, actual interface{}, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.NotSamef(t, expected, actual, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// NotSubset asserts that the specified list(array, slice...) contains not all
|
||||
// elements given in the specified subset(array, slice...).
|
||||
//
|
||||
@@ -1229,6 +1319,36 @@ func Panics(t TestingT, f assert.PanicTestFunc, msgAndArgs ...interface{}) {
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// PanicsWithError asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// assert.PanicsWithError(t, "crazy error", func(){ GoCrazy() })
|
||||
func PanicsWithError(t TestingT, errString string, f assert.PanicTestFunc, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.PanicsWithError(t, errString, f, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// PanicsWithErrorf asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// assert.PanicsWithErrorf(t, "crazy error", func(){ GoCrazy() }, "error message %s", "formatted")
|
||||
func PanicsWithErrorf(t TestingT, errString string, f assert.PanicTestFunc, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.PanicsWithErrorf(t, errString, f, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// PanicsWithValue asserts that the code inside the specified PanicTestFunc panics, and that
|
||||
// the recovered panic value equals the expected panic value.
|
||||
//
|
||||
@@ -1410,6 +1530,28 @@ func WithinDurationf(t TestingT, expected time.Time, actual time.Time, delta tim
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// YAMLEq asserts that two YAML strings are equivalent.
|
||||
func YAMLEq(t TestingT, expected string, actual string, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.YAMLEq(t, expected, actual, msgAndArgs...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func YAMLEqf(t TestingT, expected string, actual string, msg string, args ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
if assert.YAMLEqf(t, expected, actual, msg, args...) {
|
||||
return
|
||||
}
|
||||
t.FailNow()
|
||||
}
|
||||
|
||||
// Zero asserts that i is the zero value for its type.
|
||||
func Zero(t TestingT, i interface{}, msgAndArgs ...interface{}) {
|
||||
if h, ok := t.(tHelper); ok {
|
||||
|
||||
+134
-22
@@ -54,7 +54,8 @@ func (a *Assertions) Containsf(s interface{}, contains interface{}, msg string,
|
||||
Containsf(a.t, s, contains, msg, args...)
|
||||
}
|
||||
|
||||
// DirExists checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExists checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func (a *Assertions) DirExists(path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -62,7 +63,8 @@ func (a *Assertions) DirExists(path string, msgAndArgs ...interface{}) {
|
||||
DirExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
// DirExistsf checks whether a directory exists in the given path. It also fails
|
||||
// if the path is a file rather a directory or there is an error checking whether it exists.
|
||||
func (a *Assertions) DirExistsf(path string, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -310,7 +312,8 @@ func (a *Assertions) Falsef(value bool, msg string, args ...interface{}) {
|
||||
Falsef(a.t, value, msg, args...)
|
||||
}
|
||||
|
||||
// FileExists checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExists checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func (a *Assertions) FileExists(path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -318,7 +321,8 @@ func (a *Assertions) FileExists(path string, msgAndArgs ...interface{}) {
|
||||
FileExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if the path points to a directory or there is an error when trying to check the file.
|
||||
// FileExistsf checks whether a file exists in the given path. It also fails if
|
||||
// the path points to a directory or there is an error when trying to check the file.
|
||||
func (a *Assertions) FileExistsf(path string, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -522,7 +526,7 @@ func (a *Assertions) Implementsf(interfaceObject interface{}, object interface{}
|
||||
|
||||
// InDelta asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// a.InDelta(math.Pi, (22 / 7.0), 0.01)
|
||||
// a.InDelta(math.Pi, 22/7.0, 0.01)
|
||||
func (a *Assertions) InDelta(expected interface{}, actual interface{}, delta float64, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -564,7 +568,7 @@ func (a *Assertions) InDeltaSlicef(expected interface{}, actual interface{}, del
|
||||
|
||||
// InDeltaf asserts that the two numerals are within delta of each other.
|
||||
//
|
||||
// a.InDeltaf(math.Pi, (22 / 7.0, "error message %s", "formatted"), 0.01)
|
||||
// a.InDeltaf(math.Pi, 22/7.0, 0.01, "error message %s", "formatted")
|
||||
func (a *Assertions) InDeltaf(expected interface{}, actual interface{}, delta float64, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
@@ -640,22 +644,6 @@ func (a *Assertions) JSONEqf(expected string, actual string, msg string, args ..
|
||||
JSONEqf(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// YAMLEq asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEq(expected string, actual string, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
YAMLEq(a.t, expected, actual, msgAndArgs...)
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEqf(expected string, actual string, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
YAMLEqf(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// Len asserts that the specified object has specific length.
|
||||
// Len also fails if the object has a type that len() not accept.
|
||||
//
|
||||
@@ -728,6 +716,28 @@ func (a *Assertions) Lessf(e1 interface{}, e2 interface{}, msg string, args ...i
|
||||
Lessf(a.t, e1, e2, msg, args...)
|
||||
}
|
||||
|
||||
// Never asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// a.Never(func() bool { return false; }, time.Second, 10*time.Millisecond)
|
||||
func (a *Assertions) Never(condition func() bool, waitFor time.Duration, tick time.Duration, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
Never(a.t, condition, waitFor, tick, msgAndArgs...)
|
||||
}
|
||||
|
||||
// Neverf asserts that the given condition doesn't satisfy in waitFor time,
|
||||
// periodically checking the target function each tick.
|
||||
//
|
||||
// a.Neverf(func() bool { return false; }, time.Second, 10*time.Millisecond, "error message %s", "formatted")
|
||||
func (a *Assertions) Neverf(condition func() bool, waitFor time.Duration, tick time.Duration, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
Neverf(a.t, condition, waitFor, tick, msg, args...)
|
||||
}
|
||||
|
||||
// Nil asserts that the specified object is nil.
|
||||
//
|
||||
// a.Nil(err)
|
||||
@@ -748,6 +758,24 @@ func (a *Assertions) Nilf(object interface{}, msg string, args ...interface{}) {
|
||||
Nilf(a.t, object, msg, args...)
|
||||
}
|
||||
|
||||
// NoDirExists checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func (a *Assertions) NoDirExists(path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
NoDirExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// NoDirExistsf checks whether a directory does not exist in the given path.
|
||||
// It fails if the path points to an existing _directory_ only.
|
||||
func (a *Assertions) NoDirExistsf(path string, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
NoDirExistsf(a.t, path, msg, args...)
|
||||
}
|
||||
|
||||
// NoError asserts that a function returned no error (i.e. `nil`).
|
||||
//
|
||||
// actualObj, err := SomeFunction()
|
||||
@@ -774,6 +802,24 @@ func (a *Assertions) NoErrorf(err error, msg string, args ...interface{}) {
|
||||
NoErrorf(a.t, err, msg, args...)
|
||||
}
|
||||
|
||||
// NoFileExists checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func (a *Assertions) NoFileExists(path string, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
NoFileExists(a.t, path, msgAndArgs...)
|
||||
}
|
||||
|
||||
// NoFileExistsf checks whether a file does not exist in a given path. It fails
|
||||
// if the path points to an existing _file_ only.
|
||||
func (a *Assertions) NoFileExistsf(path string, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
NoFileExistsf(a.t, path, msg, args...)
|
||||
}
|
||||
|
||||
// NotContains asserts that the specified string, list(array, slice...) or map does NOT contain the
|
||||
// specified substring or element.
|
||||
//
|
||||
@@ -914,6 +960,32 @@ func (a *Assertions) NotRegexpf(rx interface{}, str interface{}, msg string, arg
|
||||
NotRegexpf(a.t, rx, str, msg, args...)
|
||||
}
|
||||
|
||||
// NotSame asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// a.NotSame(ptr1, ptr2)
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func (a *Assertions) NotSame(expected interface{}, actual interface{}, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
NotSame(a.t, expected, actual, msgAndArgs...)
|
||||
}
|
||||
|
||||
// NotSamef asserts that two pointers do not reference the same object.
|
||||
//
|
||||
// a.NotSamef(ptr1, ptr2, "error message %s", "formatted")
|
||||
//
|
||||
// Both arguments must be pointer variables. Pointer variable sameness is
|
||||
// determined based on the equality of both type and value.
|
||||
func (a *Assertions) NotSamef(expected interface{}, actual interface{}, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
NotSamef(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// NotSubset asserts that the specified list(array, slice...) contains not all
|
||||
// elements given in the specified subset(array, slice...).
|
||||
//
|
||||
@@ -962,6 +1034,30 @@ func (a *Assertions) Panics(f assert.PanicTestFunc, msgAndArgs ...interface{}) {
|
||||
Panics(a.t, f, msgAndArgs...)
|
||||
}
|
||||
|
||||
// PanicsWithError asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// a.PanicsWithError("crazy error", func(){ GoCrazy() })
|
||||
func (a *Assertions) PanicsWithError(errString string, f assert.PanicTestFunc, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
PanicsWithError(a.t, errString, f, msgAndArgs...)
|
||||
}
|
||||
|
||||
// PanicsWithErrorf asserts that the code inside the specified PanicTestFunc
|
||||
// panics, and that the recovered panic value is an error that satisfies the
|
||||
// EqualError comparison.
|
||||
//
|
||||
// a.PanicsWithErrorf("crazy error", func(){ GoCrazy() }, "error message %s", "formatted")
|
||||
func (a *Assertions) PanicsWithErrorf(errString string, f assert.PanicTestFunc, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
PanicsWithErrorf(a.t, errString, f, msg, args...)
|
||||
}
|
||||
|
||||
// PanicsWithValue asserts that the code inside the specified PanicTestFunc panics, and that
|
||||
// the recovered panic value equals the expected panic value.
|
||||
//
|
||||
@@ -1104,6 +1200,22 @@ func (a *Assertions) WithinDurationf(expected time.Time, actual time.Time, delta
|
||||
WithinDurationf(a.t, expected, actual, delta, msg, args...)
|
||||
}
|
||||
|
||||
// YAMLEq asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEq(expected string, actual string, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
YAMLEq(a.t, expected, actual, msgAndArgs...)
|
||||
}
|
||||
|
||||
// YAMLEqf asserts that two YAML strings are equivalent.
|
||||
func (a *Assertions) YAMLEqf(expected string, actual string, msg string, args ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
h.Helper()
|
||||
}
|
||||
YAMLEqf(a.t, expected, actual, msg, args...)
|
||||
}
|
||||
|
||||
// Zero asserts that i is the zero value for its type.
|
||||
func (a *Assertions) Zero(i interface{}, msgAndArgs ...interface{}) {
|
||||
if h, ok := a.t.(tHelper); ok {
|
||||
|
||||
+1
-1
@@ -26,4 +26,4 @@ type BoolAssertionFunc func(TestingT, bool, ...interface{})
|
||||
// for table driven tests.
|
||||
type ErrorAssertionFunc func(TestingT, error, ...interface{})
|
||||
|
||||
//go:generate go run ../_codegen/main.go -output-package=require -template=require.go.tmpl -include-format-funcs
|
||||
//go:generate sh -c "cd ../_codegen && go build && cd - && ../_codegen/_codegen -output-package=require -template=require.go.tmpl -include-format-funcs"
|
||||
|
||||
+6
-1
@@ -7,6 +7,7 @@ import (
|
||||
"reflect"
|
||||
"regexp"
|
||||
"runtime/debug"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -80,11 +81,12 @@ func (suite *Suite) Run(name string, subtest func()) bool {
|
||||
// Run takes a testing suite and runs all of the tests attached
|
||||
// to it.
|
||||
func Run(t *testing.T, suite TestingSuite) {
|
||||
testsSync := &sync.WaitGroup{}
|
||||
suite.SetT(t)
|
||||
defer failOnPanic(t)
|
||||
|
||||
suiteSetupDone := false
|
||||
|
||||
|
||||
methodFinder := reflect.TypeOf(suite)
|
||||
tests := []testing.InternalTest{}
|
||||
for index := 0; index < methodFinder.NumMethod(); index++ {
|
||||
@@ -103,6 +105,7 @@ func Run(t *testing.T, suite TestingSuite) {
|
||||
}
|
||||
defer func() {
|
||||
if tearDownAllSuite, ok := suite.(TearDownAllSuite); ok {
|
||||
testsSync.Wait()
|
||||
tearDownAllSuite.TearDownSuite()
|
||||
}
|
||||
}()
|
||||
@@ -111,6 +114,7 @@ func Run(t *testing.T, suite TestingSuite) {
|
||||
test := testing.InternalTest{
|
||||
Name: method.Name,
|
||||
F: func(t *testing.T) {
|
||||
defer testsSync.Done()
|
||||
parentT := suite.T()
|
||||
suite.SetT(t)
|
||||
defer failOnPanic(t)
|
||||
@@ -134,6 +138,7 @@ func Run(t *testing.T, suite TestingSuite) {
|
||||
},
|
||||
}
|
||||
tests = append(tests, test)
|
||||
testsSync.Add(1)
|
||||
}
|
||||
runTests(t, tests)
|
||||
}
|
||||
|
||||
Vendored
+23
-1
@@ -393,6 +393,28 @@ github.com/kr/logfmt
|
||||
github.com/kr/pty
|
||||
# github.com/leodido/go-urn v1.1.0
|
||||
github.com/leodido/go-urn
|
||||
# github.com/lestrrat-go/iter v0.0.0-20200422075355-fc1769541911
|
||||
github.com/lestrrat-go/iter/arrayiter
|
||||
github.com/lestrrat-go/iter/mapiter
|
||||
# github.com/lestrrat-go/jwx v1.0.2
|
||||
github.com/lestrrat-go/jwx/internal/base64
|
||||
github.com/lestrrat-go/jwx/internal/iter
|
||||
github.com/lestrrat-go/jwx/internal/option
|
||||
github.com/lestrrat-go/jwx/internal/pool
|
||||
github.com/lestrrat-go/jwx/jwa
|
||||
github.com/lestrrat-go/jwx/jwk
|
||||
github.com/lestrrat-go/jwx/jws
|
||||
github.com/lestrrat-go/jwx/jws/sign
|
||||
github.com/lestrrat-go/jwx/jws/verify
|
||||
github.com/lestrrat-go/jwx/jwt
|
||||
github.com/lestrrat-go/jwx/jwt/internal/types
|
||||
github.com/lestrrat-go/jwx/jwt/openid
|
||||
# github.com/lestrrat/go-jwx v0.0.0-20180221005942-b7d4802280ae
|
||||
github.com/lestrrat/go-jwx/internal/base64
|
||||
github.com/lestrrat/go-jwx/jwa
|
||||
github.com/lestrrat/go-jwx/jwk
|
||||
# github.com/lestrrat/go-pdebug v0.0.0-20180220043741-569c97477ae8
|
||||
github.com/lestrrat/go-pdebug
|
||||
# github.com/libvirt/libvirt-go-xml v5.2.0+incompatible
|
||||
github.com/libvirt/libvirt-go-xml
|
||||
# github.com/ma314smith/signedxml v0.0.0-20200410192636-c342a2d0ae60
|
||||
@@ -528,7 +550,7 @@ github.com/spaolacci/murmur3
|
||||
github.com/spf13/pflag
|
||||
# github.com/stretchr/objx v0.2.0
|
||||
github.com/stretchr/objx
|
||||
# github.com/stretchr/testify v1.4.0
|
||||
# github.com/stretchr/testify v1.5.1
|
||||
github.com/stretchr/testify/assert
|
||||
github.com/stretchr/testify/mock
|
||||
github.com/stretchr/testify/require
|
||||
|
||||
Reference in New Issue
Block a user