PostGeneratorViewModel.swift 12 KB

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