update vendor

This commit is contained in:
Qiu Jian
2020-06-29 14:41:49 +08:00
parent 2183490a53
commit 182c96a160
99 changed files with 11616 additions and 146 deletions
+4 -1
View File
@@ -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
+12
View File
@@ -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=
+6 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+50
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1 @@
package jws
+76
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 [ISO86012004] 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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1 @@
package jwt
+387
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1,99 @@
# go-pdebug
[![Build Status](https://travis-ci.org/lestrrat/go-pdebug.svg?branch=master)](https://travis-ci.org/lestrrat/go-pdebug)
[![GoDoc](https://godoc.org/github.com/lestrrat/go-pdebug?status.svg)](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
![optimized](https://pbs.twimg.com/media/CbiqhzLUUAIN_7o.png)
# 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
View File
@@ -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
View File
@@ -0,0 +1,6 @@
// +build debug0
package pdebug
var Trace = true
+43
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
}
+23 -1
View File
@@ -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