CommentWriterViewModel.swift 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253
  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 var hasEverPurchasedPremium = false
  16. private let injectedGenerationService: (any CommentGenerationServiceProtocol)?
  17. private let historyManager: AIHistoryManager
  18. private let freeUsageManager: AIFreeUsageManager
  19. var onPaywallRequired: (() -> Void)?
  20. private(set) var remainingFreeAIUses: Int
  21. private let emptyDraft = CommentDraft()
  22. init(
  23. hasPremiumAccess: Bool = false,
  24. generationService: (any CommentGenerationServiceProtocol)? = nil,
  25. historyManager: AIHistoryManager = .shared,
  26. freeUsageManager: AIFreeUsageManager = .shared
  27. ) {
  28. self.hasPremiumAccess = hasPremiumAccess
  29. self.injectedGenerationService = generationService
  30. self.historyManager = historyManager
  31. self.freeUsageManager = freeUsageManager
  32. self.remainingFreeAIUses = freeUsageManager.remainingUses
  33. }
  34. func setSubscriptionState(hasPremiumAccess: Bool, hasEverPurchasedPremium: Bool) {
  35. self.hasPremiumAccess = hasPremiumAccess
  36. self.hasEverPurchasedPremium = hasEverPurchasedPremium
  37. refreshFreeUsageState()
  38. }
  39. func setPremiumAccess(_ hasPremiumAccess: Bool) {
  40. setSubscriptionState(
  41. hasPremiumAccess: hasPremiumAccess,
  42. hasEverPurchasedPremium: hasEverPurchasedPremium
  43. )
  44. }
  45. func refreshFreeUsageState() {
  46. remainingFreeAIUses = freeUsageManager.remainingUses
  47. }
  48. private var generationService: any CommentGenerationServiceProtocol {
  49. injectedGenerationService ?? CommentGenerationServiceFactory.make(usesLiveAI: usesLiveAI)
  50. }
  51. var formattedSubreddit: String {
  52. let name = PostDraftValidator.normalizedSubreddit(draft.subreddit)
  53. guard !name.isEmpty else { return "r/subreddit" }
  54. return "r/\(name)"
  55. }
  56. var usesLiveAI: Bool {
  57. AIConfiguration.usesLiveAI(
  58. hasPremiumAccess: hasPremiumAccess,
  59. canUseFreeTrial: freeUsageManager.canUseFreeLiveAI(
  60. hasEverPurchasedPremium: hasEverPurchasedPremium
  61. )
  62. )
  63. }
  64. var generationEngineSubtitle: String {
  65. if usesLiveAI {
  66. if hasPremiumAccess {
  67. return "Write Reddit-ready comments with AI"
  68. }
  69. let remaining = remainingFreeAIUses
  70. let useLabel = remaining == 1 ? "use" : "uses"
  71. return "Write Reddit-ready comments with AI (\(remaining) free \(useLabel) left)"
  72. }
  73. if hasEverPurchasedPremium {
  74. return "Upgrade to Premium for unlimited AI generations"
  75. }
  76. return "Write Reddit-ready comments with AI templates"
  77. }
  78. var hasDraftChanges: Bool {
  79. draft != emptyDraft
  80. }
  81. var canGenerate: Bool {
  82. CommentDraftValidator.canGenerate(from: draft, isGenerating: isGenerating)
  83. }
  84. var canExport: Bool {
  85. let body = activeCommentBody.trimmingCharacters(in: .whitespaces)
  86. guard !body.isEmpty, body.count <= CommentDraftValidator.maxCommentLength else { return false }
  87. return (try? CommentDraftValidator.validateForGeneration(from: draft)) != nil
  88. }
  89. var activeCommentBody: String {
  90. if let selectedID = selectedVariantID,
  91. let variant = variants.first(where: { $0.id == selectedID }) {
  92. return variant.body
  93. }
  94. return draft.body
  95. }
  96. var bestVariant: CommentVariant? {
  97. variants.max(by: { $0.score < $1.score })
  98. }
  99. var needsOverwriteConfirmation: Bool {
  100. !draft.body.trimmingCharacters(in: .whitespaces).isEmpty
  101. }
  102. func selectCommentType(_ type: CommentType) {
  103. draft.commentType = type
  104. if type == .topLevel {
  105. draft.parentComment = ""
  106. }
  107. clearMessages()
  108. }
  109. func selectVariant(_ variant: CommentVariant) {
  110. selectedVariantID = variant.id
  111. clearMessages()
  112. }
  113. func applyVariant(_ variant: CommentVariant) {
  114. draft.body = variant.body
  115. selectedVariantID = variant.id
  116. successMessage = "Variant applied."
  117. errorMessage = nil
  118. }
  119. func generateComment() async {
  120. guard canGenerate else { return }
  121. if requiresPaywallForGeneration() {
  122. onPaywallRequired?()
  123. return
  124. }
  125. if needsOverwriteConfirmation {
  126. showOverwriteConfirmation = true
  127. return
  128. }
  129. await performGenerate(replaceExisting: true)
  130. }
  131. func confirmOverwriteAndGenerate() async {
  132. showOverwriteConfirmation = false
  133. await performGenerate(replaceExisting: true)
  134. }
  135. func cancelOverwriteConfirmation() {
  136. showOverwriteConfirmation = false
  137. }
  138. func copyToClipboard() {
  139. let content = activeCommentBody.trimmingCharacters(in: .whitespaces)
  140. guard !content.isEmpty else {
  141. errorMessage = CommentDraftValidationError.emptyComment.localizedDescription
  142. successMessage = nil
  143. return
  144. }
  145. guard content.count <= CommentDraftValidator.maxCommentLength else {
  146. errorMessage = CommentDraftValidationError.commentTooLong.localizedDescription
  147. successMessage = nil
  148. return
  149. }
  150. NSPasteboard.general.clearContents()
  151. NSPasteboard.general.setString(content, forType: .string)
  152. successMessage = "Comment copied to clipboard."
  153. errorMessage = nil
  154. }
  155. func resetDraft() {
  156. draft = CommentDraft()
  157. selectedTab = .compose
  158. variants = []
  159. selectedVariantID = nil
  160. clearMessages()
  161. }
  162. func restore(from entry: AIHistoryEntry) {
  163. guard case .commentWriter(let storedDraft, let storedVariants, let selectedID) = entry.payload else {
  164. return
  165. }
  166. draft = storedDraft.commentDraft
  167. variants = storedVariants.map(\.commentVariant)
  168. selectedVariantID = selectedID
  169. selectedTab = variants.isEmpty ? .compose : .preview
  170. clearMessages()
  171. successMessage = "Restored from history."
  172. }
  173. func clearMessages() {
  174. errorMessage = nil
  175. successMessage = nil
  176. }
  177. func notifyDraftEdited() {
  178. clearMessages()
  179. }
  180. private func performGenerate(replaceExisting: Bool) async {
  181. if requiresPaywallForGeneration() {
  182. onPaywallRequired?()
  183. return
  184. }
  185. isGenerating = true
  186. errorMessage = nil
  187. successMessage = nil
  188. let usedLiveAI = usesLiveAI
  189. do {
  190. let result = try await generationService.generateComment(from: draft)
  191. if replaceExisting || draft.body.trimmingCharacters(in: .whitespaces).isEmpty {
  192. draft.body = String(result.body.prefix(CommentDraftValidator.maxCommentLength))
  193. }
  194. variants = result.variants
  195. selectedVariantID = result.variants.first?.id
  196. selectedTab = .preview
  197. if usedLiveAI, !hasPremiumAccess, !hasEverPurchasedPremium {
  198. freeUsageManager.recordUse(for: .commentWriter)
  199. refreshFreeUsageState()
  200. }
  201. historyManager.save(AIHistoryEntry.fromCommentWriter(
  202. draft: draft,
  203. variants: variants,
  204. selectedVariantID: selectedVariantID
  205. ))
  206. let engine = usedLiveAI ? "AI" : "local AI templates"
  207. successMessage = "Generated comment with \(result.variants.count) variants using \(engine)."
  208. } catch {
  209. errorMessage = UserFacingError.message(for: error)
  210. }
  211. isGenerating = false
  212. }
  213. private func requiresPaywallForGeneration() -> Bool {
  214. !hasPremiumAccess && (hasEverPurchasedPremium || remainingFreeAIUses == 0)
  215. }
  216. }