diff --git a/packages/shorebird_cli/lib/src/auth/auth.dart b/packages/shorebird_cli/lib/src/auth/auth.dart index 45120398..379b6e7e 100644 --- a/packages/shorebird_cli/lib/src/auth/auth.dart +++ b/packages/shorebird_cli/lib/src/auth/auth.dart @@ -1,7 +1,7 @@ import 'dart:convert'; import 'dart:io'; -import 'package:googleapis_auth/auth_io.dart'; +import 'package:googleapis_auth/auth_io.dart' as oauth2; import 'package:http/http.dart' as http; import 'package:path/path.dart' as p; import 'package:shorebird_cli/src/auth/jwt.dart'; @@ -10,7 +10,7 @@ import 'package:shorebird_cli/src/config/config.dart'; export 'package:shorebird_cli/src/auth/models/models.dart' show User; -final _clientId = ClientId( +final _clientId = oauth2.ClientId( /// Shorebird CLI's OAuth 2.0 identifier. '523302233293-eia5antm0tgvek240t46orctktiabrek.apps.googleusercontent.com', @@ -28,22 +28,50 @@ final _clientId = ClientId( ); final _scopes = ['openid', 'https://www.googleapis.com/auth/userinfo.email']; -typedef ObtainAccessCredentials = Future Function( - ClientId clientId, +typedef ObtainAccessCredentials = Future Function( + oauth2.ClientId clientId, List scopes, http.Client client, void Function(String) userPrompt, ); +typedef RefreshCredentials = Future Function( + oauth2.ClientId clientId, + oauth2.AccessCredentials credentials, + http.Client client, +); + +typedef OnRefreshCredentials = void Function( + oauth2.AccessCredentials credentials, +); + class AuthenticatedClient extends http.BaseClient { - AuthenticatedClient({required this.token, required http.Client httpClient}) - : _baseClient = httpClient; + AuthenticatedClient({ + required oauth2.AccessCredentials credentials, + required http.Client httpClient, + required OnRefreshCredentials onRefreshCredentials, + RefreshCredentials refreshCredentials = oauth2.refreshCredentials, + }) : _credentials = credentials, + _baseClient = httpClient, + _onRefreshCredentials = onRefreshCredentials, + _refreshCredentials = refreshCredentials; final http.Client _baseClient; - final String token; + final OnRefreshCredentials _onRefreshCredentials; + final RefreshCredentials _refreshCredentials; + oauth2.AccessCredentials _credentials; @override - Future send(http.BaseRequest request) { + Future send(http.BaseRequest request) async { + if (_credentials.accessToken.hasExpired) { + _credentials = await _refreshCredentials( + _clientId, + _credentials, + _baseClient, + ); + _onRefreshCredentials(_credentials); + } + final token = _credentials.idToken; request.headers['Authorization'] = 'Bearer $token'; return _baseClient.send(request); } @@ -54,8 +82,8 @@ class Auth { http.Client? httpClient, ObtainAccessCredentials? obtainAccessCredentials, }) : _httpClient = httpClient ?? http.Client(), - _obtainAccessCredentials = - obtainAccessCredentials ?? obtainAccessCredentialsViaUserConsent { + _obtainAccessCredentials = obtainAccessCredentials ?? + oauth2.obtainAccessCredentialsViaUserConsent { _loadCredentials(); } @@ -66,9 +94,13 @@ class Auth { final credentialsFilePath = p.join(shorebirdConfigDir, _credentialsFileName); http.Client get client { - final token = _credentials?.idToken; - if (token == null) return _httpClient; - return AuthenticatedClient(token: token, httpClient: _httpClient); + final credentials = _credentials; + if (credentials == null) return _httpClient; + return AuthenticatedClient( + credentials: credentials, + httpClient: _httpClient, + onRefreshCredentials: _flushCredentials, + ); } Future login(void Function(String) prompt) async { @@ -91,7 +123,7 @@ class Auth { void logout() => _clearCredentials(); - AccessCredentials? _credentials; + oauth2.AccessCredentials? _credentials; User? _user; @@ -105,7 +137,7 @@ class Auth { if (credentialsFile.existsSync()) { try { final contents = credentialsFile.readAsStringSync(); - _credentials = AccessCredentials.fromJson( + _credentials = oauth2.AccessCredentials.fromJson( json.decode(contents) as Map, ); _user = _credentials?.toUser(); @@ -113,7 +145,7 @@ class Auth { } } - void _flushCredentials(AccessCredentials credentials) { + void _flushCredentials(oauth2.AccessCredentials credentials) { File(credentialsFilePath) ..createSync(recursive: true) ..writeAsStringSync(json.encode(credentials.toJson())); @@ -134,7 +166,7 @@ class Auth { } } -extension on AccessCredentials { +extension on oauth2.AccessCredentials { User toUser() { final token = idToken; diff --git a/packages/shorebird_cli/test/src/auth/auth_test.dart b/packages/shorebird_cli/test/src/auth/auth_test.dart index bc176d3e..602c5ee9 100644 --- a/packages/shorebird_cli/test/src/auth/auth_test.dart +++ b/packages/shorebird_cli/test/src/auth/auth_test.dart @@ -17,8 +17,12 @@ void main() { const idToken = '''eyJhbGciOiJSUzI1NiIsImN0eSI6IkpXVCJ9.eyJlbWFpbCI6InRlc3RAZW1haWwuY29tIn0.pD47BhF3MBLyIpfsgWCzP9twzC1HJxGukpcR36DqT6yfiOMHTLcjDbCjRLAnklWEHiT0BQTKTfhs8IousU90Fm5bVKObudfKu8pP5iZZ6Ls4ohDjTrXky9j3eZpZjwv8CnttBVgRfMJG-7YASTFRYFcOLUpnb4Zm5R6QdoCDUYg'''; const email = 'test@email.com'; - final credentials = AccessCredentials( - AccessToken('Bearer', 'accessToken', DateTime.now().toUtc()), + final validCredentials = AccessCredentials( + AccessToken( + 'Bearer', + 'accessToken', + DateTime.now().add(const Duration(minutes: 10)).toUtc(), + ), '', [], idToken: idToken, @@ -38,11 +42,79 @@ void main() { auth = Auth( httpClient: httpClient, obtainAccessCredentials: (clientId, scopes, client, userPrompt) async { - return credentials; + return validCredentials; }, )..logout(); }); + group('AuthenticatedClient', () { + test('refreshes and uses new token when credentials are expired.', + () async { + when(() => httpClient.send(any())).thenAnswer( + (_) async => http.StreamedResponse( + const Stream.empty(), + HttpStatus.ok, + ), + ); + + final onRefreshCredentialsCalls = []; + final expiredCredentials = AccessCredentials( + AccessToken( + 'Bearer', + 'accessToken', + DateTime.now().subtract(const Duration(minutes: 1)).toUtc(), + ), + '', + [], + idToken: 'expiredIdToken', + ); + + final client = AuthenticatedClient( + credentials: expiredCredentials, + httpClient: httpClient, + onRefreshCredentials: onRefreshCredentialsCalls.add, + refreshCredentials: (clientId, credentials, client) async => + validCredentials, + ); + + await client.get(Uri.parse('https://example.com')); + + expect( + onRefreshCredentialsCalls, + equals([ + isA().having((c) => c.idToken, 'token', idToken) + ]), + ); + final captured = verify(() => httpClient.send(captureAny())).captured; + expect(captured, hasLength(1)); + final request = captured.first as http.BaseRequest; + expect(request.headers['Authorization'], equals('Bearer $idToken')); + }); + + test('uses valid token when credentials valid.', () async { + when(() => httpClient.send(any())).thenAnswer( + (_) async => http.StreamedResponse( + const Stream.empty(), + HttpStatus.ok, + ), + ); + final onRefreshCredentialsCalls = []; + final client = AuthenticatedClient( + credentials: validCredentials, + httpClient: httpClient, + onRefreshCredentials: onRefreshCredentialsCalls.add, + ); + + await client.get(Uri.parse('https://example.com')); + + expect(onRefreshCredentialsCalls, isEmpty); + final captured = verify(() => httpClient.send(captureAny())).captured; + expect(captured, hasLength(1)); + final request = captured.first as http.BaseRequest; + expect(request.headers['Authorization'], equals('Bearer $idToken')); + }); + }); + group('client', () { test( 'returns an authenticated client '