| 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 | |
| 5 | use g1t_contracts::time::rfc3339; |
| 6 | use g1t_contracts::work::*; |
| 7 | use g1t_contracts::{FailureCode, Outcome, PrincipalKind, new_id}; |
| 8 | use g1t_kit::now_ms; |
| 9 | use serde::Deserialize; |
| 10 | use worker::Result; |
| 11 | |
| 12 | use crate::Work; |
| 13 | |
| 14 | /// The longest message an agent is sent. |
| 15 | const MAX_MESSAGE_CHARS: usize = 4000; |
| 16 | |
| 17 | #[derive(Deserialize)] |
| 18 | struct MessageRow { |
| 19 | id: String, |
| 20 | author_name: String, |
| 21 | body: String, |
| 22 | created_at: String, |
| 23 | delivered_at: Option<String>, |
| 24 | } |
| 25 | |
| 26 | impl 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 | |
| 38 | impl 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 | } |