PostGeneratorViewModel.swift 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367
  1. import AppKit
  2. import SwiftUI
  3. @MainActor
  4. @Observable
  5. final class PostGeneratorViewModel {
  6. var draft = PostDraft()
  7. var selectedTab: PostGeneratorTab = .compose
  8. var isGenerating = false
  9. var errorMessage: String?
  10. var successMessage: String?
  11. var showImageImporter = false
  12. var showOverwriteConfirmation = false
  13. /// True after a successful Generate or History restore. Cleared on Reset.
  14. /// Controls whether title/body result fields are shown in Compose.
  15. private(set) var hasGeneratedContent = false
  16. private var hasPremiumAccess = false
  17. private var hasEverPurchasedPremium = false
  18. private var generationService: any PostGenerationServiceProtocol
  19. private let injectedGenerationService: (any PostGenerationServiceProtocol)?
  20. private let historyManager: AIHistoryManager
  21. private let freeUsageManager: AIFreeUsageManager
  22. var onPaywallRequired: (() -> Void)?
  23. private(set) var remainingFreeAIUses: Int
  24. private var imageAccessURL: URL?
  25. private var isAccessingImageResource = false
  26. private let emptyDraft = PostDraft()
  27. init(
  28. hasPremiumAccess: Bool = false,
  29. generationService: (any PostGenerationServiceProtocol)? = nil,
  30. historyManager: AIHistoryManager = .shared,
  31. freeUsageManager: AIFreeUsageManager = .shared
  32. ) {
  33. self.hasPremiumAccess = hasPremiumAccess
  34. self.injectedGenerationService = generationService
  35. self.historyManager = historyManager
  36. self.freeUsageManager = freeUsageManager
  37. self.remainingFreeAIUses = freeUsageManager.remainingUses
  38. self.hasEverPurchasedPremium = false
  39. self.generationService = generationService
  40. ?? PostGenerationServiceFactory.make(usesLiveAI: Self.liveAIEnabled(
  41. hasPremiumAccess: hasPremiumAccess,
  42. hasEverPurchasedPremium: false,
  43. freeUsageManager: freeUsageManager
  44. ))
  45. }
  46. func setSubscriptionState(hasPremiumAccess: Bool, hasEverPurchasedPremium: Bool) {
  47. self.hasPremiumAccess = hasPremiumAccess
  48. self.hasEverPurchasedPremium = hasEverPurchasedPremium
  49. refreshFreeUsageState()
  50. }
  51. func setPremiumAccess(_ hasPremiumAccess: Bool) {
  52. setSubscriptionState(
  53. hasPremiumAccess: hasPremiumAccess,
  54. hasEverPurchasedPremium: hasEverPurchasedPremium
  55. )
  56. }
  57. func refreshFreeUsageState() {
  58. remainingFreeAIUses = freeUsageManager.remainingUses
  59. refreshGenerationService()
  60. }
  61. private func refreshGenerationService() {
  62. if injectedGenerationService == nil {
  63. generationService = PostGenerationServiceFactory.make(usesLiveAI: usesLiveAI)
  64. }
  65. }
  66. private static func liveAIEnabled(
  67. hasPremiumAccess: Bool,
  68. hasEverPurchasedPremium: Bool,
  69. freeUsageManager: AIFreeUsageManager
  70. ) -> Bool {
  71. AIConfiguration.usesLiveAI(
  72. hasPremiumAccess: hasPremiumAccess,
  73. canUseFreeTrial: freeUsageManager.canUseFreeLiveAI(
  74. hasEverPurchasedPremium: hasEverPurchasedPremium
  75. )
  76. )
  77. }
  78. var formattedSubreddit: String {
  79. let name = PostDraftValidator.normalizedSubreddit(draft.subreddit)
  80. guard !name.isEmpty else { return "r/subreddit" }
  81. return "r/\(name)"
  82. }
  83. var usesLiveAI: Bool {
  84. AIConfiguration.usesLiveAI(
  85. hasPremiumAccess: hasPremiumAccess,
  86. canUseFreeTrial: freeUsageManager.canUseFreeLiveAI(
  87. hasEverPurchasedPremium: hasEverPurchasedPremium
  88. )
  89. )
  90. }
  91. var generationEngineSubtitle: String {
  92. if usesLiveAI {
  93. if hasPremiumAccess {
  94. return "Create Reddit-ready posts with AI"
  95. }
  96. let remaining = remainingFreeAIUses
  97. let useLabel = remaining == 1 ? "use" : "uses"
  98. return "Create Reddit-ready posts with AI (\(remaining) free \(useLabel) left)"
  99. }
  100. if hasEverPurchasedPremium {
  101. return "Upgrade to Premium for unlimited AI generations"
  102. }
  103. return "Create Reddit-ready posts with AI templates"
  104. }
  105. var hasDraftChanges: Bool {
  106. draft != emptyDraft
  107. }
  108. var canGenerate: Bool {
  109. PostDraftValidator.canGenerate(from: draft, isGenerating: isGenerating)
  110. }
  111. var canExport: Bool {
  112. PostDraftValidator.canExport(draft)
  113. }
  114. var canAddPollOption: Bool {
  115. draft.pollOptions.count < 6
  116. }
  117. var canRemovePollOption: Bool {
  118. draft.pollOptions.count > 2
  119. }
  120. var needsOverwriteConfirmation: Bool {
  121. !draft.title.trimmingCharacters(in: .whitespaces).isEmpty
  122. || !draft.body.trimmingCharacters(in: .whitespaces).isEmpty
  123. || draft.pollOptions.contains { !$0.text.trimmingCharacters(in: .whitespaces).isEmpty }
  124. }
  125. func selectPostType(_ type: RedditPostType) {
  126. draft.postType = type
  127. clearMessages()
  128. }
  129. func addPollOption() {
  130. guard canAddPollOption else { return }
  131. draft.pollOptions.append(PollOption())
  132. clearMessages()
  133. }
  134. func removePollOption(_ option: PollOption) {
  135. guard canRemovePollOption else { return }
  136. draft.pollOptions.removeAll { $0.id == option.id }
  137. clearMessages()
  138. }
  139. func updatePollOption(id: UUID, text: String) {
  140. guard let index = draft.pollOptions.firstIndex(where: { $0.id == id }) else { return }
  141. draft.pollOptions[index].text = text
  142. clearMessages()
  143. }
  144. func setImage(from url: URL) {
  145. releaseImageAccess()
  146. guard url.startAccessingSecurityScopedResource() else {
  147. errorMessage = "Couldn't access the selected image. Try choosing the file again."
  148. successMessage = nil
  149. return
  150. }
  151. imageAccessURL = url
  152. isAccessingImageResource = true
  153. draft.imageFileURL = url
  154. clearMessages()
  155. }
  156. func handleImageImportFailure(_ error: Error) {
  157. errorMessage = UserFacingError.message(for: error)
  158. successMessage = nil
  159. }
  160. func removeImage() {
  161. releaseImageAccess()
  162. draft.imageFileURL = nil
  163. clearMessages()
  164. }
  165. func generatePost() async {
  166. guard canGenerate else { return }
  167. if requiresPaywallForGeneration() {
  168. onPaywallRequired?()
  169. return
  170. }
  171. if needsOverwriteConfirmation {
  172. showOverwriteConfirmation = true
  173. return
  174. }
  175. await performGenerate(replaceExisting: true)
  176. }
  177. func confirmOverwriteAndGenerate() async {
  178. showOverwriteConfirmation = false
  179. await performGenerate(replaceExisting: true)
  180. }
  181. func cancelOverwriteConfirmation() {
  182. showOverwriteConfirmation = false
  183. }
  184. func copyToClipboard() {
  185. do {
  186. try PostDraftValidator.validateForExport(draft)
  187. let content = exportText()
  188. NSPasteboard.general.clearContents()
  189. NSPasteboard.general.setString(content, forType: .string)
  190. successMessage = "Copied to clipboard."
  191. errorMessage = nil
  192. } catch {
  193. errorMessage = UserFacingError.message(for: error)
  194. successMessage = nil
  195. }
  196. }
  197. func resetDraft() {
  198. releaseImageAccess()
  199. draft = PostDraft()
  200. hasGeneratedContent = false
  201. selectedTab = .compose
  202. clearMessages()
  203. }
  204. func restore(from entry: AIHistoryEntry) {
  205. guard case .postGenerator(let storedDraft) = entry.payload else { return }
  206. releaseImageAccess()
  207. draft = storedDraft.postDraft
  208. hasGeneratedContent = true
  209. selectedTab = .preview
  210. clearMessages()
  211. successMessage = "Restored from history."
  212. }
  213. func clearMessages() {
  214. errorMessage = nil
  215. successMessage = nil
  216. }
  217. func notifyDraftEdited() {
  218. clearMessages()
  219. }
  220. func exportText() -> String {
  221. var lines: [String] = []
  222. lines.append("Subreddit: \(formattedSubreddit)")
  223. lines.append("Type: \(draft.postType.title)")
  224. if !draft.title.isEmpty {
  225. lines.append("Title: \(draft.title)")
  226. }
  227. switch draft.postType {
  228. case .text:
  229. if !draft.body.isEmpty { lines.append("\n\(draft.body)") }
  230. case .image:
  231. if let url = draft.imageFileURL {
  232. lines.append("Image: \(url.lastPathComponent)")
  233. }
  234. if !draft.body.isEmpty { lines.append("Caption: \(draft.body)") }
  235. case .link:
  236. if !draft.linkURL.isEmpty { lines.append("URL: \(draft.linkURL)") }
  237. if !draft.body.isEmpty { lines.append("\n\(draft.body)") }
  238. case .video:
  239. if !draft.videoURL.isEmpty { lines.append("Video URL: \(draft.videoURL)") }
  240. if !draft.body.isEmpty { lines.append("\n\(draft.body)") }
  241. case .poll:
  242. lines.append("Duration: \(draft.pollDuration.label)")
  243. for (index, option) in draft.pollOptions.enumerated() where !option.text.isEmpty {
  244. lines.append("Option \(index + 1): \(option.text)")
  245. }
  246. if !draft.body.isEmpty { lines.append("\n\(draft.body)") }
  247. }
  248. var tags: [String] = []
  249. if draft.isNSFW { tags.append("NSFW") }
  250. if draft.isSpoiler { tags.append("Spoiler") }
  251. if draft.isOC { tags.append("OC") }
  252. if !tags.isEmpty { lines.append("Tags: \(tags.joined(separator: ", "))") }
  253. if !draft.flair.isEmpty { lines.append("Flair: \(draft.flair)") }
  254. return lines.joined(separator: "\n")
  255. }
  256. private func performGenerate(replaceExisting: Bool) async {
  257. if requiresPaywallForGeneration() {
  258. onPaywallRequired?()
  259. return
  260. }
  261. isGenerating = true
  262. errorMessage = nil
  263. successMessage = nil
  264. await Task.yield()
  265. let usedLiveAI = usesLiveAI
  266. do {
  267. let result = try await generationService.generatePost(from: draft)
  268. applyGeneratedPost(result, replaceExisting: replaceExisting)
  269. selectedTab = .preview
  270. if usedLiveAI, !hasPremiumAccess, !hasEverPurchasedPremium {
  271. freeUsageManager.recordUse(for: .postGenerator)
  272. refreshFreeUsageState()
  273. }
  274. let engine = usedLiveAI ? "AI" : "local AI templates"
  275. successMessage = "Post generated successfully using \(engine)."
  276. let entry = AIHistoryEntry.fromPostGenerator(draft: draft)
  277. Task { @MainActor in
  278. historyManager.save(entry)
  279. }
  280. } catch {
  281. errorMessage = UserFacingError.message(for: error)
  282. }
  283. isGenerating = false
  284. }
  285. private func applyGeneratedPost(_ result: GeneratedPost, replaceExisting: Bool) {
  286. if replaceExisting || draft.title.trimmingCharacters(in: .whitespaces).isEmpty {
  287. draft.title = String(result.title.prefix(PostDraftValidator.maxTitleLength))
  288. }
  289. if replaceExisting || draft.body.trimmingCharacters(in: .whitespaces).isEmpty {
  290. draft.body = String(result.body.prefix(PostDraftValidator.maxBodyLength))
  291. }
  292. if let flair = result.suggestedFlair, draft.flair.trimmingCharacters(in: .whitespaces).isEmpty {
  293. draft.flair = flair
  294. }
  295. if draft.postType == .poll, let options = result.pollOptions, !options.isEmpty {
  296. if replaceExisting || draft.pollOptions.allSatisfy({ $0.text.trimmingCharacters(in: .whitespaces).isEmpty }) {
  297. draft.pollOptions = options.map { PollOption(text: $0) }
  298. }
  299. }
  300. hasGeneratedContent = true
  301. }
  302. private func requiresPaywallForGeneration() -> Bool {
  303. !hasPremiumAccess && (hasEverPurchasedPremium || remainingFreeAIUses == 0)
  304. }
  305. private func releaseImageAccess() {
  306. if isAccessingImageResource, let imageAccessURL {
  307. imageAccessURL.stopAccessingSecurityScopedResource()
  308. }
  309. imageAccessURL = nil
  310. isAccessingImageResource = false
  311. }
  312. }