anvilsign in

collin/anvil

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