Skip to content

Commit ef150ec

Browse files
authored
Merge pull request #92 from asiniscalchi/refactor-froid
Refactor froid
2 parents 867d30e + 135839a commit ef150ec

26 files changed

Lines changed: 1409 additions & 1591 deletions

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,3 +3,4 @@
33
*.sqlite3
44
data/
55
froid.wiki
6+
.antigravitycli/

src/app.rs

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,7 @@ use crate::{
4444
prompts::{PromptKey, PromptRepository, PromptSource},
4545
version,
4646
workers::{
47-
ReconciliationWorker,
47+
ReconciliationWorker, ReconciliationWorkerConfig,
4848
daily_review::{DailyReviewDeliveryWorker, TelegramDailyReviewSender},
4949
embedding::EmbeddingCycle,
5050
extraction::ExtractionCycle,
@@ -440,7 +440,7 @@ fn spawn_daily_review_delivery_worker(
440440
return Ok(false);
441441
};
442442

443-
let worker = DailyReviewDeliveryWorker::new(
443+
let cycle = DailyReviewDeliveryWorker::new(
444444
JournalRepository::new(pool.clone()),
445445
crate::journal::review::repository::DailyReviewRepository::new(pool.clone()),
446446
daily_review_service,
@@ -450,6 +450,12 @@ fn spawn_daily_review_delivery_worker(
450450
),
451451
config.daily_review_delivery.clone(),
452452
);
453+
let worker_config = ReconciliationWorkerConfig {
454+
enabled: config.daily_review_delivery.enabled,
455+
batch_size: 1,
456+
interval: config.daily_review_delivery.interval,
457+
};
458+
let worker = ReconciliationWorker::new(cycle, worker_config);
453459
let token = shutdown.clone();
454460
workers.spawn(async move {
455461
worker.run_forever(token).await;
@@ -477,7 +483,7 @@ fn spawn_weekly_review_delivery_worker(
477483
return Ok(());
478484
};
479485

480-
let worker = WeeklyReviewDeliveryWorker::new(
486+
let cycle = WeeklyReviewDeliveryWorker::new(
481487
JournalRepository::new(pool.clone()),
482488
crate::journal::week_review::repository::WeeklyReviewRepository::new(pool.clone()),
483489
weekly_review_service,
@@ -487,6 +493,12 @@ fn spawn_weekly_review_delivery_worker(
487493
),
488494
config.weekly_review_delivery.clone(),
489495
);
496+
let worker_config = ReconciliationWorkerConfig {
497+
enabled: config.weekly_review_delivery.enabled,
498+
batch_size: 1,
499+
interval: config.weekly_review_delivery.interval,
500+
};
501+
let worker = ReconciliationWorker::new(cycle, worker_config);
490502
let token = shutdown.clone();
491503
workers.spawn(async move {
492504
worker.run_forever(token).await;

src/dashboard/api.rs

Lines changed: 105 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -51,14 +51,91 @@ struct ExportEnvelope {
5151
messages: Vec<ExportedMessage>,
5252
}
5353

54-
async fn export_messages(State(repo): State<JournalRepository>) -> Response {
55-
let records = match repo.fetch_all_for_export().await {
56-
Ok(records) => records,
57-
Err(err) => {
58-
error!(error = %err, "failed to fetch journal entries for export");
59-
return (StatusCode::INTERNAL_SERVER_ERROR, "failed to load messages").into_response();
54+
#[derive(Debug)]
55+
enum DashboardError {
56+
ExportRepository(sqlx::Error),
57+
Serialization(serde_json::Error),
58+
UnsupportedVersion {
59+
version: u32,
60+
},
61+
ImportConflict {
62+
source: String,
63+
source_conversation_id: String,
64+
source_message_id: String,
65+
},
66+
ImportRepository(sqlx::Error),
67+
}
68+
69+
impl IntoResponse for DashboardError {
70+
fn into_response(self) -> Response {
71+
match self {
72+
Self::ExportRepository(err) => {
73+
error!(error = %err, "failed to fetch journal entries for export");
74+
(StatusCode::INTERNAL_SERVER_ERROR, "failed to load messages").into_response()
75+
}
76+
Self::Serialization(err) => {
77+
error!(error = %err, "failed to serialize journal entries for export");
78+
(
79+
StatusCode::INTERNAL_SERVER_ERROR,
80+
"failed to serialize messages",
81+
)
82+
.into_response()
83+
}
84+
Self::UnsupportedVersion { version } => {
85+
(
86+
StatusCode::BAD_REQUEST,
87+
Json(ImportError {
88+
error: format!(
89+
"unsupported export version {} (supported {}..={})",
90+
version, MIN_SUPPORTED_IMPORT_VERSION, EXPORT_FORMAT_VERSION
91+
),
92+
conflict: None,
93+
}),
94+
)
95+
.into_response()
96+
}
97+
Self::ImportConflict {
98+
source,
99+
source_conversation_id,
100+
source_message_id,
101+
} => {
102+
(
103+
StatusCode::CONFLICT,
104+
Json(ImportError {
105+
error: format!(
106+
"import aborted: entry ({source}, {source_conversation_id}, {source_message_id}) collides with an existing message"
107+
),
108+
conflict: Some(ConflictDetails {
109+
source: source.clone(),
110+
source_conversation_id: source_conversation_id.clone(),
111+
source_message_id: source_message_id.clone(),
112+
}),
113+
}),
114+
)
115+
.into_response()
116+
}
117+
Self::ImportRepository(err) => {
118+
error!(error = %err, "failed to import journal entries");
119+
(
120+
StatusCode::INTERNAL_SERVER_ERROR,
121+
Json(ImportError {
122+
error: "failed to import messages".to_string(),
123+
conflict: None,
124+
}),
125+
)
126+
.into_response()
127+
}
60128
}
61-
};
129+
}
130+
}
131+
132+
async fn export_messages(
133+
State(repo): State<JournalRepository>,
134+
) -> Result<impl IntoResponse, DashboardError> {
135+
let records = repo
136+
.fetch_all_for_export()
137+
.await
138+
.map_err(DashboardError::ExportRepository)?;
62139

63140
let envelope = ExportEnvelope {
64141
version: EXPORT_FORMAT_VERSION,
@@ -76,24 +153,14 @@ async fn export_messages(State(repo): State<JournalRepository>) -> Response {
76153
.collect(),
77154
};
78155

79-
let body = match serde_json::to_vec(&envelope) {
80-
Ok(body) => body,
81-
Err(err) => {
82-
error!(error = %err, "failed to serialize journal entries for export");
83-
return (
84-
StatusCode::INTERNAL_SERVER_ERROR,
85-
"failed to serialize messages",
86-
)
87-
.into_response();
88-
}
89-
};
156+
let body = serde_json::to_vec(&envelope).map_err(DashboardError::Serialization)?;
90157

91158
let filename = format!(
92159
"froid-messages-{}.json",
93160
envelope.exported_at.format("%Y-%m-%d")
94161
);
95162

96-
(
163+
Ok((
97164
StatusCode::OK,
98165
[
99166
(header::CONTENT_TYPE, "application/json".to_string()),
@@ -103,8 +170,7 @@ async fn export_messages(State(repo): State<JournalRepository>) -> Response {
103170
),
104171
],
105172
body,
106-
)
107-
.into_response()
173+
))
108174
}
109175

110176
#[derive(Deserialize)]
@@ -149,23 +215,15 @@ struct ImportError {
149215
async fn import_messages(
150216
State(repo): State<JournalRepository>,
151217
Json(envelope): Json<ImportEnvelope>,
152-
) -> Response {
218+
) -> Result<impl IntoResponse, DashboardError> {
153219
if envelope.version < MIN_SUPPORTED_IMPORT_VERSION || envelope.version > EXPORT_FORMAT_VERSION {
154-
return (
155-
StatusCode::BAD_REQUEST,
156-
Json(ImportError {
157-
error: format!(
158-
"unsupported export version {} (supported {}..={})",
159-
envelope.version, MIN_SUPPORTED_IMPORT_VERSION, EXPORT_FORMAT_VERSION
160-
),
161-
conflict: None,
162-
}),
163-
)
164-
.into_response();
220+
return Err(DashboardError::UnsupportedVersion {
221+
version: envelope.version,
222+
});
165223
}
166224

167225
if envelope.messages.is_empty() {
168-
return (StatusCode::OK, Json(ImportResult { imported: 0 })).into_response();
226+
return Ok((StatusCode::OK, Json(ImportResult { imported: 0 })).into_response());
169227
}
170228

171229
let records: Vec<JournalEntryRecord> = envelope
@@ -181,38 +239,20 @@ async fn import_messages(
181239
})
182240
.collect();
183241

184-
match repo.bulk_import(&records).await {
185-
Ok(imported) => (StatusCode::OK, Json(ImportResult { imported })).into_response(),
186-
Err(BulkImportError::Conflict {
242+
let imported = repo.bulk_import(&records).await.map_err(|err| match err {
243+
BulkImportError::Conflict {
187244
source,
188245
source_conversation_id,
189246
source_message_id,
190-
}) => (
191-
StatusCode::CONFLICT,
192-
Json(ImportError {
193-
error: format!(
194-
"import aborted: entry ({source}, {source_conversation_id}, {source_message_id}) collides with an existing message"
195-
),
196-
conflict: Some(ConflictDetails {
197-
source,
198-
source_conversation_id,
199-
source_message_id,
200-
}),
201-
}),
202-
)
203-
.into_response(),
204-
Err(BulkImportError::Database(err)) => {
205-
error!(error = %err, "failed to import journal entries");
206-
(
207-
StatusCode::INTERNAL_SERVER_ERROR,
208-
Json(ImportError {
209-
error: "failed to import messages".to_string(),
210-
conflict: None,
211-
}),
212-
)
213-
.into_response()
214-
}
215-
}
247+
} => DashboardError::ImportConflict {
248+
source,
249+
source_conversation_id,
250+
source_message_id,
251+
},
252+
BulkImportError::Database(e) => DashboardError::ImportRepository(e),
253+
})?;
254+
255+
Ok((StatusCode::OK, Json(ImportResult { imported })).into_response())
216256
}
217257

218258
#[derive(Serialize)]
@@ -564,7 +604,7 @@ mod tests {
564604
let parsed: Value = serde_json::from_slice(&body).unwrap();
565605
assert_eq!(parsed["imported"], 2);
566606

567-
let entries = repo.fetch_recent("7", 10).await.unwrap();
607+
let entries = repo.fetch_recent(10).await.unwrap();
568608
assert_eq!(entries.len(), 2);
569609
}
570610

@@ -616,7 +656,7 @@ mod tests {
616656
assert_eq!(parsed["conflict"]["source_conversation_id"], "42");
617657
assert_eq!(parsed["conflict"]["source_message_id"], "dup");
618658

619-
let entries = repo.fetch_recent("7", 10).await.unwrap();
659+
let entries = repo.fetch_recent(10).await.unwrap();
620660
assert_eq!(entries.len(), 1);
621661
assert_eq!(entries[0].entry.text, "existing");
622662
}

src/journal/analyzer/journal.rs

Lines changed: 7 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,7 @@ fn map_storage_error(err: sqlx::Error) -> AnalyzerError {
6363
impl JournalReadService for DefaultJournalReadService {
6464
async fn get_recent(
6565
&self,
66-
ctx: &UserContext,
66+
_ctx: &UserContext,
6767
request: GetRecentRequest,
6868
) -> Result<Vec<JournalEntryView>, AnalyzerError> {
6969
let limit = validate_limit(request.limit, MAX_RECENT_LIMIT)?;
@@ -72,14 +72,14 @@ impl JournalReadService for DefaultJournalReadService {
7272
let entries = match (request.from_date, request.to_date_exclusive) {
7373
(None, None) => self
7474
.repository
75-
.fetch_recent(&ctx.user_id, limit)
75+
.fetch_recent(limit)
7676
.await
7777
.map_err(map_storage_error)?,
7878
(from, to) => {
7979
let from = from.unwrap_or(chrono::NaiveDate::from_ymd_opt(1970, 1, 1).unwrap());
8080
let to = to.unwrap_or(chrono::NaiveDate::from_ymd_opt(9999, 1, 1).unwrap());
8181
self.repository
82-
.fetch_in_range(&ctx.user_id, from, to, limit)
82+
.fetch_in_range(from, to, limit)
8383
.await
8484
.map_err(map_storage_error)?
8585
}
@@ -90,7 +90,7 @@ impl JournalReadService for DefaultJournalReadService {
9090

9191
async fn search_text(
9292
&self,
93-
ctx: &UserContext,
93+
_ctx: &UserContext,
9494
request: SearchTextRequest,
9595
) -> Result<Vec<JournalEntryView>, AnalyzerError> {
9696
let limit = validate_limit(request.limit, MAX_TEXT_SEARCH_LIMIT)?;
@@ -104,13 +104,7 @@ impl JournalReadService for DefaultJournalReadService {
104104

105105
let entries = self
106106
.repository
107-
.search_text(
108-
&ctx.user_id,
109-
trimmed,
110-
request.from_date,
111-
request.to_date_exclusive,
112-
limit,
113-
)
107+
.search_text(trimmed, request.from_date, request.to_date_exclusive, limit)
114108
.await
115109
.map_err(map_storage_error)?;
116110

@@ -147,12 +141,12 @@ impl JournalReadService for DefaultJournalReadService {
147141

148142
async fn get_by_id(
149143
&self,
150-
ctx: &UserContext,
144+
_ctx: &UserContext,
151145
id: &str,
152146
) -> Result<Option<JournalEntryView>, AnalyzerError> {
153147
let mut rows = self
154148
.repository
155-
.fetch_by_ids(&ctx.user_id, &[id.to_string()])
149+
.fetch_by_ids(&[id.to_string()])
156150
.await
157151
.map_err(map_storage_error)?;
158152

0 commit comments

Comments
 (0)