diff --git a/workout-logger/lib/main.dart b/workout-logger/lib/main.dart index f1b6cd1..91aa03a 100644 --- a/workout-logger/lib/main.dart +++ b/workout-logger/lib/main.dart @@ -9,7 +9,8 @@ import 'package:provider/provider.dart'; import 'services/storage_service.dart'; import 'services/ml_service.dart'; -import 'services/gemini_service.dart'; +import 'services/ai/gemini_ai_service.dart'; +import 'services/ai/coach_tool_service.dart'; import 'services/health_connect_service.dart'; import 'services/interfaces/storage_service_interface.dart'; import 'services/interfaces/ml_service_interface.dart'; @@ -21,6 +22,7 @@ import 'services/managers/program_manager.dart'; import 'services/managers/history_manager.dart'; import 'services/managers/health_sync_manager.dart'; import 'services/managers/pr_manager.dart'; +import 'services/managers/conversation_manager.dart'; import 'theme/app_theme.dart'; import 'screens/home_screen.dart'; import 'screens/onboarding_screen.dart'; @@ -62,7 +64,10 @@ class WorkoutLoggerApp extends StatelessWidget { static final HistoryManager _historyManager = HistoryManager(_storageService, healthSyncManager: _healthSyncManager); static final PRManager _prManager = PRManager(_storageService); - static final GeminiService _geminiService = GeminiService(); + static final GeminiAiService _geminiService = + GeminiAiService(storage: _storageService); + static final ConversationManager _conversationManager = + ConversationManager(_storageService); const WorkoutLoggerApp({super.key}); @@ -89,7 +94,15 @@ class WorkoutLoggerApp extends StatelessWidget { // Provided as ChangeNotifier so HistoryScreen rebuilds on sync badge changes. ChangeNotifierProvider.value(value: _historyManager), ChangeNotifierProvider.value(value: _prManager), - ChangeNotifierProvider.value(value: _geminiService), + // GeminiAiService is the single AI backend instance. It's a ChangeNotifier + // (settings UI watches isConfigured/model), so it's provided as such. + // Consumers that should depend on the abstraction (the coach ViewModel, + // program generator) receive it typed as IAiService at construction — + // the future firebase_ai swap point — without a separate provider. + ChangeNotifierProvider.value(value: _geminiService), + ChangeNotifierProvider.value( + value: _conversationManager, + ), // WorkoutProvider receives dependencies via constructor injection ChangeNotifierProvider( create: (_) => WorkoutProvider( @@ -99,6 +112,13 @@ class WorkoutLoggerApp extends StatelessWidget { programManager: _programManager, ), ), + // CoachToolService backs AI tool calls; reads from WorkoutProvider + PRManager. + Provider( + create: (ctx) => CoachToolService( + ctx.read(), + ctx.read(), + ), + ), ], child: MaterialApp( title: 'Workout Logger', @@ -135,12 +155,17 @@ class _AppInitializerState extends State { final historyManager = context.read(); final prManager = context.read(); final api = context.read(); - final gemini = context.read(); + final gemini = context.read(); try { await provider.init(); await settings.init(); gemini.init(settings.geminiApiKey, model: settings.geminiModel); + try { + await gemini.loadUsage(); + } catch (e, st) { + debugPrint('gemini.loadUsage failed: $e\n$st'); + } await historyManager.loadSessions(); await prManager.load(); await prManager.backfillFromSessions(historyManager.sessions); diff --git a/workout-logger/lib/models/models.dart b/workout-logger/lib/models/models.dart index 2a5b643..ecf7551 100644 --- a/workout-logger/lib/models/models.dart +++ b/workout-logger/lib/models/models.dart @@ -812,3 +812,101 @@ class TrainingProgram { createdAt == _sentinel ? this.createdAt : createdAt as DateTime, ); } + +// ==================== AI Coach Chat ==================== + +/// A single message in an AI coach conversation. +class ChatMessage { + final String id; + final String role; // 'user' | 'model' + final String text; + final DateTime timestamp; + + ChatMessage({ + String? id, + required this.role, + required this.text, + DateTime? timestamp, + }) : id = id ?? _uuid.v4(), + timestamp = timestamp ?? DateTime.now(); + + Map toJson() => { + 'id': id, + 'role': role, + 'text': text, + 'timestamp': timestamp.toIso8601String(), + }; + + factory ChatMessage.fromJson(Map json) => ChatMessage( + id: json['id'] as String?, + role: json['role'] as String, + text: json['text'] as String, + timestamp: DateTime.parse(json['timestamp'] as String), + ); + + ChatMessage copyWith({ + Object? role = _sentinel, + Object? text = _sentinel, + Object? timestamp = _sentinel, + }) => ChatMessage( + id: id, + role: role == _sentinel ? this.role : role as String, + text: text == _sentinel ? this.text : text as String, + timestamp: timestamp == _sentinel ? this.timestamp : timestamp as DateTime, + ); +} + +/// A persisted AI coach conversation: an ordered list of [ChatMessage]s. +class Conversation { + final String id; + final String title; + final DateTime createdAt; + final DateTime updatedAt; + final List messages; + + Conversation({ + String? id, + required this.title, + DateTime? createdAt, + DateTime? updatedAt, + List? messages, + }) : id = id ?? _uuid.v4(), + createdAt = createdAt ?? DateTime.now(), + updatedAt = updatedAt ?? createdAt ?? DateTime.now(), + messages = messages ?? const []; + + Map toJson() => { + 'id': id, + 'title': title, + 'createdAt': createdAt.toIso8601String(), + 'updatedAt': updatedAt.toIso8601String(), + 'messages': messages.map((m) => m.toJson()).toList(), + }; + + factory Conversation.fromJson(Map json) => Conversation( + id: json['id'] as String?, + title: json['title'] as String, + createdAt: DateTime.parse(json['createdAt'] as String), + updatedAt: json['updatedAt'] != null + ? DateTime.parse(json['updatedAt'] as String) + : DateTime.parse(json['createdAt'] as String), + messages: (json['messages'] as List) + .map((m) => ChatMessage.fromJson(m as Map)) + .toList(), + ); + + Conversation copyWith({ + Object? title = _sentinel, + Object? createdAt = _sentinel, + Object? updatedAt = _sentinel, + Object? messages = _sentinel, + }) => Conversation( + id: id, + title: title == _sentinel ? this.title : title as String, + createdAt: createdAt == _sentinel ? this.createdAt : createdAt as DateTime, + updatedAt: updatedAt == _sentinel ? this.updatedAt : updatedAt as DateTime, + messages: messages == _sentinel + ? this.messages + : messages as List, + ); +} diff --git a/workout-logger/lib/screens/ai_coach_screen.dart b/workout-logger/lib/screens/ai_coach_screen.dart index 9a57341..79f8145 100644 --- a/workout-logger/lib/screens/ai_coach_screen.dart +++ b/workout-logger/lib/screens/ai_coach_screen.dart @@ -1,48 +1,58 @@ -// ai_coach_screen.dart — Conversational AI workout coach powered by Gemini +// ai_coach_screen.dart — Conversational AI workout coach (View). +// +// This is a lean View: all orchestration (streaming, tool calls, persistence, +// system-prompt building) lives in AiCoachViewModel. The widget only renders +// state, forwards user intents, and holds UI-local controllers. import 'package:flutter/material.dart'; import 'package:flutter/services.dart'; -import 'package:google_generative_ai/google_generative_ai.dart'; import 'package:provider/provider.dart'; import 'package:google_fonts/google_fonts.dart'; +import 'package:gpt_markdown/gpt_markdown.dart'; -import '../services/gemini_service.dart'; -import '../services/gemini_context_builder.dart'; -import '../services/workout_provider.dart'; +import '../models/models.dart'; +import '../viewmodels/ai_coach_view_model.dart'; +import '../services/ai/gemini_ai_service.dart'; +import '../services/ai/coach_tool_service.dart'; +import '../services/managers/conversation_manager.dart'; import '../services/settings_provider.dart'; -import '../services/interfaces/ml_service_interface.dart'; import '../theme/app_theme.dart'; import 'widgets/rf_widgets.dart'; import 'profile_screen.dart'; -// ── Data ────────────────────────────────────────────────────────────────────── - -class _ChatMessage { - const _ChatMessage({required this.role, required this.text}); - final String role; // 'user' | 'model' - final String text; -} - -// ── Screen ──────────────────────────────────────────────────────────────────── - -class AiCoachScreen extends StatefulWidget { +/// Public entry point. Owns the screen-scoped [AiCoachViewModel]. +class AiCoachScreen extends StatelessWidget { const AiCoachScreen({super.key, this.seedPrompt}); /// Optional question to auto-send on open (e.g. deep-linked from analytics). - /// The coach system prompt already carries the user's data, so a seed needs - /// no extra context. final String? seedPrompt; @override - State createState() => _AiCoachScreenState(); + Widget build(BuildContext context) { + return ChangeNotifierProvider( + create: (ctx) => AiCoachViewModel( + ai: ctx.read(), + coachTools: ctx.read(), + conversations: ctx.read(), + settings: ctx.read(), + )..loadConversations(), + child: _AiCoachView(seedPrompt: seedPrompt), + ); + } } -class _AiCoachScreenState extends State { +class _AiCoachView extends StatefulWidget { + const _AiCoachView({this.seedPrompt}); + final String? seedPrompt; + + @override + State<_AiCoachView> createState() => _AiCoachViewState(); +} + +class _AiCoachViewState extends State<_AiCoachView> { final _controller = TextEditingController(); final _scrollCtrl = ScrollController(); - final _messages = <_ChatMessage>[]; - bool _loading = false; - String _streamingText = ''; + AiCoachViewModel? _vm; @override void initState() { @@ -51,92 +61,41 @@ class _AiCoachScreenState extends State { if (seed != null && seed.isNotEmpty) { WidgetsBinding.instance.addPostFrameCallback((_) { if (!mounted) return; - if (!context.read().isConfigured) return; - _controller.text = seed; - _send(); + final vm = context.read(); + if (!vm.isConfigured) return; + _controller.clear(); + vm.sendMessage(seed); }); } } + @override + void didChangeDependencies() { + super.didChangeDependencies(); + // Attach a scroll-follow listener once. + final vm = context.read(); + if (!identical(vm, _vm)) { + _vm?.removeListener(_onVmChanged); + _vm = vm..addListener(_onVmChanged); + } + } + + void _onVmChanged() => _scrollToBottom(); + @override void dispose() { + _vm?.removeListener(_onVmChanged); _controller.dispose(); _scrollCtrl.dispose(); super.dispose(); } - String _buildSystemPrompt() { - final wp = context.read(); - final settings = context.read(); - final mlService = context.read(); - - final exerciseMap = {for (final e in wp.allExercises) e.id: e}; - final allSessions = wp.sessions; - final recoveryScores = mlService.computeMuscleRecoveryScores( - allSessions, - exerciseMap, - ); - final activeTargets = wp.targets.where((t) => !t.isCompleted).toList(); - - return GeminiContextBuilder.buildCoachSystemPrompt( - recentSessions: allSessions, - exerciseMap: exerciseMap, - recoveryScores: recoveryScores, - activeTargets: activeTargets, - userName: settings.userName, - unitLabel: settings.unitLabel, - ); - } - - List _buildHistory() => _messages - .map((m) => Content(m.role, [TextPart(m.text)])) - .toList(); - - Future _send() async { + void _send() { final text = _controller.text.trim(); - if (text.isEmpty || _loading) return; - + if (text.isEmpty) return; HapticFeedback.lightImpact(); _controller.clear(); - - setState(() { - _messages.add(_ChatMessage(role: 'user', text: text)); - _loading = true; - _streamingText = ''; - }); - _scrollToBottom(); - - final gemini = context.read(); - final systemPrompt = _buildSystemPrompt(); - // Build history from all messages except the one we just added. - final history = _messages.length > 1 - ? _buildHistory().sublist(0, _messages.length - 1) - : []; - - final buffer = StringBuffer(); - try { - await for (final chunk in gemini.streamCoachReply( - userMessage: text, - systemPrompt: systemPrompt, - history: history, - )) { - buffer.write(chunk); - if (mounted) { - setState(() => _streamingText = buffer.toString()); - _scrollToBottom(); - } - } - if (mounted) { - setState(() { - _messages.add(_ChatMessage(role: 'model', text: buffer.toString())); - _streamingText = ''; - _loading = false; - }); - _scrollToBottom(); - } - } catch (_) { - if (mounted) setState(() { _streamingText = ''; _loading = false; }); - } + context.read().sendMessage(text); } void _scrollToBottom() { @@ -153,7 +112,7 @@ class _AiCoachScreenState extends State { @override Widget build(BuildContext context) { - final gemini = context.watch(); + final vm = context.watch(); return Scaffold( backgroundColor: AppColors.background, @@ -163,13 +122,13 @@ class _AiCoachScreenState extends State { SafeArea( child: Column( children: [ - _buildHeader(context), + _buildHeader(context, vm), Expanded( - child: gemini.isConfigured - ? _buildChatArea() + child: vm.isConfigured + ? _buildChatArea(vm) : _buildNoKeyState(context), ), - if (gemini.isConfigured) _buildInputBar(), + if (vm.isConfigured) _buildInputBar(vm), ], ), ), @@ -178,7 +137,7 @@ class _AiCoachScreenState extends State { ); } - Widget _buildHeader(BuildContext context) { + Widget _buildHeader(BuildContext context, AiCoachViewModel vm) { return Padding( padding: const EdgeInsets.fromLTRB( AppSpacing.md, @@ -248,34 +207,42 @@ class _AiCoachScreenState extends State { ], ), ), - if (_messages.isNotEmpty) - GestureDetector( - onTap: () => setState(() => _messages.clear()), - child: Container( - padding: const EdgeInsets.symmetric(horizontal: 10, vertical: 5), - decoration: BoxDecoration( - color: AppColors.glass, - borderRadius: BorderRadius.circular(AppRadius.full), - border: Border.all(color: AppColors.glassBorder), - ), - child: Text( - 'Clear', - style: GoogleFonts.geist( - color: AppColors.textMuted, - fontSize: 11, - ), - ), - ), + if (vm.isConfigured) ...[ + _HeaderIconButton( + icon: Icons.history_rounded, + onTap: () => _openHistory(context, vm), ), + const SizedBox(width: AppSpacing.sm), + _HeaderIconButton( + icon: Icons.add_rounded, + onTap: () { + HapticFeedback.lightImpact(); + vm.newConversation(); + }, + ), + ], ], ), ); } - Widget _buildChatArea() { - final hasMessages = _messages.isNotEmpty || _loading; + Future _openHistory(BuildContext context, AiCoachViewModel vm) async { + HapticFeedback.lightImpact(); + await showModalBottomSheet( + context: context, + backgroundColor: AppColors.surface, + isScrollControlled: true, + shape: const RoundedRectangleBorder( + borderRadius: BorderRadius.vertical(top: Radius.circular(AppRadius.lg)), + ), + builder: (_) => _ConversationsSheet(vm: vm), + ); + } - if (!hasMessages) return _buildWelcome(); + Widget _buildChatArea(AiCoachViewModel vm) { + final messages = vm.messages; + final hasContent = messages.isNotEmpty || vm.isLoading; + if (!hasContent) return _buildWelcome(); return ListView.builder( controller: _scrollCtrl, @@ -285,20 +252,18 @@ class _AiCoachScreenState extends State { AppSpacing.md, AppSpacing.sm, ), - itemCount: _messages.length + (_loading ? 1 : 0), + itemCount: messages.length + (vm.isLoading ? 1 : 0), itemBuilder: (_, i) { - if (i == _messages.length) { - // Streaming bubble - return _StreamingBubble(text: _streamingText); + if (i == messages.length) { + return _StreamingBubble(text: vm.streamingText); } - return _MessageBubble(message: _messages[i]); + return _MessageBubble(message: messages[i]); }, ); } Widget _buildWelcome() { - final settings = context.read(); - final name = settings.userName; + final name = context.read().userName; return Center( child: Padding( padding: const EdgeInsets.all(AppSpacing.xl), @@ -331,9 +296,7 @@ class _AiCoachScreenState extends State { ), const SizedBox(height: AppSpacing.lg), Text( - name != null && name.isNotEmpty - ? 'Hey $name 👋' - : 'Your AI Coach', + name != null && name.isNotEmpty ? 'Hey $name 👋' : 'Your AI Coach', style: GoogleFonts.geist( color: AppColors.textPrimary, fontSize: 22, @@ -356,11 +319,20 @@ class _AiCoachScreenState extends State { spacing: AppSpacing.sm, runSpacing: AppSpacing.sm, alignment: WrapAlignment.center, - children: const [ - _SuggestionChip('What should I train today?'), - _SuggestionChip('How\'s my recovery?'), - _SuggestionChip('Am I progressing on bench?'), - _SuggestionChip('Suggest a deload week'), + children: [ + for (final s in const [ + 'What should I train today?', + 'How\'s my recovery?', + 'Am I progressing on bench?', + 'Suggest a deload week', + ]) + _SuggestionChip( + label: s, + onTap: () { + _controller.text = s; + _send(); + }, + ), ], ), ], @@ -397,7 +369,8 @@ class _AiCoachScreenState extends State { ); } - Widget _buildInputBar() { + Widget _buildInputBar(AiCoachViewModel vm) { + final loading = vm.isLoading; return Container( padding: EdgeInsets.fromLTRB( AppSpacing.md, @@ -445,22 +418,22 @@ class _AiCoachScreenState extends State { ), const SizedBox(width: AppSpacing.sm), GestureDetector( - onTap: _loading ? null : _send, + onTap: loading ? null : _send, child: AnimatedContainer( duration: const Duration(milliseconds: 150), width: 44, height: 44, decoration: BoxDecoration( - gradient: _loading + gradient: loading ? null : const LinearGradient( colors: [AppColors.primary, Color(0xFF5B21B6)], begin: Alignment.topLeft, end: Alignment.bottomRight, ), - color: _loading ? AppColors.glass3 : null, + color: loading ? AppColors.glass3 : null, borderRadius: BorderRadius.circular(AppRadius.xl), - boxShadow: _loading + boxShadow: loading ? null : [ BoxShadow( @@ -470,7 +443,7 @@ class _AiCoachScreenState extends State { ), ], ), - child: _loading + child: loading ? const Center( child: SizedBox( width: 18, @@ -494,21 +467,202 @@ class _AiCoachScreenState extends State { } } +// ── Header icon button ────────────────────────────────────────────────────── + +class _HeaderIconButton extends StatelessWidget { + const _HeaderIconButton({required this.icon, required this.onTap}); + final IconData icon; + final VoidCallback onTap; + + @override + Widget build(BuildContext context) { + return GestureDetector( + onTap: onTap, + child: Container( + padding: const EdgeInsets.all(8), + decoration: BoxDecoration( + color: AppColors.glass, + borderRadius: BorderRadius.circular(AppRadius.sm), + border: Border.all(color: AppColors.glassBorder), + ), + child: Icon(icon, color: AppColors.textSoft, size: 18), + ), + ); + } +} + +// ── Conversations history sheet ─────────────────────────────────────────────── + +class _ConversationsSheet extends StatelessWidget { + const _ConversationsSheet({required this.vm}); + final AiCoachViewModel vm; + + @override + Widget build(BuildContext context) { + // Rebuild when the conversation list changes (delete, new message). + return AnimatedBuilder( + animation: vm, + builder: (context, _) { + final conversations = vm.conversations; + return SafeArea( + child: Padding( + padding: const EdgeInsets.all(AppSpacing.md), + child: Column( + mainAxisSize: MainAxisSize.min, + crossAxisAlignment: CrossAxisAlignment.start, + children: [ + Row( + children: [ + Text( + 'Conversations', + style: GoogleFonts.geist( + color: AppColors.textPrimary, + fontSize: 16, + fontWeight: FontWeight.w700, + ), + ), + const Spacer(), + GestureDetector( + onTap: () { + vm.newConversation(); + Navigator.pop(context); + }, + child: Row( + children: [ + const Icon(Icons.add_rounded, + color: AppColors.primary, size: 18), + const SizedBox(width: 4), + Text( + 'New chat', + style: GoogleFonts.geist( + color: AppColors.primary, + fontSize: 13, + fontWeight: FontWeight.w600, + ), + ), + ], + ), + ), + ], + ), + const SizedBox(height: AppSpacing.md), + if (conversations.isEmpty) + Padding( + padding: const EdgeInsets.symmetric(vertical: AppSpacing.lg), + child: Text( + 'No saved conversations yet.', + style: GoogleFonts.geist( + color: AppColors.textMuted, + fontSize: 13, + ), + ), + ) + else + ConstrainedBox( + constraints: BoxConstraints( + maxHeight: MediaQuery.of(context).size.height * 0.5, + ), + child: ListView.separated( + shrinkWrap: true, + itemCount: conversations.length, + separatorBuilder: (_, __) => + const SizedBox(height: AppSpacing.sm), + itemBuilder: (_, i) { + final c = conversations[i]; + final isActive = c.id == vm.activeConversationId; + return _ConversationTile( + conversation: c, + isActive: isActive, + onTap: () { + vm.selectConversation(c.id); + Navigator.pop(context); + }, + onDelete: () => vm.deleteConversation(c.id), + ); + }, + ), + ), + ], + ), + ), + ); + }, + ); + } +} + +class _ConversationTile extends StatelessWidget { + const _ConversationTile({ + required this.conversation, + required this.isActive, + required this.onTap, + required this.onDelete, + }); + + final Conversation conversation; + final bool isActive; + final VoidCallback onTap; + final VoidCallback onDelete; + + @override + Widget build(BuildContext context) { + return GestureDetector( + onTap: onTap, + child: Container( + padding: const EdgeInsets.symmetric( + horizontal: AppSpacing.md, + vertical: AppSpacing.sm + 2, + ), + decoration: BoxDecoration( + color: isActive ? AppColors.primary.withValues(alpha: 0.12) : AppColors.glass3, + borderRadius: BorderRadius.circular(AppRadius.md), + border: Border.all( + color: isActive ? AppColors.primary.withValues(alpha: 0.4) : AppColors.glassBorder, + ), + ), + child: Row( + children: [ + const Icon(Icons.chat_bubble_outline_rounded, + color: AppColors.textMuted, size: 16), + const SizedBox(width: AppSpacing.sm), + Expanded( + child: Text( + conversation.title.isEmpty ? 'New chat' : conversation.title, + maxLines: 1, + overflow: TextOverflow.ellipsis, + style: GoogleFonts.geist( + color: AppColors.textPrimary, + fontSize: 13, + fontWeight: FontWeight.w500, + ), + ), + ), + GestureDetector( + onTap: onDelete, + child: const Padding( + padding: EdgeInsets.only(left: AppSpacing.sm), + child: Icon(Icons.delete_outline_rounded, + color: AppColors.textFaint, size: 18), + ), + ), + ], + ), + ), + ); + } +} + // ── Suggestion chip ─────────────────────────────────────────────────────────── class _SuggestionChip extends StatelessWidget { - const _SuggestionChip(this.label); + const _SuggestionChip({required this.label, required this.onTap}); final String label; + final VoidCallback onTap; @override Widget build(BuildContext context) { return GestureDetector( - onTap: () { - final state = context.findAncestorStateOfType<_AiCoachScreenState>(); - if (state == null) return; - state._controller.text = label; - state._send(); - }, + onTap: onTap, child: Container( padding: const EdgeInsets.symmetric(horizontal: 14, vertical: 8), decoration: BoxDecoration( @@ -533,7 +687,7 @@ class _SuggestionChip extends StatelessWidget { class _MessageBubble extends StatelessWidget { const _MessageBubble({required this.message}); - final _ChatMessage message; + final ChatMessage message; @override Widget build(BuildContext context) { @@ -583,14 +737,16 @@ class _MessageBubble extends StatelessWidget { ] : null, ), - child: Text( - message.text, - style: GoogleFonts.geist( - color: AppColors.textPrimary, - fontSize: 14, - height: 1.55, - ), - ), + child: isUser + ? Text( + message.text, + style: GoogleFonts.geist( + color: AppColors.textPrimary, + fontSize: 14, + height: 1.55, + ), + ) + : _CoachMarkdown(text: message.text), ), ), ], @@ -630,14 +786,7 @@ class _StreamingBubble extends StatelessWidget { ), child: text.isEmpty ? const RFLoadingDots() - : Text( - text, - style: GoogleFonts.geist( - color: AppColors.textPrimary, - fontSize: 14, - height: 1.55, - ), - ), + : _CoachMarkdown(text: text), ), ), ], @@ -646,6 +795,24 @@ class _StreamingBubble extends StatelessWidget { } } +/// Markdown renderer for coach replies, styled to the app theme. +class _CoachMarkdown extends StatelessWidget { + const _CoachMarkdown({required this.text}); + final String text; + + @override + Widget build(BuildContext context) { + return GptMarkdown( + text, + style: GoogleFonts.geist( + color: AppColors.textPrimary, + fontSize: 14, + height: 1.55, + ), + ); + } +} + class _AiAvatar extends StatelessWidget { @override Widget build(BuildContext context) { diff --git a/workout-logger/lib/screens/ai_program_generator_screen.dart b/workout-logger/lib/screens/ai_program_generator_screen.dart index da7cc5e..ded3dc5 100644 --- a/workout-logger/lib/screens/ai_program_generator_screen.dart +++ b/workout-logger/lib/screens/ai_program_generator_screen.dart @@ -6,7 +6,7 @@ import 'package:provider/provider.dart'; import 'package:google_fonts/google_fonts.dart'; import '../models/models.dart'; -import '../services/gemini_service.dart'; +import '../services/ai/gemini_ai_service.dart'; import '../services/workout_provider.dart'; import '../services/managers/program_manager.dart'; import '../theme/app_theme.dart'; @@ -45,7 +45,7 @@ class _AiProgramGeneratorScreenState extends State { final prompt = _promptCtrl.text.trim(); if (prompt.isEmpty) return; - final gemini = context.read(); + final gemini = context.read(); if (!gemini.isConfigured) { setState(() { _error = 'Add your Gemini API key in Profile → AI Features first.'; }); return; diff --git a/workout-logger/lib/screens/home_screen.dart b/workout-logger/lib/screens/home_screen.dart index e57ab32..b55a3e9 100644 --- a/workout-logger/lib/screens/home_screen.dart +++ b/workout-logger/lib/screens/home_screen.dart @@ -9,7 +9,7 @@ import 'package:google_fonts/google_fonts.dart'; import '../models/models.dart'; import '../services/workout_provider.dart'; import '../services/settings_provider.dart'; -import '../services/gemini_service.dart'; +import '../services/ai/gemini_ai_service.dart'; import '../services/gemini_context_builder.dart'; import '../theme/app_theme.dart'; import 'workout_flow_screen.dart'; @@ -1207,7 +1207,7 @@ class _WeeklyInsightsCardState extends State<_WeeklyInsightsCard> { bool _loading = false; Future _refresh() async { - final gemini = context.read(); + final gemini = context.read(); if (!gemini.isConfigured) return; setState(() => _loading = true); @@ -1250,7 +1250,7 @@ class _WeeklyInsightsCardState extends State<_WeeklyInsightsCard> { @override Widget build(BuildContext context) { - final gemini = context.watch(); + final gemini = context.watch(); final settings = context.watch(); if (!gemini.isConfigured) return const SizedBox.shrink(); diff --git a/workout-logger/lib/screens/widgets/exercise_progress_view.dart b/workout-logger/lib/screens/widgets/exercise_progress_view.dart index d7283c6..640a664 100644 --- a/workout-logger/lib/screens/widgets/exercise_progress_view.dart +++ b/workout-logger/lib/screens/widgets/exercise_progress_view.dart @@ -11,7 +11,7 @@ import 'package:google_fonts/google_fonts.dart'; import '../../models/models.dart'; import '../../services/workout_provider.dart'; import '../../services/settings_provider.dart'; -import '../../services/gemini_service.dart'; +import '../../services/ai/gemini_ai_service.dart'; import '../../theme/app_theme.dart'; import 'rf_widgets.dart'; import '../ai_coach_screen.dart'; @@ -1489,7 +1489,7 @@ class _AskCoachButton extends StatelessWidget { @override Widget build(BuildContext context) { - final gemini = context.watch(); + final gemini = context.watch(); if (!gemini.isConfigured) return const SizedBox.shrink(); final isPlateauing = diff --git a/workout-logger/lib/screens/widgets/muscle_detail_sheet.dart b/workout-logger/lib/screens/widgets/muscle_detail_sheet.dart index e66011d..5fa8f97 100644 --- a/workout-logger/lib/screens/widgets/muscle_detail_sheet.dart +++ b/workout-logger/lib/screens/widgets/muscle_detail_sheet.dart @@ -8,7 +8,7 @@ import 'package:google_fonts/google_fonts.dart'; import '../../services/workout_provider.dart'; import '../../services/settings_provider.dart'; -import '../../services/gemini_service.dart'; +import '../../services/ai/gemini_ai_service.dart'; import '../../services/interfaces/ml_service_interface.dart'; import '../../data/exercise_database.dart'; import '../../theme/app_theme.dart'; @@ -252,7 +252,7 @@ class _AiInsightSectionState extends State<_AiInsightSection> { Future _fetchInsight() async { setState(() => _loading = true); - final gemini = context.read(); + final gemini = context.read(); final settings = context.read(); final mlService = context.read(); final provider = widget.provider; @@ -298,7 +298,7 @@ class _AiInsightSectionState extends State<_AiInsightSection> { @override Widget build(BuildContext context) { - final gemini = context.watch(); + final gemini = context.watch(); if (!gemini.isConfigured) return const SizedBox.shrink(); return Column( diff --git a/workout-logger/lib/screens/widgets/profile_sections.dart b/workout-logger/lib/screens/widgets/profile_sections.dart index cbeb430..d86bfd0 100644 --- a/workout-logger/lib/screens/widgets/profile_sections.dart +++ b/workout-logger/lib/screens/widgets/profile_sections.dart @@ -6,7 +6,7 @@ import 'package:google_fonts/google_fonts.dart'; import 'package:provider/provider.dart'; import '../../services/settings_provider.dart'; -import '../../services/gemini_service.dart'; +import '../../services/ai/gemini_ai_service.dart'; import '../../theme/app_theme.dart'; import 'rf_widgets.dart'; @@ -719,7 +719,7 @@ class _AiSettingsSectionState extends State { setState(() => _saving = true); final key = _ctrl.text.trim(); final settings = context.read(); - final gemini = context.read(); + final gemini = context.read(); try { await settings.setGeminiApiKey(key); gemini.updateApiKey(key); @@ -730,14 +730,14 @@ class _AiSettingsSectionState extends State { Future _selectModel(String modelId) async { final settings = context.read(); - final gemini = context.read(); + final gemini = context.read(); await settings.setGeminiModel(modelId); gemini.updateModel(modelId); } @override Widget build(BuildContext context) { - final gemini = context.watch(); + final gemini = context.watch(); final settings = context.watch(); return _ProfileSection( icon: Icons.auto_awesome_rounded, @@ -899,12 +899,99 @@ class _AiSettingsSectionState extends State { ), ), ), + const SizedBox(height: AppSpacing.md), + Row( + children: [ + const _SectionLabel('TOKEN USAGE'), + const Spacer(), + if (gemini.aiRequestCount > 0) + GestureDetector( + onTap: () => context.read().resetUsage(), + child: Text( + 'Reset', + style: GoogleFonts.geist( + color: AppColors.accent, + fontSize: 11, + fontWeight: FontWeight.w600, + ), + ), + ), + ], + ), + const SizedBox(height: AppSpacing.sm), + Container( + padding: const EdgeInsets.all(AppSpacing.md), + decoration: BoxDecoration( + color: AppColors.glass, + borderRadius: BorderRadius.circular(AppRadius.sm), + border: Border.all(color: AppColors.glassBorder), + ), + child: Column( + children: [ + _UsageRow(label: 'Total tokens', value: _formatInt(gemini.totalTokensUsed)), + const SizedBox(height: 6), + _UsageRow(label: 'Input (prompt)', value: _formatInt(gemini.promptTokensUsed)), + const SizedBox(height: 6), + _UsageRow(label: 'Output (response)', value: _formatInt(gemini.responseTokensUsed)), + const SizedBox(height: 6), + _UsageRow(label: 'Requests', value: _formatInt(gemini.aiRequestCount)), + ], + ), + ), + const SizedBox(height: AppSpacing.sm), + Text( + 'Cumulative billable tokens across coach, program builder & insights.', + style: GoogleFonts.geist( + color: AppColors.textFaint, + fontSize: 11, + fontStyle: FontStyle.italic, + ), + ), ], ), ); } } +/// One label/value line in the token-usage card. +class _UsageRow extends StatelessWidget { + const _UsageRow({required this.label, required this.value}); + final String label; + final String value; + + @override + Widget build(BuildContext context) { + return Row( + mainAxisAlignment: MainAxisAlignment.spaceBetween, + children: [ + Text( + label, + style: GoogleFonts.geist(color: AppColors.textMuted, fontSize: 12), + ), + Text( + value, + style: GoogleFonts.geistMono( + color: AppColors.textPrimary, + fontSize: 12, + fontWeight: FontWeight.w600, + ), + ), + ], + ); + } +} + +/// Format an int with thousands separators (e.g. 12345 → "12,345"). +String _formatInt(int n) { + final s = n.toString(); + final buf = StringBuffer(); + for (var i = 0; i < s.length; i++) { + if (i > 0 && (s.length - i) % 3 == 0) buf.write(','); + buf.write(s[i]); + } + return buf.toString(); +} + class _ComingSoonBadge extends StatelessWidget { const _ComingSoonBadge(); diff --git a/workout-logger/lib/screens/widgets/targets_tab.dart b/workout-logger/lib/screens/widgets/targets_tab.dart index 1865d4b..5cc1744 100644 --- a/workout-logger/lib/screens/widgets/targets_tab.dart +++ b/workout-logger/lib/screens/widgets/targets_tab.dart @@ -9,7 +9,7 @@ import 'package:intl/intl.dart'; import '../../models/models.dart'; import '../../services/workout_provider.dart'; import '../../services/settings_provider.dart'; -import '../../services/gemini_service.dart'; +import '../../services/ai/gemini_ai_service.dart'; import '../../theme/app_theme.dart'; import '../../data/exercise_database.dart'; import 'rf_widgets.dart'; @@ -224,7 +224,7 @@ class _TargetCardWithAiState extends State<_TargetCardWithAi> { Future _fetchNudge() async { setState(() => _loadingNudge = true); - final gemini = context.read(); + final gemini = context.read(); final settings = context.read(); final t = widget.target; @@ -257,7 +257,7 @@ class _TargetCardWithAiState extends State<_TargetCardWithAi> { @override Widget build(BuildContext context) { final settings = context.watch(); - final gemini = context.watch(); + final gemini = context.watch(); final t = widget.target; final pct = t.progressPercentage.clamp(0.0, 100.0); final etaStr = t.estimatedCompletionDate != null @@ -486,7 +486,7 @@ class _CreateTargetSheetState extends State<_CreateTargetSheet> { final provider = context.read(); final settings = context.read(); - final gemini = context.read(); + final gemini = context.read(); final exerciseName = provider.getExerciseName(_selectedExerciseId!); final growth = provider.getGrowthModel(_selectedExerciseId!); final oneRM = provider.getBestOneRM(_selectedExerciseId!); @@ -520,7 +520,7 @@ class _CreateTargetSheetState extends State<_CreateTargetSheet> { Widget build(BuildContext context) { final exercises = ExerciseDatabase.getAll(); final bottom = MediaQuery.of(context).viewInsets.bottom; - final gemini = context.watch(); + final gemini = context.watch(); return Padding( padding: EdgeInsets.fromLTRB( diff --git a/workout-logger/lib/services/ai/coach_tool_service.dart b/workout-logger/lib/services/ai/coach_tool_service.dart new file mode 100644 index 0000000..846c3c9 --- /dev/null +++ b/workout-logger/lib/services/ai/coach_tool_service.dart @@ -0,0 +1,458 @@ +// coach_tool_service.dart — DB-backed function-calling tools for the AI coach. +// +// Exposes a set of read-only query functions the model can call to ground its +// answers in the user's real data. Every tool reuses existing parameterized +// query methods on WorkoutProvider / PRManager — no new analytics logic lives +// here, only the schema + arg parsing + JSON shaping. + +import 'package:google_generative_ai/google_generative_ai.dart'; + +import '../../models/models.dart'; +import '../workout_provider.dart'; +import '../managers/pr_manager.dart'; + +class AmbiguousMatchException implements Exception { + const AmbiguousMatchException(this.candidates); + final List candidates; +} + +class CoachToolService { + final WorkoutProvider _wp; + final PRManager _pr; + + CoachToolService(this._wp, this._pr); + + /// Tool declarations advertised to the model. + List buildTools() => [ + Tool(functionDeclarations: [ + FunctionDeclaration( + 'get_exercise_performance', + 'Get how a specific exercise has progressed: per-session volume ' + 'trend, growth slope, best estimated 1RM, last logged sets, and ' + 'personal record. Use for questions like "how is my bench press ' + 'progressing".', + Schema.object( + properties: { + 'exercise_name': Schema.string( + description: + 'Name of the exercise, e.g. "Bench Press" or "Squat".', + ), + 'days': Schema.integer( + description: + 'Optional. Only consider sessions from the last N days.', + nullable: true, + ), + }, + requiredProperties: ['exercise_name'], + ), + ), + FunctionDeclaration( + 'get_workouts_in_range', + 'Summarize workouts in a date range: session count, total volume, ' + 'and a per-session breakdown. Use for "what did I do last week" ' + 'or "how many workouts in the last 3 months".', + Schema.object( + properties: { + 'start_date': Schema.string( + description: 'Optional ISO date (YYYY-MM-DD) range start.', + nullable: true, + ), + 'end_date': Schema.string( + description: 'Optional ISO date (YYYY-MM-DD) range end.', + nullable: true, + ), + 'days': Schema.integer( + description: + 'Optional. Last N days; overrides start/end when set. ' + 'Defaults to 30 if no dates are provided.', + nullable: true, + ), + }, + ), + ), + FunctionDeclaration( + 'get_routine_performance', + 'Get how a named routine is performing: number of sessions logged ' + 'against it, total volume, volume trend over time, and the ' + 'exercises it contains.', + Schema.object( + properties: { + 'routine_name': Schema.string( + description: 'Name of the routine, e.g. "Push Day".', + ), + 'days': Schema.integer( + description: + 'Optional. Only consider sessions from the last N days.', + nullable: true, + ), + }, + requiredProperties: ['routine_name'], + ), + ), + FunctionDeclaration( + 'get_personal_records', + 'Get personal records (best weight, reps, and single-set volume). ' + 'Pass an exercise name for one exercise, or omit for all PRs.', + Schema.object( + properties: { + 'exercise_name': Schema.string( + description: 'Optional exercise name to filter to.', + nullable: true, + ), + }, + ), + ), + FunctionDeclaration( + 'get_goal_progress', + 'Get progress toward training goals/targets: current vs target ' + 'value, percent complete, and estimated completion date.', + Schema.object( + properties: { + 'exercise_name': Schema.string( + description: 'Optional exercise name to filter goals to.', + nullable: true, + ), + }, + ), + ), + FunctionDeclaration( + 'get_muscle_recovery', + 'Get current per-muscle-group recovery status (percent recovered ' + 'and whether each is ready, recovering, or fatigued). Use for ' + '"what can I train today".', + Schema.object(properties: {}), + ), + ]), + ]; + + /// Dispatch a model function call to the matching query and return a + /// JSON-serializable result map. + Future> handleCall(FunctionCall call) async { + switch (call.name) { + case 'get_exercise_performance': + return _exercisePerformance(call.args); + case 'get_workouts_in_range': + return _workoutsInRange(call.args); + case 'get_routine_performance': + return _routinePerformance(call.args); + case 'get_personal_records': + return _personalRecords(call.args); + case 'get_goal_progress': + return _goalProgress(call.args); + case 'get_muscle_recovery': + return _muscleRecovery(); + default: + return {'error': 'Unknown tool: ${call.name}'}; + } + } + + // ── Tool implementations ─────────────────────────────────────────────────── + + Map _exercisePerformance(Map args) { + final name = (args['exercise_name'] as String?)?.trim() ?? ''; + final Exercise exercise; + try { + final resolved = _resolveExercise(name); + if (resolved == null) { + return { + 'error': 'No exercise found matching "$name".', + 'available_examples': _exampleExerciseNames(), + }; + } + exercise = resolved; + } on AmbiguousMatchException catch (e) { + return { + 'error': 'Multiple exercises match "$name". Did you mean one of:', + 'ambiguous_matches': e.candidates, + }; + } + + final days = (args['days'] as num?)?.toInt(); + final cutoff = + days != null ? DateTime.now().subtract(Duration(days: days)) : null; + + final progression = _wp + .getVolumeProgression(exercise.id) + .where((p) => cutoff == null || !p.date.isBefore(cutoff)) + .toList(); + + final growth = _wp.getGrowthModel(exercise.id); + final lastLog = _wp.getLastSessionForExercise(exercise.id); + final pr = _pr.getRecord(exercise.id); + + return { + 'exercise': exercise.name, + 'session_count': progression.length, + if (days != null) 'window_days': days, + 'volume_trend': [ + for (final p in progression.length > 40 + ? progression.sublist(progression.length - 40) + : progression) + {'date': _d(p.date), 'volume': _round(p.volume)}, + ], + 'growth': growth == null + ? null + : { + 'slope_per_session': _round(growth.slope), + 'r2': _round(growth.r2), + 'trend': growth.slope > 0 + ? 'improving' + : growth.slope < 0 + ? 'declining' + : 'flat', + }, + 'best_estimated_1rm': _roundOrNull(_wp.getBestOneRM(exercise.id)), + 'last_session': lastLog == null + ? null + : [ + for (final s in lastLog.sets) + {'weight': _round(s.weight), 'reps': s.reps}, + ], + 'personal_record': pr == null + ? null + : { + 'best_weight': _round(pr.bestWeight), + 'best_reps': pr.bestReps, + 'best_volume': _round(pr.bestVolume), + 'achieved_at': _d(pr.achievedAt), + }, + }; + } + + Map _workoutsInRange(Map args) { + final now = DateTime.now(); + final days = (args['days'] as num?)?.toInt(); + final startArg = DateTime.tryParse((args['start_date'] as String?) ?? ''); + final endArg = DateTime.tryParse((args['end_date'] as String?) ?? ''); + + final DateTime start; + final DateTime end; + if (days != null) { + start = now.subtract(Duration(days: days)); + end = now; + } else if (startArg != null || endArg != null) { + start = startArg ?? now.subtract(const Duration(days: 30)); + end = endArg ?? now; + } else { + start = now.subtract(const Duration(days: 30)); + end = now; + } + + final sessions = _wp.sessions + .where((s) => !s.date.isBefore(start) && !s.date.isAfter(end)) + .toList() + ..sort((a, b) => b.date.compareTo(a.date)); + + final totalVolume = + sessions.fold(0, (sum, s) => sum + s.totalVolume); + + return { + 'start_date': _d(start), + 'end_date': _d(end), + 'session_count': sessions.length, + 'total_volume': _round(totalVolume), + 'sessions': [ + for (final s in sessions.take(40)) + { + 'date': _d(s.date), + 'duration_min': s.duration, + 'exercise_count': s.exercises.length, + 'volume': _round(s.totalVolume), + 'exercises': [ + for (final e in s.exercises) _wp.getExerciseName(e.exerciseId), + ], + }, + ], + }; + } + + Map _routinePerformance(Map args) { + final name = (args['routine_name'] as String?)?.trim() ?? ''; + final Routine routine; + try { + final resolved = _resolveRoutine(name); + if (resolved == null) { + return { + 'error': 'No routine found matching "$name".', + 'available_routines': [for (final r in _wp.routines) r.name], + }; + } + routine = resolved; + } on AmbiguousMatchException catch (e) { + return { + 'error': 'Multiple routines match "$name". Did you mean one of:', + 'ambiguous_matches': e.candidates, + }; + } + + final days = (args['days'] as num?)?.toInt(); + final cutoff = + days != null ? DateTime.now().subtract(Duration(days: days)) : null; + + final sessions = _wp.sessions + .where((s) => s.routineId == routine.id) + .where((s) => cutoff == null || !s.date.isBefore(cutoff)) + .toList() + ..sort((a, b) => a.date.compareTo(b.date)); + + final totalVolume = + sessions.fold(0, (sum, s) => sum + s.totalVolume); + + return { + 'routine': routine.name, + 'exercises': [for (final id in routine.exerciseIds) _wp.getExerciseName(id)], + 'session_count': sessions.length, + if (days != null) 'window_days': days, + 'total_volume': _round(totalVolume), + 'volume_over_time': [ + for (final s in sessions.length > 40 + ? sessions.sublist(sessions.length - 40) + : sessions) + {'date': _d(s.date), 'volume': _round(s.totalVolume)}, + ], + }; + } + + Map _personalRecords(Map args) { + final name = (args['exercise_name'] as String?)?.trim(); + if (name != null && name.isNotEmpty) { + final Exercise exercise; + try { + final resolved = _resolveExercise(name); + if (resolved == null) { + return {'error': 'No exercise found matching "$name".'}; + } + exercise = resolved; + } on AmbiguousMatchException catch (e) { + return { + 'error': 'Multiple exercises match "$name". Did you mean one of:', + 'ambiguous_matches': e.candidates, + }; + } + final pr = _pr.getRecord(exercise.id); + return { + 'exercise': exercise.name, + 'personal_record': pr == null + ? null + : { + 'best_weight': _round(pr.bestWeight), + 'best_reps': pr.bestReps, + 'best_volume': _round(pr.bestVolume), + 'achieved_at': _d(pr.achievedAt), + }, + }; + } + + return { + 'records': [ + for (final pr in _pr.allRecords) + { + 'exercise': _wp.getExerciseName(pr.exerciseId), + 'best_weight': _round(pr.bestWeight), + 'best_reps': pr.bestReps, + 'best_volume': _round(pr.bestVolume), + 'achieved_at': _d(pr.achievedAt), + }, + ], + }; + } + + Map _goalProgress(Map args) { + final name = (args['exercise_name'] as String?)?.trim(); + Iterable targets = _wp.targets; + if (name != null && name.isNotEmpty) { + final Exercise exercise; + try { + final resolved = _resolveExercise(name); + if (resolved == null) { + return {'error': 'No exercise found matching "$name".'}; + } + exercise = resolved; + } on AmbiguousMatchException catch (e) { + return { + 'error': 'Multiple exercises match "$name". Did you mean one of:', + 'ambiguous_matches': e.candidates, + }; + } + targets = targets.where((t) => t.exerciseId == exercise.id); + } + + return { + 'goals': [ + for (final t in targets) + { + 'exercise': _wp.getExerciseName(t.exerciseId), + 'type': t.targetType, + 'current_value': _round(t.currentValue), + 'target_value': _round(t.targetValue), + 'progress_percent': _round(t.progressPercentage), + 'completed': t.isCompleted, + 'estimated_completion': t.estimatedCompletionDate == null + ? null + : _d(t.estimatedCompletionDate!), + }, + ], + }; + } + + Map _muscleRecovery() { + final scores = _wp.getMuscleRecoveryScores(); + final entries = scores.entries.toList() + ..sort((a, b) => a.value.recoveryPercent.compareTo(b.value.recoveryPercent)); + return { + 'muscles': [ + for (final e in entries) + { + 'muscle': _wp.getMuscleGroupName(e.key), + 'recovery_percent': e.value.recoveryPercent, + 'status': e.value.isRecovered + ? 'ready' + : e.value.isUnderRecovered + ? 'fatigued' + : 'recovering', + }, + ], + }; + } + + // ── Helpers ──────────────────────────────────────────────────────────────── + + Exercise? _resolveExercise(String query) { + final q = query.toLowerCase().trim(); + if (q.isEmpty) return null; + final all = _wp.allExercises; + for (final e in all) { + if (e.name.toLowerCase() == q) return e; + } + final partials = [for (final e in all) if (e.name.toLowerCase().contains(q)) e]; + if (partials.isEmpty) return null; + if (partials.length == 1) return partials.first; + throw AmbiguousMatchException([for (final e in partials) e.name]); + } + + Routine? _resolveRoutine(String query) { + final q = query.toLowerCase().trim(); + if (q.isEmpty) return null; + for (final r in _wp.routines) { + if (r.name.toLowerCase() == q) return r; + } + final partials = [ + for (final r in _wp.routines) if (r.name.toLowerCase().contains(q)) r + ]; + if (partials.isEmpty) return null; + if (partials.length == 1) return partials.first; + throw AmbiguousMatchException([for (final r in partials) r.name]); + } + + List _exampleExerciseNames() => + _wp.allExercises.take(8).map((e) => e.name).toList(); + + String _d(DateTime dt) { + final m = dt.month.toString().padLeft(2, '0'); + final d = dt.day.toString().padLeft(2, '0'); + return '${dt.year}-$m-$d'; + } + + double _round(double v) => (v * 10).round() / 10; + double? _roundOrNull(double? v) => v == null ? null : _round(v); +} diff --git a/workout-logger/lib/services/gemini_service.dart b/workout-logger/lib/services/ai/gemini_ai_service.dart similarity index 51% rename from workout-logger/lib/services/gemini_service.dart rename to workout-logger/lib/services/ai/gemini_ai_service.dart index e87be84..08f811c 100644 --- a/workout-logger/lib/services/gemini_service.dart +++ b/workout-logger/lib/services/ai/gemini_ai_service.dart @@ -1,11 +1,18 @@ -// gemini_service.dart — Gemini AI integration (coach chat, program gen, insights) +// gemini_ai_service.dart — google_generative_ai implementation of IAiService. +// +// Backs the AI coach chat (streaming + tool calling), program generation, and +// insights. Uses a user-supplied Google AI Studio API key (free-tier friendly). +// Implements [IAiService] so the backend can be swapped (e.g. firebase_ai) +// without touching consumers. import 'dart:convert'; import 'package:flutter/foundation.dart'; import 'package:google_generative_ai/google_generative_ai.dart'; import 'package:uuid/uuid.dart'; -import '../models/models.dart'; +import '../../models/models.dart'; +import '../interfaces/ai_service_interface.dart'; +import '../interfaces/storage_service_interface.dart'; // Ordered list of available Gemini models shown in the picker. const kGeminiModels = [ @@ -15,20 +22,116 @@ const kGeminiModels = [ ('gemini-3.5-flash', 'Gemini 3.5 Flash'), ]; -const kDefaultGeminiModel = 'gemini-2.5-flash'; +// Default to a fast, free-tier 3.x model. gemini-3.5-flash is selectable and +// preferable when heavy tool-calling reliability matters. +const kDefaultGeminiModel = 'gemini-3.1-flash-lite'; + +// Upper bound on tool-resolution rounds per user turn, to bound runaway loops. +const int _kMaxToolRounds = 5; + +class GeminiAiService extends ChangeNotifier implements IAiService { + // Optional storage so cumulative token usage survives restarts. + final IStorageService? _storage; + + GeminiAiService({IStorageService? storage}) : _storage = storage; + + static const String _usageKey = 'aiTokenUsage'; -class GeminiService extends ChangeNotifier { String _apiKey = ''; String _model = kDefaultGeminiModel; + // Cumulative token usage across all AI calls (persisted). + int _promptTokens = 0; + int _responseTokens = 0; + int _totalTokens = 0; + int _requestCount = 0; + + @override bool get isConfigured => _apiKey.isNotEmpty; + + @override String get currentModel => _model; + /// Cumulative input (prompt) tokens billed across all AI calls. + int get promptTokensUsed => _promptTokens; + + /// Cumulative output (response) tokens across all AI calls. + int get responseTokensUsed => _responseTokens; + + /// Cumulative total tokens (prompt + response) across all AI calls. + int get totalTokensUsed => _totalTokens; + + /// Number of AI requests recorded. + int get aiRequestCount => _requestCount; + void init(String apiKey, {String model = kDefaultGeminiModel}) { _apiKey = apiKey.trim(); _model = model; } + /// Load persisted cumulative token usage (call once at startup). + Future loadUsage() async { + final raw = await _storage?.getSetting(_usageKey); + if (raw == null || raw.isEmpty) return; + try { + final m = jsonDecode(raw) as Map; + _promptTokens = (m['prompt'] as num?)?.toInt() ?? 0; + _responseTokens = (m['response'] as num?)?.toInt() ?? 0; + _totalTokens = (m['total'] as num?)?.toInt() ?? 0; + _requestCount = (m['requests'] as num?)?.toInt() ?? 0; + notifyListeners(); + } catch (_) { + // Ignore corrupt usage data. + } + } + + /// Reset cumulative token usage to zero. + Future resetUsage() async { + _promptTokens = 0; + _responseTokens = 0; + _totalTokens = 0; + _requestCount = 0; + await _persistUsage(); + notifyListeners(); + } + + /// Accumulate one request's token counts. Exposed for testing; normally + /// fed from a response's [UsageMetadata] via [_recordUsage]. + @visibleForTesting + Future recordUsage({ + required int prompt, + required int response, + required int total, + }) async { + _promptTokens += prompt; + _responseTokens += response; + _totalTokens += total; + _requestCount += 1; + await _persistUsage(); + notifyListeners(); + } + + void _recordUsage(UsageMetadata? m) { + if (m == null) return; + final p = m.promptTokenCount ?? 0; + final r = m.candidatesTokenCount ?? 0; + recordUsage(prompt: p, response: r, total: m.totalTokenCount ?? (p + r)); + } + + Future _persistUsage() async { + final storage = _storage; + if (storage == null) return; + await storage.saveSetting( + _usageKey, + jsonEncode({ + 'prompt': _promptTokens, + 'response': _responseTokens, + 'total': _totalTokens, + 'requests': _requestCount, + }), + ); + } + void updateApiKey(String key) { _apiKey = key.trim(); notifyListeners(); @@ -39,35 +142,73 @@ class GeminiService extends ChangeNotifier { notifyListeners(); } - GenerativeModel _makeModel({bool jsonMode = false, String? system}) { + GenerativeModel _makeModel({ + bool jsonMode = false, + String? system, + List? tools, + }) { return GenerativeModel( model: _model, apiKey: _apiKey, systemInstruction: system != null ? Content.system(system) : null, + tools: tools, generationConfig: jsonMode ? GenerationConfig(responseMimeType: 'application/json') : null, ); } - // ── Coach chat (streaming) ───────────────────────────────────────────────── - // [history] is the prior conversation as alternating user/model Content objects. + // ── Coach chat (streaming + optional tool-call loop) ─────────────────────── + // [history] is the prior conversation as alternating user/model Content. + // When [tools] + [onToolCall] are supplied, function calls the model emits + // are dispatched and their results fed back until a text answer is produced. + @override Stream streamCoachReply({ required String userMessage, required String systemPrompt, required List history, + List? tools, + Future> Function(FunctionCall call)? onToolCall, }) async* { if (!isConfigured) { yield 'Please add your Gemini API key in Profile → AI Features to get started.'; return; } try { - final session = _makeModel(system: systemPrompt).startChat(history: history); - await for (final chunk - in session.sendMessageStream(Content.text(userMessage))) { - final t = chunk.text; - if (t != null && t.isNotEmpty) yield t; + final chat = _makeModel(system: systemPrompt, tools: tools) + .startChat(history: history); + + Content next = Content.text(userMessage); + + for (var round = 0; round < _kMaxToolRounds; round++) { + final calls = []; + UsageMetadata? roundUsage; + await for (final chunk in chat.sendMessageStream(next)) { + final t = chunk.text; + if (t != null && t.isNotEmpty) yield t; + calls.addAll(chunk.functionCalls); + if (chunk.usageMetadata != null) roundUsage = chunk.usageMetadata; + } + // The final chunk of each round carries that round's cumulative usage. + _recordUsage(roundUsage); + + // No tools requested (or no handler) → the streamed text is the answer. + if (calls.isEmpty || onToolCall == null) return; + + // Resolve every requested call and feed the results back as one turn. + final responses = []; + for (final call in calls) { + try { + final result = await onToolCall(call); + responses.add(FunctionResponse(call.name, result)); + } catch (e) { + responses.add(FunctionResponse(call.name, {'error': '$e'})); + } + } + next = Content.functionResponses(responses); } + // Exhausted the tool-round budget without a final text answer. + yield '\n\n_(Stopped after $_kMaxToolRounds tool steps — try rephrasing.)_'; } on GenerativeAIException catch (e) { yield 'AI error: ${e.message}'; } catch (e) { @@ -76,6 +217,7 @@ class GeminiService extends ChangeNotifier { } // ── Program generator (structured JSON output) ──────────────────────────── + @override Future generateProgram({ required String userPrompt, required List allExercises, @@ -143,8 +285,9 @@ Required JSON schema (follow exactly): try { final response = await _makeModel(jsonMode: true, system: systemPrompt) .generateContent([Content.text(prompt)]); + _recordUsage(response.usageMetadata); final raw = response.text ?? ''; - if (raw.isEmpty) throw FormatException('Empty response from Gemini.'); + if (raw.isEmpty) throw const FormatException('Empty response from Gemini.'); final data = jsonDecode(raw) as Map; // Ensure a fresh UUID so it never collides with an existing program. @@ -160,6 +303,7 @@ Required JSON schema (follow exactly): } // ── Weekly insights (single-shot text) ──────────────────────────────────── + @override Future generateWeeklyInsights(String contextText) async { if (!isConfigured) { return 'Add your Gemini API key in Profile → AI Features to unlock insights.'; @@ -173,6 +317,7 @@ Required JSON schema (follow exactly): try { final response = await _makeModel(system: systemPrompt) .generateContent([Content.text(contextText)]); + _recordUsage(response.usageMetadata); return response.text?.trim() ?? 'No insights generated.'; } on GenerativeAIException catch (e) { return 'AI error: ${e.message}'; @@ -182,9 +327,7 @@ Required JSON schema (follow exactly): } // ── Generic one-shot insight (contextual) ───────────────────────────────── - // Thin, tool-agnostic helper for on-demand contextual insights (muscle - // drill-down, target suggestions, stalled-target nudges). Kept generic so a - // future function-calling path can be added additively over [_makeModel]. + @override Future generateInsight(String system, String context) async { if (!isConfigured) { return 'Add your Gemini API key in Profile → AI Features to unlock insights.'; @@ -192,6 +335,7 @@ Required JSON schema (follow exactly): try { final response = await _makeModel(system: system) .generateContent([Content.text(context)]); + _recordUsage(response.usageMetadata); return response.text?.trim() ?? 'No insight generated.'; } on GenerativeAIException catch (e) { return 'AI error: ${e.message}'; diff --git a/workout-logger/lib/services/gemini_context_builder.dart b/workout-logger/lib/services/gemini_context_builder.dart index 362af1e..a017bb6 100644 --- a/workout-logger/lib/services/gemini_context_builder.dart +++ b/workout-logger/lib/services/gemini_context_builder.dart @@ -1,88 +1,49 @@ // gemini_context_builder.dart — Builds rich context strings from app data for Gemini prompts. import '../models/models.dart'; -import 'interfaces/ml_service_interface.dart'; class GeminiContextBuilder { const GeminiContextBuilder._(); // ── Coach system prompt ──────────────────────────────────────────────────── + // + // Deliberately STATIC (no per-turn workout data) so the prefix stays + // byte-identical across a conversation and Gemini's implicit prompt caching + // can engage. All live data is fetched on demand via the coach tools + // (see CoachToolService), not embedded here. static String buildCoachSystemPrompt({ - required List recentSessions, - required Map exerciseMap, - required Map recoveryScores, - required List activeTargets, String? userName, String unitLabel = 'kg', + DateTime? now, }) { + final n = now ?? DateTime.now(); + final today = '${n.year}-${n.month.toString().padLeft(2, '0')}-' + '${n.day.toString().padLeft(2, '0')}'; + final buf = StringBuffer() ..writeln( 'You are an expert personal trainer embedded in RepForge, a workout tracking app.', ) ..writeln( 'Answer concisely (under 180 words unless a plan is requested). ' - 'Be encouraging and specific — always reference the user\'s actual data.', + 'Be encouraging and specific.', + ) + ..writeln('Today is $today. Use this when interpreting relative dates ' + '("last week", "3 months ago").') + ..writeln( + 'This prompt contains NO workout data. To answer anything about the ' + 'user\'s training — exercise progression, workouts in a date range, ' + 'routine performance, personal records, goal progress, or muscle ' + 'recovery — CALL THE PROVIDED TOOLS rather than guessing or inventing ' + 'numbers. Pass ISO dates (YYYY-MM-DD) or a day count to the tools.', + ) + ..writeln( + 'Weights are in $unitLabel. Format replies with Markdown (lists, bold, ' + 'tables) where it aids clarity.', ); if (userName != null && userName.isNotEmpty) { - buf.writeln('\nUser: $userName'); - } - - // Recent sessions - buf.writeln('\n--- RECENT SESSIONS (last 14 days) ---'); - final cutoff = DateTime.now().subtract(const Duration(days: 14)); - final recent = recentSessions - .where((s) => s.date.isAfter(cutoff)) - .toList() - ..sort((a, b) => b.date.compareTo(a.date)); - - if (recent.isEmpty) { - buf.writeln('No sessions in the last 14 days.'); - } else { - for (final s in recent.take(8)) { - final date = '${_weekday(s.date.weekday)} ${s.date.day}/${s.date.month}'; - final exParts = s.exercises.map((e) { - final name = exerciseMap[e.exerciseId]?.name ?? e.exerciseId; - final sets = e.sets - .map((ws) => '${ws.weight}$unitLabel×${ws.reps}') - .join(', '); - return '$name [$sets]'; - }); - buf.writeln('$date: ${exParts.join(' | ')}'); - } - } - - // Muscle recovery - buf.writeln('\n--- MUSCLE RECOVERY ---'); - if (recoveryScores.isEmpty) { - buf.writeln('No recovery data yet.'); - } else { - final sorted = recoveryScores.entries.toList() - ..sort((a, b) => a.value.recoveryPercent.compareTo(b.value.recoveryPercent)); - for (final e in sorted) { - final name = e.key.replaceAll('_', ' '); - final pct = e.value.recoveryPercent; - final tag = e.value.isRecovered - ? 'ready' - : e.value.isUnderRecovered - ? 'fatigued' - : 'recovering'; - buf.writeln('$name: $pct% ($tag)'); - } - } - - // Active goals - buf.writeln('\n--- ACTIVE GOALS ---'); - if (activeTargets.isEmpty) { - buf.writeln('No active goals set.'); - } else { - for (final t in activeTargets) { - final name = exerciseMap[t.exerciseId]?.name ?? t.exerciseId; - final progress = t.progressPercentage.toStringAsFixed(0); - buf.writeln( - '$name: ${t.currentValue}$unitLabel → ${t.targetValue}$unitLabel ($progress%)', - ); - } + buf.writeln('\nThe user\'s name is $userName.'); } return buf.toString(); diff --git a/workout-logger/lib/services/interfaces/ai_service_interface.dart b/workout-logger/lib/services/interfaces/ai_service_interface.dart new file mode 100644 index 0000000..6f0d2d1 --- /dev/null +++ b/workout-logger/lib/services/interfaces/ai_service_interface.dart @@ -0,0 +1,52 @@ +// Abstract AI Service Interface (Dependency Inversion Principle) +// +// Defines the contract for the conversational AI / generation backend. +// High-level modules (the coach ViewModel, program generator) depend on this +// abstraction rather than a concrete SDK, so the backend can be swapped (e.g. +// google_generative_ai today → firebase_ai later) without touching consumers. +// +// The signatures intentionally use the google_generative_ai content model +// (Content / Tool / FunctionCall). firebase_ai exposes an almost identical +// shape, so a future backend swap is a mechanical adapter rather than a rewrite. + +import 'package:google_generative_ai/google_generative_ai.dart'; + +import '../../models/models.dart'; + +/// Contract for the AI backend used across RepForge (coach chat, program +/// generation, insights). Implemented by [GeminiAiService] today. +abstract class IAiService { + /// True once an API key (or equivalent credential) has been supplied. + bool get isConfigured; + + /// The model identifier currently in use (e.g. `gemini-3.1-flash-lite`). + String get currentModel; + + /// Stream a coach reply token-by-token. + /// + /// When [tools] and [onToolCall] are provided, the implementation runs a + /// tool-call loop: any function calls the model emits are dispatched through + /// [onToolCall] and their results fed back, until the model produces a final + /// natural-language answer. Only text is yielded to the caller. + Stream streamCoachReply({ + required String userMessage, + required String systemPrompt, + required List history, + List? tools, + Future> Function(FunctionCall call)? onToolCall, + }); + + /// Generate a structured multi-week training program from a natural-language + /// prompt, constrained to the provided exercise catalogue. + Future generateProgram({ + required String userPrompt, + required List allExercises, + }); + + /// One-shot weekly training summary in conversational prose. + Future generateWeeklyInsights(String contextText); + + /// Generic one-shot contextual insight given a [system] instruction and + /// [context] payload. + Future generateInsight(String system, String context); +} diff --git a/workout-logger/lib/services/interfaces/storage_service_interface.dart b/workout-logger/lib/services/interfaces/storage_service_interface.dart index 9274ea1..34132f6 100644 --- a/workout-logger/lib/services/interfaces/storage_service_interface.dart +++ b/workout-logger/lib/services/interfaces/storage_service_interface.dart @@ -73,6 +73,13 @@ abstract class IStorageService { Future getPersonalRecord(String exerciseId); Future> getAllPersonalRecords(); + // ==================== AI CONVERSATIONS ==================== + + Future saveConversation(Conversation conversation); + Future> getAllConversations(); + Future getConversation(String id); + Future deleteConversation(String id); + // ==================== EXPORT / IMPORT ==================== Future exportAllData(); diff --git a/workout-logger/lib/services/managers/conversation_manager.dart b/workout-logger/lib/services/managers/conversation_manager.dart new file mode 100644 index 0000000..ed39a52 --- /dev/null +++ b/workout-logger/lib/services/managers/conversation_manager.dart @@ -0,0 +1,118 @@ +// Conversation Manager (Single Responsibility Principle) +// +// Single source of truth for persisted AI coach conversations. Owns the +// in-memory list + the currently active conversation, and mirrors every +// mutation to storage. Does NOT talk to the AI backend — that's the +// AiCoachViewModel's job. + +import 'package:flutter/foundation.dart'; +import '../../models/models.dart'; +import '../interfaces/storage_service_interface.dart'; + +/// Manages the lifecycle of AI coach [Conversation]s (load, create, append, +/// rename, delete) backed by [IStorageService]. +class ConversationManager extends ChangeNotifier { + final IStorageService _storage; + + List _conversations = []; + Conversation? _active; + + ConversationManager(this._storage); + + /// All conversations, most-recently-updated first. + List get conversations => List.unmodifiable(_conversations); + + /// The conversation currently shown in the coach screen, or null for a + /// fresh (unsaved) chat. + Conversation? get active => _active; + + /// Messages of the active conversation (empty for a fresh chat). + List get activeMessages => _active?.messages ?? const []; + + /// Load all conversations from storage. Does not change the active one. + Future loadConversations() async { + _conversations = await _storage.getAllConversations(); + notifyListeners(); + } + + /// Begin a fresh conversation. Nothing is persisted until the first message + /// is appended (avoids littering storage with empty chats). + void startNewConversation() { + _active = null; + notifyListeners(); + } + + /// Make [id] the active conversation, if it exists. + void selectConversation(String id) { + final idx = _conversations.indexWhere((c) => c.id == id); + if (idx < 0) return; + _active = _conversations[idx]; + notifyListeners(); + } + + /// Append [message] to the active conversation, creating one if needed, + /// then persist. The conversation title is derived from the first user + /// message. Bumps `updatedAt` and re-sorts the list newest-first. + Future appendMessage(ChatMessage message) async { + final current = _active; + final Conversation updated; + + if (current == null) { + updated = Conversation( + title: _deriveTitle(message), + messages: [message], + ); + } else { + final title = current.title.isEmpty && message.role == 'user' + ? _deriveTitle(message) + : current.title; + updated = current.copyWith( + title: title, + updatedAt: DateTime.now(), + messages: [...current.messages, message], + ); + } + + _active = updated; + _upsert(updated); + await _storage.saveConversation(updated); + notifyListeners(); + } + + /// Rename a conversation. + Future renameConversation(String id, String title) async { + final idx = _conversations.indexWhere((c) => c.id == id); + if (idx < 0) return; + final updated = _conversations[idx].copyWith( + title: title.trim(), + updatedAt: DateTime.now(), + ); + if (_active?.id == id) _active = updated; + _upsert(updated); + await _storage.saveConversation(updated); + notifyListeners(); + } + + /// Delete a conversation. Clears the active one if it was deleted. + Future deleteConversation(String id) async { + _conversations = _conversations.where((c) => c.id != id).toList(); + if (_active?.id == id) _active = null; + await _storage.deleteConversation(id); + notifyListeners(); + } + + // ── Helpers ──────────────────────────────────────────────────────────────── + + void _upsert(Conversation conversation) { + final next = _conversations.where((c) => c.id != conversation.id).toList() + ..add(conversation) + ..sort((a, b) => b.updatedAt.compareTo(a.updatedAt)); + _conversations = next; + } + + String _deriveTitle(ChatMessage message) { + final text = message.text.trim().replaceAll(RegExp(r'\s+'), ' '); + if (text.isEmpty) return 'New chat'; + return text.length <= 40 ? text : '${text.substring(0, 40).trim()}…'; + } +} diff --git a/workout-logger/lib/services/storage_service.dart b/workout-logger/lib/services/storage_service.dart index 40cdc39..873de63 100644 --- a/workout-logger/lib/services/storage_service.dart +++ b/workout-logger/lib/services/storage_service.dart @@ -25,6 +25,7 @@ class StorageService implements IStorageService { static const String _settingsBox = 'settings'; static const String _trainingProgramsBox = 'training_programs'; static const String _personalRecordsBox = 'personal_records'; + static const String _aiConversationsBox = 'ai_conversations'; late Box _sessionsBox; late Box _routinesBoxInstance; @@ -34,6 +35,7 @@ class StorageService implements IStorageService { late Box _settingsBoxInstance; late Box _trainingProgramsBoxInstance; late Box _personalRecordsBoxInstance; + late Box _aiConversationsBoxInstance; String _appVersion = const String.fromEnvironment( 'APP_VERSION', @@ -73,6 +75,9 @@ class StorageService implements IStorageService { _personalRecordsBoxInstance = await Hive.openBox( _personalRecordsBox, ); + _aiConversationsBoxInstance = await Hive.openBox( + _aiConversationsBox, + ); // Initialize default muscle groups if empty if (_muscleGroupsBoxInstance.isEmpty) { @@ -364,6 +369,9 @@ class StorageService implements IStorageService { 'customExercises': _customExercisesBoxInstance.values .map(_normalizeExportValue) .toList(growable: false), + 'conversations': _aiConversationsBoxInstance.values + .map(_normalizeExportValue) + .toList(growable: false), 'settings': settingsMap, 'exportDate': DateTime.now().toIso8601String(), 'appVersion': _appVersion, @@ -455,6 +463,20 @@ class StorageService implements IStorageService { } } } + + // Import AI conversations (merge: skip if id already exists) + final conversations = data['conversations']; + if (conversations is List) { + for (var item in conversations) { + final map = _normalizeImportItem(item); + if (map == null) continue; + final conversation = Conversation.fromJson(map); + final existing = await getConversation(conversation.id); + if (existing == null) { + await saveConversation(conversation); + } + } + } } // ==================== TRAINING PROGRAMS ==================== @@ -515,6 +537,39 @@ class StorageService implements IStorageService { return records; } + // ==================== AI CONVERSATIONS ==================== + + @override + Future saveConversation(Conversation conversation) async { + await _aiConversationsBoxInstance.put( + conversation.id, + jsonEncode(conversation.toJson()), + ); + } + + @override + Future> getAllConversations() async { + final conversations = []; + for (final json in _aiConversationsBoxInstance.values) { + conversations.add(Conversation.fromJson(jsonDecode(json))); + } + // Most recently updated first. + conversations.sort((a, b) => b.updatedAt.compareTo(a.updatedAt)); + return conversations; + } + + @override + Future getConversation(String id) async { + final json = _aiConversationsBoxInstance.get(id); + if (json == null) return null; + return Conversation.fromJson(jsonDecode(json)); + } + + @override + Future deleteConversation(String id) async { + await _aiConversationsBoxInstance.delete(id); + } + // ==================== STATS ==================== @override diff --git a/workout-logger/lib/viewmodels/ai_coach_view_model.dart b/workout-logger/lib/viewmodels/ai_coach_view_model.dart new file mode 100644 index 0000000..036154c --- /dev/null +++ b/workout-logger/lib/viewmodels/ai_coach_view_model.dart @@ -0,0 +1,143 @@ +// ai_coach_view_model.dart — orchestration for the AI coach screen. +// +// Owns all coach logic so the View stays dumb: builds the system prompt, +// drives the streaming tool-call loop via IAiService + CoachToolService, and +// persists each turn through ConversationManager. Exposes immutable state. + +import 'package:flutter/foundation.dart'; +import 'package:google_generative_ai/google_generative_ai.dart' show Content, TextPart; + +import '../models/models.dart'; +import '../services/interfaces/ai_service_interface.dart'; +import '../services/ai/coach_tool_service.dart'; +import '../services/managers/conversation_manager.dart'; +import '../services/settings_provider.dart'; +import '../services/gemini_context_builder.dart'; + +class AiCoachViewModel extends ChangeNotifier { + final IAiService _ai; + final CoachToolService _coachTools; + final ConversationManager _conversations; + final SettingsProvider _settings; + + bool _loading = false; + String _streamingText = ''; + + AiCoachViewModel({ + required IAiService ai, + required CoachToolService coachTools, + required ConversationManager conversations, + required SettingsProvider settings, + }) : _ai = ai, + _coachTools = coachTools, + _conversations = conversations, + _settings = settings { + // Forward conversation-store changes so the View only watches the VM. + _conversations.addListener(notifyListeners); + } + + @override + void dispose() { + _conversations.removeListener(notifyListeners); + super.dispose(); + } + + // ── Exposed state (immutable snapshots) ──────────────────────────────────── + + bool get isConfigured => _ai.isConfigured; + bool get isLoading => _loading; + String get streamingText => _streamingText; + List get messages => _conversations.activeMessages; + List get conversations => _conversations.conversations; + String? get activeConversationId => _conversations.active?.id; + + // ── Commands ─────────────────────────────────────────────────────────────── + + /// Load the persisted conversation list (call when the screen opens). + Future loadConversations() => _conversations.loadConversations(); + + /// Start a fresh, unsaved conversation. + void newConversation() { + if (_loading) return; + _conversations.startNewConversation(); + } + + /// Switch to an existing conversation. + void selectConversation(String id) { + if (_loading) return; + _conversations.selectConversation(id); + } + + /// Delete a conversation. + Future deleteConversation(String id) => + _conversations.deleteConversation(id); + + /// Send a user message and stream the coach's reply (running the tool-call + /// loop). Both the user message and the final reply are persisted. + Future sendMessage(String text) async { + final trimmed = text.trim(); + if (trimmed.isEmpty || _loading) return; + + _loading = true; + _streamingText = ''; + notifyListeners(); + + // Persist the user message first; history is derived from the store. + await _conversations.appendMessage( + ChatMessage(role: 'user', text: trimmed), + ); + + final systemPrompt = _buildSystemPrompt(); + final history = _buildHistory(); + + final buffer = StringBuffer(); + try { + await for (final chunk in _ai.streamCoachReply( + userMessage: trimmed, + systemPrompt: systemPrompt, + history: history, + tools: _coachTools.buildTools(), + onToolCall: _coachTools.handleCall, + )) { + buffer.write(chunk); + _streamingText = buffer.toString(); + notifyListeners(); + } + final reply = buffer.toString().trim(); + if (reply.isNotEmpty) { + await _conversations.appendMessage( + ChatMessage(role: 'model', text: reply), + ); + } + } catch (e) { + buffer.write('\n\n_Error: ${e}_'); + final errText = buffer.toString().trim(); + if (errText.isNotEmpty) { + await _conversations.appendMessage( + ChatMessage(role: 'model', text: errText), + ); + } + } finally { + _streamingText = ''; + _loading = false; + notifyListeners(); + } + } + + // ── Internals ────────────────────────────────────────────────────────────── + + // Static prompt — live data is fetched by the model via the coach tools, + // keeping the prefix stable for implicit prompt caching. + String _buildSystemPrompt() => GeminiContextBuilder.buildCoachSystemPrompt( + userName: _settings.userName, + unitLabel: _settings.unitLabel, + ); + + /// Prior turns (everything before the user message just appended). + List _buildHistory() { + final msgs = _conversations.activeMessages; + final prior = + msgs.length > 1 ? msgs.sublist(0, msgs.length - 1) : []; + return prior.map((m) => Content(m.role, [TextPart(m.text)])).toList(); + } +} diff --git a/workout-logger/pubspec.yaml b/workout-logger/pubspec.yaml index 65ed2c6..d4a544d 100644 --- a/workout-logger/pubspec.yaml +++ b/workout-logger/pubspec.yaml @@ -65,6 +65,7 @@ dependencies: file_picker: ^10.3.10 path_provider: ^2.1.5 share_plus: ^12.0.1 + gpt_markdown: ^1.1.7 dev_dependencies: flutter_test: diff --git a/workout-logger/test/ai_coach_view_model_test.dart b/workout-logger/test/ai_coach_view_model_test.dart new file mode 100644 index 0000000..2832a2a --- /dev/null +++ b/workout-logger/test/ai_coach_view_model_test.dart @@ -0,0 +1,148 @@ +// Unit tests for AiCoachViewModel — verifies orchestration (send → stream → +// persist) using a fake IAiService, so the View has no logic left to test. + +import 'package:flutter_test/flutter_test.dart'; +import 'package:google_generative_ai/google_generative_ai.dart' + show Content, Tool, FunctionCall; +import 'package:repforge/models/models.dart'; +import 'package:repforge/services/interfaces/ai_service_interface.dart'; +import 'package:repforge/services/ai/coach_tool_service.dart'; +import 'package:repforge/services/managers/conversation_manager.dart'; +import 'package:repforge/services/managers/program_manager.dart'; +import 'package:repforge/services/managers/pr_manager.dart'; +import 'package:repforge/services/workout_provider.dart'; +import 'package:repforge/services/settings_provider.dart'; +import 'package:repforge/viewmodels/ai_coach_view_model.dart'; +import 'test_utils/mock_storage_service.dart'; + +/// Scripted IAiService: yields fixed chunks; optionally invokes a tool first. +class _FakeAiService implements IAiService { + _FakeAiService({this.chunks = const ['Hello ', 'world'], this.invokeTool = false}); + + final List chunks; + final bool invokeTool; + int toolCallsMade = 0; + + @override + bool get isConfigured => true; + + @override + String get currentModel => 'fake-model'; + + @override + Stream streamCoachReply({ + required String userMessage, + required String systemPrompt, + required List history, + List? tools, + Future> Function(FunctionCall call)? onToolCall, + }) async* { + if (invokeTool && onToolCall != null) { + await onToolCall(FunctionCall('get_muscle_recovery', {})); + toolCallsMade++; + } + for (final c in chunks) { + yield c; + } + } + + @override + Future generateProgram({ + required String userPrompt, + required List allExercises, + }) => + throw UnimplementedError(); + + @override + Future generateWeeklyInsights(String contextText) async => ''; + + @override + Future generateInsight(String system, String context) async => ''; +} + +void main() { + group('AiCoachViewModel', () { + late MockStorageService storage; + late WorkoutProvider provider; + late ConversationManager conversations; + late SettingsProvider settings; + late PRManager pr; + + Future buildVm(_FakeAiService ai) async { + provider = WorkoutProvider( + storage, + programManager: ProgramManager(storage), + ); + await provider.init(); + pr = PRManager(storage); + settings = SettingsProvider(storage); + conversations = ConversationManager(storage); + return AiCoachViewModel( + ai: ai, + coachTools: CoachToolService(provider, pr), + conversations: conversations, + settings: settings, + ); + } + + setUp(() { + storage = MockStorageService(); + }); + + test('sendMessage appends user + model messages and persists', () async { + final vm = await buildVm(_FakeAiService()); + + await vm.sendMessage('How am I doing?'); + + expect(vm.messages, hasLength(2)); + expect(vm.messages[0].role, 'user'); + expect(vm.messages[0].text, 'How am I doing?'); + expect(vm.messages[1].role, 'model'); + expect(vm.messages[1].text, 'Hello world'); + expect(vm.isLoading, isFalse); + expect(vm.streamingText, isEmpty); + + // Persisted. + final stored = await storage.getAllConversations(); + expect(stored, hasLength(1)); + expect(stored.first.messages, hasLength(2)); + }); + + test('blank or whitespace messages are ignored', () async { + final vm = await buildVm(_FakeAiService()); + await vm.sendMessage(' '); + expect(vm.messages, isEmpty); + }); + + test('runs the tool-call loop via CoachToolService', () async { + final ai = _FakeAiService(invokeTool: true, chunks: const ['done']); + final vm = await buildVm(ai); + + await vm.sendMessage('what can I train?'); + + expect(ai.toolCallsMade, 1); + expect(vm.messages.last.text, 'done'); + }); + + test('newConversation then selectConversation swaps active state', + () async { + final vm = await buildVm(_FakeAiService()); + + await vm.sendMessage('first chat'); + final firstId = vm.activeConversationId; + expect(firstId, isNotNull); + + vm.newConversation(); + expect(vm.messages, isEmpty); + + await vm.sendMessage('second chat'); + final secondId = vm.activeConversationId; + expect(secondId, isNot(firstId)); + expect(vm.conversations, hasLength(2)); + + vm.selectConversation(firstId!); + expect(vm.activeConversationId, firstId); + expect(vm.messages.first.text, 'first chat'); + }); + }); +} diff --git a/workout-logger/test/analytics_screen_test.dart b/workout-logger/test/analytics_screen_test.dart index 038c2a9..fbfb123 100644 --- a/workout-logger/test/analytics_screen_test.dart +++ b/workout-logger/test/analytics_screen_test.dart @@ -11,7 +11,7 @@ import 'package:repforge/models/models.dart'; import 'package:repforge/screens/analytics_screen.dart'; import 'package:repforge/services/workout_provider.dart'; import 'package:repforge/services/settings_provider.dart'; -import 'package:repforge/services/gemini_service.dart'; +import 'package:repforge/services/ai/gemini_ai_service.dart'; import 'package:repforge/services/managers/program_manager.dart'; import 'package:repforge/services/managers/pr_manager.dart'; import 'package:repforge/services/interfaces/ml_service_interface.dart'; @@ -31,7 +31,7 @@ Widget _wrap({ ChangeNotifierProvider.value(value: workoutProvider), ChangeNotifierProvider.value(value: sp), ChangeNotifierProvider.value(value: prManager), - ChangeNotifierProvider.value(value: GeminiService()), + ChangeNotifierProvider.value(value: GeminiAiService()), Provider.value(value: MockMLService()), ], child: const MaterialApp(home: AnalyticsScreen()), diff --git a/workout-logger/test/coach_tool_service_test.dart b/workout-logger/test/coach_tool_service_test.dart new file mode 100644 index 0000000..4cecf55 --- /dev/null +++ b/workout-logger/test/coach_tool_service_test.dart @@ -0,0 +1,170 @@ +// Unit tests for CoachToolService — each tool returns expected JSON shapes, +// backed by a seeded WorkoutProvider + PRManager. + +import 'package:flutter_test/flutter_test.dart'; +import 'package:google_generative_ai/google_generative_ai.dart' show FunctionCall; +import 'package:repforge/models/models.dart'; +import 'package:repforge/services/workout_provider.dart'; +import 'package:repforge/services/managers/program_manager.dart'; +import 'package:repforge/services/managers/pr_manager.dart'; +import 'package:repforge/services/ai/coach_tool_service.dart'; +import 'test_utils/mock_storage_service.dart'; + +void main() { + group('CoachToolService', () { + late MockStorageService storage; + late WorkoutProvider provider; + late PRManager pr; + late CoachToolService tools; + + WorkoutSession benchSession(DateTime date, double weight, {String? routineId}) { + return WorkoutSession( + id: 'sess-${date.millisecondsSinceEpoch}', + date: date, + routineId: routineId, + duration: 45, + exercises: [ + ExerciseLog( + exerciseId: 'bench_press', + sets: [ + WorkoutSet(weight: weight, reps: 8), + WorkoutSet(weight: weight, reps: 8), + ], + ), + ], + ); + } + + setUp(() async { + storage = MockStorageService(); + + // Two bench sessions on different days → enough for a growth model. + final now = DateTime.now(); + storage.addMockRoutine( + Routine(id: 'r1', name: 'Push Day', exerciseIds: ['bench_press']), + ); + storage.addMockSession( + benchSession(now.subtract(const Duration(days: 10)), 60, routineId: 'r1'), + ); + storage.addMockSession( + benchSession(now.subtract(const Duration(days: 3)), 65, routineId: 'r1'), + ); + + provider = WorkoutProvider( + storage, + programManager: ProgramManager(storage), + ); + await provider.init(); + + pr = PRManager(storage); + await pr.backfillFromSessions(provider.sessions); + + tools = CoachToolService(provider, pr); + }); + + test('exposes the expected tool declarations', () { + final declared = tools + .buildTools() + .expand((t) => t.functionDeclarations ?? []) + .map((f) => f.name) + .toSet(); + expect( + declared, + containsAll([ + 'get_exercise_performance', + 'get_workouts_in_range', + 'get_routine_performance', + 'get_personal_records', + 'get_goal_progress', + 'get_muscle_recovery', + ]), + ); + }); + + test('get_exercise_performance returns trend + PR for a known exercise', + () async { + final result = await tools.handleCall( + FunctionCall('get_exercise_performance', {'exercise_name': 'Bench Press'}), + ); + + expect(result['exercise'], 'Bench Press'); + expect(result['session_count'], 2); + expect(result['volume_trend'], isA>()); + expect((result['volume_trend'] as List), isNotEmpty); + expect(result['personal_record'], isNotNull); + }); + + test('get_exercise_performance returns an error for an unknown exercise', + () async { + final result = await tools.handleCall( + FunctionCall('get_exercise_performance', {'exercise_name': 'Nonexistent'}), + ); + expect(result['error'], isNotNull); + expect(result['available_examples'], isA>()); + }); + + test('get_workouts_in_range summarizes sessions in the window', () async { + final result = await tools.handleCall( + FunctionCall('get_workouts_in_range', {'days': 30}), + ); + expect(result['session_count'], 2); + expect(result['total_volume'], isA()); + expect((result['total_volume'] as num) > 0, isTrue); + }); + + test('get_routine_performance returns sessions logged against the routine', + () async { + final result = await tools.handleCall( + FunctionCall('get_routine_performance', {'routine_name': 'Push Day'}), + ); + expect(result['routine'], 'Push Day'); + expect(result['session_count'], 2); + expect(result['exercises'], contains('Bench Press')); + }); + + test('get_routine_performance errors for an unknown routine', () async { + final result = await tools.handleCall( + FunctionCall('get_routine_performance', {'routine_name': 'Leg Day'}), + ); + expect(result['error'], isNotNull); + expect(result['available_routines'], contains('Push Day')); + }); + + test('get_personal_records returns all records when unfiltered', () async { + final result = await tools.handleCall( + FunctionCall('get_personal_records', {}), + ); + final records = result['records'] as List; + expect(records, isNotEmpty); + expect((records.first as Map)['exercise'], 'Bench Press'); + }); + + test('get_goal_progress reflects active targets', () async { + await provider.createTarget( + exerciseId: 'bench_press', + type: 'weight', + targetValue: 100, + ); + + final result = await tools.handleCall( + FunctionCall('get_goal_progress', {'exercise_name': 'Bench Press'}), + ); + final goals = result['goals'] as List; + expect(goals, hasLength(1)); + expect((goals.first as Map)['type'], 'weight'); + expect((goals.first as Map)['target_value'], 100); + }); + + test('get_muscle_recovery returns per-muscle status', () async { + final result = await tools.handleCall( + FunctionCall('get_muscle_recovery', {}), + ); + final muscles = result['muscles'] as List; + expect(muscles, isNotEmpty); + final first = muscles.first as Map; + expect(first['muscle'], isA()); + expect(first['recovery_percent'], isA()); + expect(first['status'], isA()); + }); + }); +} diff --git a/workout-logger/test/conversation_manager_test.dart b/workout-logger/test/conversation_manager_test.dart new file mode 100644 index 0000000..e533313 --- /dev/null +++ b/workout-logger/test/conversation_manager_test.dart @@ -0,0 +1,117 @@ +// Unit tests for ConversationManager — persistence + active-conversation logic. + +import 'package:flutter_test/flutter_test.dart'; +import 'package:repforge/models/models.dart'; +import 'package:repforge/services/managers/conversation_manager.dart'; +import 'test_utils/mock_storage_service.dart'; + +void main() { + group('ConversationManager', () { + late MockStorageService storage; + late ConversationManager manager; + + setUp(() { + storage = MockStorageService(); + manager = ConversationManager(storage); + }); + + test('appendMessage creates a conversation and persists it', () async { + await manager.appendMessage( + ChatMessage(role: 'user', text: 'How is my bench press?'), + ); + + expect(manager.active, isNotNull); + expect(manager.activeMessages, hasLength(1)); + + // Persisted to storage. + final stored = await storage.getAllConversations(); + expect(stored, hasLength(1)); + expect(stored.first.messages.first.text, 'How is my bench press?'); + }); + + test('title is derived from the first user message', () async { + await manager.appendMessage( + ChatMessage(role: 'user', text: 'Plan my next push day please'), + ); + expect(manager.active!.title, 'Plan my next push day please'); + }); + + test('long first message title is truncated', () async { + final long = 'a' * 80; + await manager.appendMessage(ChatMessage(role: 'user', text: long)); + expect(manager.active!.title.length, lessThanOrEqualTo(41)); + expect(manager.active!.title.endsWith('…'), isTrue); + }); + + test('multiple messages append to the same active conversation', () async { + await manager.appendMessage(ChatMessage(role: 'user', text: 'hi')); + await manager.appendMessage(ChatMessage(role: 'model', text: 'hello!')); + + expect(manager.activeMessages, hasLength(2)); + final stored = await storage.getAllConversations(); + expect(stored, hasLength(1)); + expect(stored.first.messages, hasLength(2)); + }); + + test('reload restores conversations from storage', () async { + await manager.appendMessage(ChatMessage(role: 'user', text: 'first')); + + final fresh = ConversationManager(storage); + await fresh.loadConversations(); + expect(fresh.conversations, hasLength(1)); + expect(fresh.conversations.first.messages.first.text, 'first'); + }); + + test('conversations are sorted most-recently-updated first', () async { + // Small delays keep updatedAt timestamps distinct (millisecond clock). + await manager.appendMessage(ChatMessage(role: 'user', text: 'older')); + final olderId = manager.active!.id; + + await Future.delayed(const Duration(milliseconds: 5)); + manager.startNewConversation(); + await manager.appendMessage(ChatMessage(role: 'user', text: 'newer')); + final newerId = manager.active!.id; + + expect(manager.conversations.first.id, newerId); + + // Touching the older one bumps it to the front. + await Future.delayed(const Duration(milliseconds: 5)); + manager.selectConversation(olderId); + await manager.appendMessage(ChatMessage(role: 'model', text: 'reply')); + expect(manager.conversations.first.id, olderId); + }); + + test('startNewConversation clears the active conversation', () async { + await manager.appendMessage(ChatMessage(role: 'user', text: 'hi')); + expect(manager.active, isNotNull); + + manager.startNewConversation(); + expect(manager.active, isNull); + expect(manager.activeMessages, isEmpty); + // The prior conversation is still saved. + expect(manager.conversations, hasLength(1)); + }); + + test('deleteConversation removes it and clears active when needed', () async { + await manager.appendMessage(ChatMessage(role: 'user', text: 'hi')); + final id = manager.active!.id; + + await manager.deleteConversation(id); + + expect(manager.active, isNull); + expect(manager.conversations, isEmpty); + expect(await storage.getAllConversations(), isEmpty); + }); + + test('renameConversation updates the title and persists', () async { + await manager.appendMessage(ChatMessage(role: 'user', text: 'hi')); + final id = manager.active!.id; + + await manager.renameConversation(id, 'My chat'); + + expect(manager.active!.title, 'My chat'); + final stored = await storage.getConversation(id); + expect(stored!.title, 'My chat'); + }); + }); +} diff --git a/workout-logger/test/exercise_progress_view_test.dart b/workout-logger/test/exercise_progress_view_test.dart index 2913cbe..f8f1ba4 100644 --- a/workout-logger/test/exercise_progress_view_test.dart +++ b/workout-logger/test/exercise_progress_view_test.dart @@ -17,7 +17,7 @@ import 'package:provider/provider.dart'; import 'package:repforge/models/models.dart'; import 'package:repforge/services/workout_provider.dart'; import 'package:repforge/services/settings_provider.dart'; -import 'package:repforge/services/gemini_service.dart'; +import 'package:repforge/services/ai/gemini_ai_service.dart'; import 'package:repforge/services/managers/program_manager.dart'; import 'package:repforge/services/interfaces/ml_service_interface.dart'; import 'package:repforge/screens/widgets/exercise_progress_view.dart'; @@ -38,7 +38,7 @@ Widget _wrap({ providers: [ ChangeNotifierProvider.value(value: provider), ChangeNotifierProvider.value(value: sp), - ChangeNotifierProvider.value(value: GeminiService()), + ChangeNotifierProvider.value(value: GeminiAiService()), Provider.value(value: MockMLService()), ], child: MaterialApp(home: Scaffold(body: child)), diff --git a/workout-logger/test/gemini_ai_service_usage_test.dart b/workout-logger/test/gemini_ai_service_usage_test.dart new file mode 100644 index 0000000..046bdae --- /dev/null +++ b/workout-logger/test/gemini_ai_service_usage_test.dart @@ -0,0 +1,59 @@ +// Unit tests for GeminiAiService token-usage tracking + persistence. + +import 'package:flutter_test/flutter_test.dart'; +import 'package:repforge/services/ai/gemini_ai_service.dart'; +import 'test_utils/mock_storage_service.dart'; + +void main() { + group('GeminiAiService token usage', () { + late MockStorageService storage; + late GeminiAiService service; + + setUp(() { + storage = MockStorageService(); + service = GeminiAiService(storage: storage); + }); + + test('starts at zero', () { + expect(service.totalTokensUsed, 0); + expect(service.promptTokensUsed, 0); + expect(service.responseTokensUsed, 0); + expect(service.aiRequestCount, 0); + }); + + test('recordUsage accumulates across calls', () { + service.recordUsage(prompt: 10, response: 5, total: 15); + service.recordUsage(prompt: 20, response: 10, total: 30); + + expect(service.promptTokensUsed, 30); + expect(service.responseTokensUsed, 15); + expect(service.totalTokensUsed, 45); + expect(service.aiRequestCount, 2); + }); + + test('usage is persisted and reloaded by a fresh instance', () async { + await service.recordUsage(prompt: 100, response: 40, total: 140); + + final reloaded = GeminiAiService(storage: storage); + await reloaded.loadUsage(); + + expect(reloaded.promptTokensUsed, 100); + expect(reloaded.responseTokensUsed, 40); + expect(reloaded.totalTokensUsed, 140); + expect(reloaded.aiRequestCount, 1); + }); + + test('resetUsage zeros counters and persists', () async { + await service.recordUsage(prompt: 100, response: 40, total: 140); + await service.resetUsage(); + + expect(service.totalTokensUsed, 0); + expect(service.aiRequestCount, 0); + + final reloaded = GeminiAiService(storage: storage); + await reloaded.loadUsage(); + expect(reloaded.totalTokensUsed, 0); + expect(reloaded.aiRequestCount, 0); + }); + }); +} diff --git a/workout-logger/test/test_utils/mock_storage_service.dart b/workout-logger/test/test_utils/mock_storage_service.dart index 9685c3b..2f0a9c3 100644 --- a/workout-logger/test/test_utils/mock_storage_service.dart +++ b/workout-logger/test/test_utils/mock_storage_service.dart @@ -21,6 +21,7 @@ class MockStorageService implements IStorageService { final Map _settings = {}; final List _trainingPrograms = []; final Map _personalRecords = {}; + final Map _conversations = {}; bool saveCustomExerciseCalled = false; Exercise? lastSavedExercise; @@ -279,6 +280,26 @@ class MockStorageService implements IStorageService { Future> getAllPersonalRecords() async => List.from(_personalRecords.values); + @override + Future saveConversation(Conversation conversation) async { + _conversations[conversation.id] = conversation; + } + + @override + Future> getAllConversations() async { + final list = _conversations.values.toList() + ..sort((a, b) => b.updatedAt.compareTo(a.updatedAt)); + return list; + } + + @override + Future getConversation(String id) async => _conversations[id]; + + @override + Future deleteConversation(String id) async { + _conversations.remove(id); + } + @override Future exportAllData() async => '{}';