CommentWriterViewModel.swift 5.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193
  1. import AppKit
  2. import SwiftUI
  3. @MainActor
  4. @Observable
  5. final class CommentWriterViewModel {
  6. var draft = CommentDraft()
  7. var selectedTab: CommentWriterTab = .compose
  8. var isGenerating = false
  9. var errorMessage: String?
  10. var successMessage: String?
  11. var variants: [CommentVariant] = []
  12. var selectedVariantID: UUID?
  13. var showOverwriteConfirmation = false
  14. private var hasPremiumAccess = false
  15. private let injectedGenerationService: (any CommentGenerationServiceProtocol)?
  16. private let historyManager: AIHistoryManager
  17. private let emptyDraft = CommentDraft()
  18. init(
  19. hasPremiumAccess: Bool = false,
  20. generationService: (any CommentGenerationServiceProtocol)? = nil,
  21. historyManager: AIHistoryManager = .shared
  22. ) {
  23. self.hasPremiumAccess = hasPremiumAccess
  24. self.injectedGenerationService = generationService
  25. self.historyManager = historyManager
  26. }
  27. func setPremiumAccess(_ hasPremiumAccess: Bool) {
  28. self.hasPremiumAccess = hasPremiumAccess
  29. }
  30. private var generationService: any CommentGenerationServiceProtocol {
  31. injectedGenerationService ?? CommentGenerationServiceFactory.make(hasPremiumAccess: hasPremiumAccess)
  32. }
  33. var formattedSubreddit: String {
  34. let name = PostDraftValidator.normalizedSubreddit(draft.subreddit)
  35. guard !name.isEmpty else { return "r/subreddit" }
  36. return "r/\(name)"
  37. }
  38. var usesLiveAI: Bool {
  39. AIConfiguration.usesLiveAI(hasPremiumAccess: hasPremiumAccess)
  40. }
  41. var hasDraftChanges: Bool {
  42. draft != emptyDraft
  43. }
  44. var canGenerate: Bool {
  45. CommentDraftValidator.canGenerate(from: draft, isGenerating: isGenerating)
  46. }
  47. var canExport: Bool {
  48. let body = activeCommentBody.trimmingCharacters(in: .whitespaces)
  49. guard !body.isEmpty, body.count <= CommentDraftValidator.maxCommentLength else { return false }
  50. return (try? CommentDraftValidator.validateForGeneration(from: draft)) != nil
  51. }
  52. var activeCommentBody: String {
  53. if let selectedID = selectedVariantID,
  54. let variant = variants.first(where: { $0.id == selectedID }) {
  55. return variant.body
  56. }
  57. return draft.body
  58. }
  59. var bestVariant: CommentVariant? {
  60. variants.max(by: { $0.score < $1.score })
  61. }
  62. var needsOverwriteConfirmation: Bool {
  63. !draft.body.trimmingCharacters(in: .whitespaces).isEmpty
  64. }
  65. func selectCommentType(_ type: CommentType) {
  66. draft.commentType = type
  67. if type == .topLevel {
  68. draft.parentComment = ""
  69. }
  70. clearMessages()
  71. }
  72. func selectVariant(_ variant: CommentVariant) {
  73. selectedVariantID = variant.id
  74. clearMessages()
  75. }
  76. func applyVariant(_ variant: CommentVariant) {
  77. draft.body = variant.body
  78. selectedVariantID = variant.id
  79. successMessage = "Variant applied."
  80. errorMessage = nil
  81. }
  82. func generateComment() async {
  83. guard canGenerate else { return }
  84. if needsOverwriteConfirmation {
  85. showOverwriteConfirmation = true
  86. return
  87. }
  88. await performGenerate(replaceExisting: true)
  89. }
  90. func confirmOverwriteAndGenerate() async {
  91. showOverwriteConfirmation = false
  92. await performGenerate(replaceExisting: true)
  93. }
  94. func cancelOverwriteConfirmation() {
  95. showOverwriteConfirmation = false
  96. }
  97. func copyToClipboard() {
  98. let content = activeCommentBody.trimmingCharacters(in: .whitespaces)
  99. guard !content.isEmpty else {
  100. errorMessage = CommentDraftValidationError.emptyComment.localizedDescription
  101. successMessage = nil
  102. return
  103. }
  104. guard content.count <= CommentDraftValidator.maxCommentLength else {
  105. errorMessage = CommentDraftValidationError.commentTooLong.localizedDescription
  106. successMessage = nil
  107. return
  108. }
  109. NSPasteboard.general.clearContents()
  110. NSPasteboard.general.setString(content, forType: .string)
  111. successMessage = "Comment copied to clipboard."
  112. errorMessage = nil
  113. }
  114. func resetDraft() {
  115. draft = CommentDraft()
  116. selectedTab = .compose
  117. variants = []
  118. selectedVariantID = nil
  119. clearMessages()
  120. }
  121. func restore(from entry: AIHistoryEntry) {
  122. guard case .commentWriter(let storedDraft, let storedVariants, let selectedID) = entry.payload else {
  123. return
  124. }
  125. draft = storedDraft.commentDraft
  126. variants = storedVariants.map(\.commentVariant)
  127. selectedVariantID = selectedID
  128. selectedTab = variants.isEmpty ? .compose : .preview
  129. clearMessages()
  130. successMessage = "Restored from history."
  131. }
  132. func clearMessages() {
  133. errorMessage = nil
  134. successMessage = nil
  135. }
  136. func notifyDraftEdited() {
  137. clearMessages()
  138. }
  139. private func performGenerate(replaceExisting: Bool) async {
  140. isGenerating = true
  141. errorMessage = nil
  142. successMessage = nil
  143. do {
  144. let result = try await generationService.generateComment(from: draft)
  145. if replaceExisting || draft.body.trimmingCharacters(in: .whitespaces).isEmpty {
  146. draft.body = String(result.body.prefix(CommentDraftValidator.maxCommentLength))
  147. }
  148. variants = result.variants
  149. selectedVariantID = result.variants.first?.id
  150. selectedTab = .preview
  151. historyManager.save(AIHistoryEntry.fromCommentWriter(
  152. draft: draft,
  153. variants: variants,
  154. selectedVariantID: selectedVariantID
  155. ))
  156. let engine = usesLiveAI ? "AI" : "local AI templates"
  157. successMessage = "Generated comment with \(result.variants.count) variants using \(engine)."
  158. } catch {
  159. errorMessage = UserFacingError.message(for: error)
  160. }
  161. isGenerating = false
  162. }
  163. }