diff --git a/scribe/lib/scribe/ai/data/model/ai_message.dart b/scribe/lib/scribe/ai/data/model/ai_message.dart index e4cedee4e..985d22529 100644 --- a/scribe/lib/scribe/ai/data/model/ai_message.dart +++ b/scribe/lib/scribe/ai/data/model/ai_message.dart @@ -2,12 +2,16 @@ import 'package:json_annotation/json_annotation.dart'; part 'ai_message.g.dart'; +enum AIRole { + @JsonValue('user') + user, + @JsonValue('system') + system, +} + @JsonSerializable() class AIMessage { - static const String aiUserRole = 'user'; - static const String aiSystemRole = 'system'; - - final String role; + final AIRole role; final String content; const AIMessage({ @@ -21,12 +25,12 @@ class AIMessage { Map toJson() => _$AIMessageToJson(this); factory AIMessage.ofUser(String content) => AIMessage( - role: aiUserRole, + role: AIRole.user, content: content, ); factory AIMessage.ofSystem(String content) => AIMessage( - role: aiSystemRole, + role: AIRole.system, content: content, ); } diff --git a/scribe/lib/scribe/ai/domain/model/prompt_data.dart b/scribe/lib/scribe/ai/domain/model/prompt_data.dart index 14509de9b..abc062039 100644 --- a/scribe/lib/scribe/ai/domain/model/prompt_data.dart +++ b/scribe/lib/scribe/ai/domain/model/prompt_data.dart @@ -54,9 +54,9 @@ class Prompt { List buildPrompt(String inputText, {String? task}) { final messages = []; for (final message in this.messages) { - if (message.role == 'system') { + if (message.role == AIRole.system) { messages.add(AIMessage.ofSystem(message.content)); - } else if (message.role == 'user') { + } else if (message.role == AIRole.user) { final userContent = _replacePlaceholders(message.content, inputText, task); messages.add(AIMessage.ofUser(userContent)); } diff --git a/scribe/test/scribe/ai/domain/model/prompt_data_test.dart b/scribe/test/scribe/ai/domain/model/prompt_data_test.dart index 836d0c460..25e9846e3 100644 --- a/scribe/test/scribe/ai/domain/model/prompt_data_test.dart +++ b/scribe/test/scribe/ai/domain/model/prompt_data_test.dart @@ -85,17 +85,17 @@ void main() { // Assert expect(prompt.name, 'test-prompt'); expect(prompt.messages.length, 2); - expect(prompt.messages.first.role, 'system'); + expect(prompt.messages.first.role, AIRole.system); expect(prompt.messages.first.content, 'System message'); - expect(prompt.messages.last.role, 'user'); + expect(prompt.messages.last.role, AIRole.user); expect(prompt.messages.last.content, 'User message with {{input}} placeholder'); }); test('buildPrompt should replace input placeholder correctly', () { // Arrange final messages = [ - const AIMessage(role: 'system', content: 'System message'), - const AIMessage(role: 'user', content: 'User message with {{input}} placeholder') + const AIMessage(role: AIRole.system, content: 'System message'), + const AIMessage(role: AIRole.user, content: 'User message with {{input}} placeholder') ]; final prompt = Prompt(name: 'test-prompt', messages: messages); @@ -104,17 +104,17 @@ void main() { // Assert expect(result.length, 2); - expect(result.first.role, 'system'); + expect(result.first.role, AIRole.system); expect(result.first.content, 'System message'); - expect(result.last.role, 'user'); + expect(result.last.role, AIRole.user); expect(result.last.content, 'User message with test input value placeholder'); }); test('buildPrompt should replace task placeholder when provided', () { // Arrange final messages = [ - const AIMessage(role: 'system', content: 'System message'), - const AIMessage(role: 'user', content: 'Task: {{task}}, Input: {{input}}') + const AIMessage(role: AIRole.system, content: 'System message'), + const AIMessage(role: AIRole.user, content: 'Task: {{task}}, Input: {{input}}') ]; final prompt = Prompt(name: 'test-prompt', messages: messages); @@ -123,17 +123,17 @@ void main() { // Assert expect(result.length, 2); - expect(result.first.role, 'system'); + expect(result.first.role, AIRole.system); expect(result.first.content, 'System message'); - expect(result.last.role, 'user'); + expect(result.last.role, AIRole.user); expect(result.last.content, 'Task: test task value, Input: test input value'); }); test('buildPrompt should not replace task placeholder when not provided', () { // Arrange final messages = [ - const AIMessage(role: 'system', content: 'System message'), - const AIMessage(role: 'user', content: 'Task: {{task}}, Input: {{input}}') + const AIMessage(role: AIRole.system, content: 'System message'), + const AIMessage(role: AIRole.user, content: 'Task: {{task}}, Input: {{input}}') ]; final prompt = Prompt(name: 'test-prompt', messages: messages); @@ -142,17 +142,17 @@ void main() { // Assert expect(result.length, 2); - expect(result.first.role, 'system'); + expect(result.first.role, AIRole.system); expect(result.first.content, 'System message'); - expect(result.last.role, 'user'); + expect(result.last.role, AIRole.user); expect(result.last.content, 'Task: {{task}}, Input: test input value'); }); test('buildPrompt should handle messages without placeholders', () { // Arrange final messages = [ - const AIMessage(role: 'system', content: 'System message'), - const AIMessage(role: 'user', content: 'User message without placeholders') + const AIMessage(role: AIRole.system, content: 'System message'), + const AIMessage(role: AIRole.user, content: 'User message without placeholders') ]; final prompt = Prompt(name: 'test-prompt', messages: messages); diff --git a/scribe/test/scribe/ai/domain/service/prompt_service_test.dart b/scribe/test/scribe/ai/domain/service/prompt_service_test.dart index c40b13d63..f308f847a 100644 --- a/scribe/test/scribe/ai/domain/service/prompt_service_test.dart +++ b/scribe/test/scribe/ai/domain/service/prompt_service_test.dart @@ -27,8 +27,8 @@ void main() { // Assert expect(messages.length, 2); - expect(messages.first.role, 'system'); - expect(messages.last.role, 'user'); + expect(messages.first.role, AIRole.system); + expect(messages.last.role, AIRole.user); expect(messages.last.content, contains('Hello, how are you?')); }); @@ -41,8 +41,8 @@ void main() { // Assert expect(messages.length, 2); - expect(messages.first.role, 'system'); - expect(messages.last.role, 'user'); + expect(messages.first.role, AIRole.system); + expect(messages.last.role, AIRole.user); expect(messages.last.content, contains('Hello, how are you?')); expect(messages.last.content, contains('Make it more casual')); }); @@ -56,8 +56,8 @@ void main() { Prompt( name: 'test-prompt', messages: [ - const AIMessage(role: 'system', content: 'System message'), - const AIMessage(role: 'user', content: 'User message with {{input}}') + const AIMessage(role: AIRole.system, content: 'System message'), + const AIMessage(role: AIRole.user, content: 'User message with {{input}}') ] ) ] @@ -81,8 +81,8 @@ void main() { Prompt( name: 'test-prompt', messages: [ - const AIMessage(role: 'system', content: 'System message'), - const AIMessage(role: 'user', content: 'User message') + const AIMessage(role: AIRole.system, content: 'System message'), + const AIMessage(role: AIRole.user, content: 'User message') ] ) ]