anvilsign in

collin/anvil

1use std::path::PathBuf;
2use std::pin::Pin;
3use std::sync::Arc;
4use std::sync::atomic::{
5 AtomicBool,
6 Ordering,
7};
8use std::task::{
9 Context,
10 Poll,
11};
12
13use tokio::io::AsyncRead;
14use tokio::time::{
15 Duration,
16 Sleep,
17 sleep,
18};
19use tokio_util::io::SyncIoBridge;
20
21use crate::error::Result;
22use crate::pack::UploadPackRequest;
23
24pub const RECEIVE_PACK_TIMEOUT: Duration = Duration::from_secs(300);
25const RECEIVE_PACK_IDLE_TIMEOUT: Duration = Duration::from_secs(30);
26
27struct TimedAsyncRead<R> {
28 inner: R,
29 timeout: Duration,
30 sleep: Option<Pin<Box<Sleep>>>,
31 interrupt: Arc<AtomicBool>,
32}
33
34impl<R> TimedAsyncRead<R> {
35 fn new(inner: R, timeout: Duration, interrupt: Arc<AtomicBool>) -> Self {
36 Self {
37 inner,
38 timeout,
39 sleep: None,
40 interrupt,
41 }
42 }
43}
44
45impl<R> AsyncRead for TimedAsyncRead<R>
46where
47 R: AsyncRead + Unpin,
48{
49 fn poll_read(
50 mut self: Pin<&mut Self>,
51 cx: &mut Context<'_>,
52 buf: &mut tokio::io::ReadBuf<'_>,
53 ) -> Poll<std::io::Result<()>> {
54 if self.sleep.is_none() {
55 self.sleep = Some(Box::pin(sleep(self.timeout)));
56 }
57
58 let before = buf.filled().len();
59 match Pin::new(&mut self.inner).poll_read(cx, buf) {
60 Poll::Ready(Ok(())) => {
61 if buf.filled().len() > before {
62 self.sleep = None;
63 }
64 Poll::Ready(Ok(()))
65 }
66 Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
67 Poll::Pending => {
68 if self
69 .sleep
70 .as_mut()
71 .expect("timeout sleep must exist")
72 .as_mut()
73 .poll(cx)
74 .is_ready()
75 {
76 self.interrupt.store(true, Ordering::Relaxed);
77 Poll::Ready(Err(std::io::Error::new(
78 std::io::ErrorKind::TimedOut,
79 "receive-pack read timed out",
80 )))
81 } else {
82 Poll::Pending
83 }
84 }
85 }
86 }
87}
88
89pub struct GitBackend {
90 repo_path: PathBuf,
91}
92
93impl GitBackend {
94 pub fn new(repo_path: PathBuf) -> Self {
95 Self { repo_path }
96 }
97
98 pub fn advertise_refs(&self) -> Result<Vec<u8>> {
99 crate::refs::advertise_refs(&self.repo_path)
100 }
101
102 pub fn advertise_receive_refs(&self) -> Result<Vec<u8>> {
103 crate::receive_pack::advertise_receive_refs(&self.repo_path)
104 }
105
106 pub async fn upload_pack(&self, request: &UploadPackRequest) -> Result<impl AsyncRead + use<>> {
107 crate::pack::generate_pack(&self.repo_path, request)
108 }
109
110 pub async fn receive_pack<R>(&self, request: R) -> Result<Vec<u8>>
111 where
112 R: AsyncRead + Unpin + Send + 'static,
113 {
114 self.receive_pack_with_timeout(request, RECEIVE_PACK_TIMEOUT)
115 .await
116 }
117
118 async fn receive_pack_with_timeout<R>(
119 &self,
120 request: R,
121 timeout_duration: Duration,
122 ) -> Result<Vec<u8>>
123 where
124 R: AsyncRead + Unpin + Send + 'static,
125 {
126 let repo_path = self.repo_path.clone();
127 let interrupt = Arc::new(AtomicBool::new(false));
128 let watchdog_interrupt = interrupt.clone();
129 let watchdog = tokio::spawn(async move {
130 sleep(timeout_duration).await;
131 watchdog_interrupt.store(true, Ordering::Relaxed);
132 });
133
134 let join = tokio::task::spawn_blocking(move || {
135 let request =
136 TimedAsyncRead::new(request, RECEIVE_PACK_IDLE_TIMEOUT, interrupt.clone());
137 let mut request = SyncIoBridge::new(request);
138 crate::receive_pack::receive_pack_with_interrupt(
139 &repo_path,
140 &mut request,
141 interrupt.as_ref(),
142 )
143 })
144 .await
145 .map_err(|e| crate::error::Error::Protocol(format!("receive-pack task panicked: {e}")));
146
147 watchdog.abort();
148 join?
149 }
150}
151
152#[cfg(test)]
153mod tests {
154 use std::process::Command;
155
156 use tempfile::TempDir;
157
158 use super::*;
159
160 fn create_repo_with_commit(root: &std::path::Path) -> PathBuf {
161 let repo_path = root.join("test.git");
162 let work_dir = root.join("work");
163 std::fs::create_dir(&work_dir).unwrap();
164 Command::new("git")
165 .args(["init", "--bare", repo_path.to_str().unwrap()])
166 .output()
167 .unwrap();
168 Command::new("git")
169 .args(["symbolic-ref", "HEAD", "refs/heads/main"])
170 .current_dir(&repo_path)
171 .output()
172 .unwrap();
173 Command::new("git")
174 .args([
175 "clone",
176 repo_path.to_str().unwrap(),
177 work_dir.to_str().unwrap(),
178 ])
179 .output()
180 .unwrap();
181 Command::new("git")
182 .current_dir(&work_dir)
183 .args(["commit", "--allow-empty", "-m", "init"])
184 .env("GIT_AUTHOR_NAME", "Test")
185 .env("GIT_AUTHOR_EMAIL", "t@t.com")
186 .env("GIT_COMMITTER_NAME", "Test")
187 .env("GIT_COMMITTER_EMAIL", "t@t.com")
188 .output()
189 .unwrap();
190 Command::new("git")
191 .current_dir(&work_dir)
192 .args(["push", "origin", "main"])
193 .output()
194 .unwrap();
195 repo_path
196 }
197
198 #[test]
199 fn backend_advertise_refs() {
200 let root = TempDir::new().unwrap();
201 let repo_path = create_repo_with_commit(root.path());
202 let backend = GitBackend::new(repo_path);
203 let output = backend.advertise_refs().unwrap();
204 let output_str = String::from_utf8_lossy(&output);
205 assert!(output_str.contains("refs/heads/main"));
206 }
207
208 #[tokio::test]
209 async fn backend_upload_pack() {
210 let root = TempDir::new().unwrap();
211 let repo_path = create_repo_with_commit(root.path());
212 let repo = gix::open(&repo_path).unwrap();
213 let head = repo.head_id().unwrap();
214
215 let backend = GitBackend::new(repo_path);
216 let request = UploadPackRequest {
217 wants: vec![head.detach()],
218 haves: vec![],
219 done: true,
220 capabilities: Default::default(),
221 shallow: Default::default(),
222 object_ids: None,
223 };
224 let reader = backend.upload_pack(&request).await.unwrap();
225 let mut buf = Vec::new();
226 tokio::io::AsyncReadExt::read_to_end(&mut tokio::io::BufReader::new(reader), &mut buf)
227 .await
228 .unwrap();
229 assert!(buf.windows(4).any(|w| w == b"PACK"));
230 }
231
232 #[tokio::test]
233 async fn backend_receive_pack_times_out_on_stalled_reader() {
234 let root = TempDir::new().unwrap();
235 let repo_path = create_repo_with_commit(root.path());
236 let backend = GitBackend::new(repo_path);
237 let (reader, _writer) = tokio::io::duplex(1);
238
239 let err = backend
240 .receive_pack_with_timeout(reader, Duration::from_millis(50))
241 .await
242 .unwrap_err();
243
244 match err {
245 crate::error::Error::Io(inner) => {
246 assert_eq!(inner.kind(), std::io::ErrorKind::TimedOut);
247 }
248 other => panic!("expected timeout io error, got {other}"),
249 }
250 }
251}