g1t/services/work/src/messages.rs

182 lines6,281 bytesCodeBlame
1//! Steering: people send the agent at work on a pull request a message,
2//! and the agent receives it at its next step. The sandbox asks for
3//! undelivered messages after each of the agent's tool calls.
4
5use g1t_contracts::time::rfc3339;
6use g1t_contracts::work::*;
7use g1t_contracts::{FailureCode, Outcome, PrincipalKind, new_id};
8use g1t_kit::now_ms;
9use serde::Deserialize;
10use worker::Result;
11
12use crate::Work;
13
14/// The longest message an agent is sent.
15const MAX_MESSAGE_CHARS: usize = 4000;
16
17#[derive(Deserialize)]
18struct MessageRow {
19 id: String,
20 author_name: String,
21 body: String,
22 created_at: String,
23 delivered_at: Option<String>,
24}
25
26impl From<MessageRow> for AgentMessage {
27 fn from(row: MessageRow) -> Self {
28 AgentMessage {
29 id: row.id,
30 author: row.author_name,
31 body: row.body,
32 created_at: row.created_at,
33 delivered_at: row.delivered_at,
34 }
35 }
36}
37
38impl Work {
39 pub(crate) async fn locate_pull(&self, a: LocatePullArgs) -> Result<Outcome<LocatedPull>> {
40 let missing = || Outcome::fail(FailureCode::NotFound, "Pull request not found.");
41 let Some(pull) = self.pull_by_id(&a.id).await? else {
42 return Ok(missing());
43 };
44 let repo: Outcome<g1t_contracts::repos::Repo> = g1t_kit::call(
45 &self.repos,
46 "get_by_id",
47 &g1t_contracts::repos::GetByIdArgs {
48 id: pull.repo_id.clone(),
49 viewer: a.viewer,
50 },
51 )
52 .await?;
53 let Outcome::Ok(repo) = repo else {
54 return Ok(missing());
55 };
56 Ok(Outcome::Ok(LocatedPull {
57 repo: g1t_contracts::repos::RepoPath {
58 namespace: repo.namespace,
59 name: repo.name,
60 },
61 number: pull.number,
62 title: pull.title,
63 status: pull.status,
64 }))
65 }
66
67 /// Every message sent to the agent on a pull request, oldest first.
68 pub(crate) async fn messages(&self, pull_id: &str) -> Result<Vec<AgentMessage>> {
69 Ok(self
70 .db
71 .prepare("SELECT * FROM agent_messages WHERE pull_id = ? ORDER BY created_at, id")
72 .bind(&[pull_id.into()])?
73 .all()
74 .await?
75 .results::<MessageRow>()?
76 .into_iter()
77 .map(AgentMessage::from)
78 .collect())
79 }
80
81 pub(crate) async fn message_agent(&self, a: MessageAgentArgs) -> Result<Outcome<AgentMessage>> {
82 let viewer = Some(a.actor.clone());
83 let (repo, pull) = match self.pull_at(&a.repo, a.number, &viewer).await? {
84 Outcome::Ok(found) => found,
85 Outcome::Fail(failure) => return Ok(Outcome::Fail(failure)),
86 };
87 if !a.actor.verified
88 || (pull.author.id != a.actor.id && !a.actor.is_member(&repo.namespace))
89 {
90 return Ok(Outcome::fail(
91 FailureCode::Forbidden,
92 "Only the pull request's author and members of the workspace can message its agent.",
93 ));
94 }
95 if !pull.status.is_active() {
96 return Ok(Outcome::fail(
97 FailureCode::Conflict,
98 "This pull request is no longer being worked on.",
99 ));
100 }
101 let body = a.body.trim();
102 if body.is_empty() {
103 return Ok(Outcome::fail(FailureCode::Invalid, "Write a message."));
104 }
105 let body: String = body.chars().take(MAX_MESSAGE_CHARS).collect();
106 let now = now_ms();
107 let message = AgentMessage {
108 id: new_id("msg", now),
109 author: a.actor.username.clone(),
110 body,
111 created_at: rfc3339(now),
112 delivered_at: None,
113 };
114 self.db
115 .prepare(
116 "INSERT INTO agent_messages (id, pull_id, author_id, author_name, body, created_at)
117 VALUES (?, ?, ?, ?, ?, ?)",
118 )
119 .bind(&[
120 message.id.as_str().into(),
121 pull.id.as_str().into(),
122 a.actor.id.as_str().into(),
123 message.author.as_str().into(),
124 message.body.as_str().into(),
125 message.created_at.as_str().into(),
126 ])?
127 .run()
128 .await?;
129 self.note(
130 &repo.id,
131 pull.number,
132 (a.actor.id.as_str(), a.actor.username.as_str()),
133 "sent the agent a message",
134 )
135 .await?;
136 Ok(Outcome::Ok(message))
137 }
138
139 /// The undelivered messages, marked delivered and recorded in the
140 /// session, for the agent's sandbox.
141 pub(crate) async fn take_messages(&self, a: TakeMessagesArgs) -> Result<Outcome<Vec<AgentMessage>>> {
142 if a.actor.kind != PrincipalKind::Agent {
143 return Ok(Outcome::fail(
144 FailureCode::Forbidden,
145 "Only g1t's agents take messages.",
146 ));
147 }
148 let viewer = Some(a.actor.clone());
149 let (_, pull) = match self.pull_at(&a.repo, a.number, &viewer).await? {
150 Outcome::Ok(found) => found,
151 Outcome::Fail(failure) => return Ok(Outcome::Fail(failure)),
152 };
153 let now = rfc3339(now_ms());
154 let taken: Vec<AgentMessage> = self
155 .db
156 .prepare(
157 "UPDATE agent_messages SET delivered_at = ?
158 WHERE pull_id = ? AND delivered_at IS NULL
159 RETURNING *",
160 )
161 .bind(&[now.as_str().into(), pull.id.as_str().into()])?
162 .all()
163 .await?
164 .results::<MessageRow>()?
165 .into_iter()
166 .map(AgentMessage::from)
167 .collect();
168 if !taken.is_empty() {
169 let entries: Vec<NewSessionEntry> = taken
170 .iter()
171 .map(|message| NewSessionEntry {
172 kind: SessionEntryKind::Prompt,
173 text: format!("Message from {}: {}", message.author, message.body),
174 tool: None,
175 commit: None,
176 })
177 .collect();
178 self.append_entries(&pull, &entries).await?;
179 }
180 Ok(Outcome::Ok(taken))
181 }
182}