CommentWriterViewModel.swift 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180
  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 let generationService: any CommentGenerationServiceProtocol
  15. private let emptyDraft = CommentDraft()
  16. init(generationService: (any CommentGenerationServiceProtocol)? = nil) {
  17. self.generationService = generationService ?? CommentGenerationServiceFactory.makeDefault()
  18. }
  19. var formattedSubreddit: String {
  20. let name = PostDraftValidator.normalizedSubreddit(draft.subreddit)
  21. guard !name.isEmpty else { return "r/subreddit" }
  22. return "r/\(name)"
  23. }
  24. var usesLiveAI: Bool {
  25. AIConfiguration.usesLiveAI
  26. }
  27. var hasDraftChanges: Bool {
  28. draft != emptyDraft
  29. }
  30. var canGenerate: Bool {
  31. CommentDraftValidator.canGenerate(from: draft, isGenerating: isGenerating)
  32. }
  33. var canExport: Bool {
  34. let body = activeCommentBody.trimmingCharacters(in: .whitespaces)
  35. guard !body.isEmpty, body.count <= CommentDraftValidator.maxCommentLength else { return false }
  36. return (try? CommentDraftValidator.validateForGeneration(from: draft)) != nil
  37. }
  38. var activeCommentBody: String {
  39. if let selectedID = selectedVariantID,
  40. let variant = variants.first(where: { $0.id == selectedID }) {
  41. return variant.body
  42. }
  43. return draft.body
  44. }
  45. var bestVariant: CommentVariant? {
  46. variants.max(by: { $0.score < $1.score })
  47. }
  48. var needsOverwriteConfirmation: Bool {
  49. !draft.body.trimmingCharacters(in: .whitespaces).isEmpty
  50. }
  51. func selectCommentType(_ type: CommentType) {
  52. draft.commentType = type
  53. if type == .topLevel {
  54. draft.parentComment = ""
  55. }
  56. clearMessages()
  57. }
  58. func selectVariant(_ variant: CommentVariant) {
  59. selectedVariantID = variant.id
  60. clearMessages()
  61. }
  62. func applyVariant(_ variant: CommentVariant) {
  63. draft.body = variant.body
  64. selectedVariantID = variant.id
  65. successMessage = "Variant applied."
  66. errorMessage = nil
  67. }
  68. func generateComment() async {
  69. guard canGenerate else { return }
  70. if needsOverwriteConfirmation {
  71. showOverwriteConfirmation = true
  72. return
  73. }
  74. await performGenerate(replaceExisting: true)
  75. }
  76. func confirmOverwriteAndGenerate() async {
  77. showOverwriteConfirmation = false
  78. await performGenerate(replaceExisting: true)
  79. }
  80. func cancelOverwriteConfirmation() {
  81. showOverwriteConfirmation = false
  82. }
  83. func copyToClipboard() {
  84. let content = activeCommentBody.trimmingCharacters(in: .whitespaces)
  85. guard !content.isEmpty else {
  86. errorMessage = CommentDraftValidationError.emptyComment.localizedDescription
  87. successMessage = nil
  88. return
  89. }
  90. guard content.count <= CommentDraftValidator.maxCommentLength else {
  91. errorMessage = CommentDraftValidationError.commentTooLong.localizedDescription
  92. successMessage = nil
  93. return
  94. }
  95. NSPasteboard.general.clearContents()
  96. NSPasteboard.general.setString(content, forType: .string)
  97. successMessage = "Comment copied to clipboard."
  98. errorMessage = nil
  99. }
  100. func openInReddit(openInApp: (URL) -> Void) {
  101. let urlString = draft.postURL.trimmingCharacters(in: .whitespaces)
  102. guard !urlString.isEmpty else {
  103. errorMessage = "Add a Reddit post URL to open the thread."
  104. successMessage = nil
  105. return
  106. }
  107. guard CommentDraftValidator.isValidRedditURL(urlString) else {
  108. errorMessage = CommentDraftValidationError.invalidPostURL.localizedDescription
  109. successMessage = nil
  110. return
  111. }
  112. guard let url = URL(string: urlString) else { return }
  113. openInApp(url)
  114. successMessage = "Opened thread in Reddit. Paste your comment when ready."
  115. errorMessage = nil
  116. }
  117. func resetDraft() {
  118. draft = CommentDraft()
  119. selectedTab = .compose
  120. variants = []
  121. selectedVariantID = nil
  122. clearMessages()
  123. }
  124. func clearMessages() {
  125. errorMessage = nil
  126. successMessage = nil
  127. }
  128. func notifyDraftEdited() {
  129. clearMessages()
  130. }
  131. private func performGenerate(replaceExisting: Bool) async {
  132. isGenerating = true
  133. errorMessage = nil
  134. successMessage = nil
  135. do {
  136. let result = try await generationService.generateComment(from: draft)
  137. if replaceExisting || draft.body.trimmingCharacters(in: .whitespaces).isEmpty {
  138. draft.body = String(result.body.prefix(CommentDraftValidator.maxCommentLength))
  139. }
  140. variants = result.variants
  141. selectedVariantID = result.variants.first?.id
  142. selectedTab = .preview
  143. let engine = usesLiveAI ? "OpenAI" : "local AI templates"
  144. successMessage = "Generated comment with \(result.variants.count) variants using \(engine)."
  145. } catch {
  146. errorMessage = error.localizedDescription
  147. }
  148. isGenerating = false
  149. }
  150. }