@@ -11,21 +11,18 @@ package msteams
1111import (
1212 "bytes"
1313 "context"
14- "crypto/rsa"
15- "encoding/base64"
1614 "encoding/json"
1715 "errors"
1816 "fmt"
1917 "io"
20- "math/big"
2118 "net/http"
2219 "net/url"
2320 "strings"
2421 "sync"
2522 "time"
2623
27- "github.com/golang-jwt/jwt/v5"
2824 "github.com/memcode-ai/memcode/internal/channels"
25+ "github.com/memcode-ai/memcode/internal/webjwt"
2926)
3027
3128// botFrameworkIssuer is the issuer every Bot Framework connector token carries.
@@ -54,15 +51,9 @@ type Channel struct {
5451 appPassword string
5552 tenantID string
5653 mediaDir string // media spool; "" disables inbound media downloads
57- metadataURL string // Bot Framework OpenID metadata; overridable in tests
5854 tokenBase string // Azure AD token endpoint base; overridable in tests
5955 client * http.Client
60-
61- // keysMu guards the JWKS cache. Keys are fetched lazily and refreshed at
62- // most once per request when an unknown kid arrives (Microsoft rotates
63- // signing keys), so a flood of bad tokens can't hammer the metadata host.
64- keysMu sync.Mutex
65- keys map [string ]* rsa.PublicKey
56+ verify * webjwt.Verifier // the shared inbound-JWT verifier; tests point its MetadataURL at a fake
6657
6758 // tokMu guards the cached outbound bearer; refreshed ~60s before expiry so
6859 // an in-flight Send never races the token's edge.
@@ -75,14 +66,20 @@ type Channel struct {
7566// the AAD tenant the bot is registered in. mediaDir is the gateway media spool
7667// inbound attachments are downloaded into; "" disables media handling.
7768func New (appID , appPassword , tenantID , mediaDir string ) * Channel {
69+ client := & http.Client {Timeout : 30 * time .Second }
7870 return & Channel {
7971 appID : appID ,
8072 appPassword : appPassword ,
8173 tenantID : tenantID ,
8274 mediaDir : mediaDir ,
83- metadataURL : defaultMetadataURL ,
8475 tokenBase : defaultTokenBase ,
85- client : & http.Client {Timeout : 30 * time .Second },
76+ client : client ,
77+ verify : & webjwt.Verifier {
78+ MetadataURL : defaultMetadataURL ,
79+ Issuer : botFrameworkIssuer ,
80+ Audience : appID ,
81+ Client : client ,
82+ },
8683 }
8784}
8885
@@ -129,7 +126,7 @@ func (c *Channel) Handler(sink channels.Sink) http.Handler {
129126 return
130127 }
131128 raw , ok := strings .CutPrefix (r .Header .Get ("Authorization" ), "Bearer " )
132- if ! ok || c .validateJWT (r .Context (), raw ) != nil {
129+ if ! ok || c .verify . Verify (r .Context (), raw ) != nil {
133130 // Unauthenticated caller: nothing is delivered. 401, not 503 — a
134131 // forged request must not be invited to retry.
135132 http .Error (w , "invalid token" , http .StatusUnauthorized )
@@ -264,122 +261,6 @@ func (c *Channel) fetch(ctx context.Context, u, bearer string) (*http.Response,
264261 return c .client .Do (req )
265262}
266263
267- // validateJWT verifies an inbound Bot Framework token: RS256 signature against
268- // the published JWKS, the Bot Framework issuer, our app id as audience, and an
269- // unexpired lifetime. An unknown kid triggers at most ONE JWKS refresh for
270- // this request — key rotation is handled, a forged-kid flood is not amplified.
271- func (c * Channel ) validateJWT (ctx context.Context , raw string ) error {
272- refreshed := false
273- keyfunc := func (t * jwt.Token ) (any , error ) {
274- kid , _ := t .Header ["kid" ].(string )
275- if kid == "" {
276- return nil , errors .New ("token missing kid" )
277- }
278- if k := c .cachedKey (kid ); k != nil {
279- return k , nil
280- }
281- if ! refreshed {
282- refreshed = true
283- if err := c .refreshKeys (ctx ); err != nil {
284- return nil , err
285- }
286- if k := c .cachedKey (kid ); k != nil {
287- return k , nil
288- }
289- }
290- return nil , fmt .Errorf ("unknown signing key %q" , kid )
291- }
292- _ , err := jwt .Parse (raw , keyfunc ,
293- jwt .WithValidMethods ([]string {"RS256" }),
294- jwt .WithIssuer (botFrameworkIssuer ),
295- jwt .WithAudience (c .appID ),
296- jwt .WithExpirationRequired (),
297- )
298- return err
299- }
300-
301- func (c * Channel ) cachedKey (kid string ) * rsa.PublicKey {
302- c .keysMu .Lock ()
303- defer c .keysMu .Unlock ()
304- return c .keys [kid ]
305- }
306-
307- // refreshKeys fetches the OpenID metadata, follows jwks_uri, and replaces the
308- // key cache. Replacing (not merging) means revoked keys actually leave.
309- func (c * Channel ) refreshKeys (ctx context.Context ) error {
310- var meta struct {
311- JWKSURI string `json:"jwks_uri"`
312- }
313- if err := c .getJSON (ctx , c .metadataURL , & meta ); err != nil {
314- return fmt .Errorf ("openid metadata: %w" , err )
315- }
316- if meta .JWKSURI == "" {
317- return errors .New ("openid metadata has no jwks_uri" )
318- }
319- var set struct {
320- Keys []struct {
321- Kty string `json:"kty"`
322- Kid string `json:"kid"`
323- N string `json:"n"`
324- E string `json:"e"`
325- } `json:"keys"`
326- }
327- if err := c .getJSON (ctx , meta .JWKSURI , & set ); err != nil {
328- return fmt .Errorf ("jwks fetch: %w" , err )
329- }
330- keys := make (map [string ]* rsa.PublicKey , len (set .Keys ))
331- for _ , k := range set .Keys {
332- if k .Kty != "RSA" || k .Kid == "" {
333- continue
334- }
335- pub , err := rsaFromJWK (k .N , k .E )
336- if err != nil {
337- continue // one malformed key must not poison the whole set
338- }
339- keys [k .Kid ] = pub
340- }
341- if len (keys ) == 0 {
342- return errors .New ("jwks contained no usable rsa keys" )
343- }
344- c .keysMu .Lock ()
345- c .keys = keys
346- c .keysMu .Unlock ()
347- return nil
348- }
349-
350- // rsaFromJWK builds an RSA public key from base64url modulus and exponent.
351- func rsaFromJWK (n64 , e64 string ) (* rsa.PublicKey , error ) {
352- nb , err := base64 .RawURLEncoding .DecodeString (n64 )
353- if err != nil {
354- return nil , err
355- }
356- eb , err := base64 .RawURLEncoding .DecodeString (e64 )
357- if err != nil {
358- return nil , err
359- }
360- e := new (big.Int ).SetBytes (eb )
361- if ! e .IsInt64 () || e .Int64 () <= 0 {
362- return nil , errors .New ("bad rsa exponent" )
363- }
364- return & rsa.PublicKey {N : new (big.Int ).SetBytes (nb ), E : int (e .Int64 ())}, nil
365- }
366-
367- func (c * Channel ) getJSON (ctx context.Context , u string , v any ) error {
368- req , err := http .NewRequestWithContext (ctx , http .MethodGet , u , nil )
369- if err != nil {
370- return err
371- }
372- resp , err := c .client .Do (req )
373- if err != nil {
374- return err
375- }
376- defer resp .Body .Close ()
377- if resp .StatusCode / 100 != 2 {
378- return fmt .Errorf ("get %s: status %d" , u , resp .StatusCode )
379- }
380- return json .NewDecoder (io .LimitReader (resp .Body , maxBody )).Decode (v )
381- }
382-
383264// token returns a valid outbound connector bearer, minting one via the Azure
384265// AD client-credentials grant when the cache is empty or within 60s of expiry
385266// (the margin keeps a token from expiring mid-Send).
0 commit comments