1use std::{hash::Hash, rc::Rc};
2
3use gpui::{
4 AnyElement, App, Bounds, ElementId, Hsla, IntoElement, Pixels, Point, SharedString, Window,
5 fill, point, px,
6};
7use gpui_component_macros::IntoPlot;
8use rust_i18n::t;
9
10use crate::{
11 ActiveTheme,
12 plot::{
13 Grid, Plot, PlotAxis, origin_point,
14 scale::{PlotValue, Scale, ScaleBand, ScaleLinear},
15 tooltip::{CrossLine, Tooltip, TooltipState},
16 },
17};
18
19use super::{
20 AXIS_GAP, MAX_BAND_WIDTH, TooltipContent, build_band_labels, caller_id, labeled_items,
21};
22
23#[derive(IntoPlot)]
24pub struct CandlestickChart<T, X, Y>
25where
26 T: 'static,
27 X: Eq + Hash + Into<SharedString> + 'static,
28 Y: PlotValue,
29{
30 data: Vec<T>,
31 x: Option<Rc<dyn Fn(&T) -> X>>,
32 open: Option<Rc<dyn Fn(&T) -> Y>>,
33 high: Option<Rc<dyn Fn(&T) -> Y>>,
34 low: Option<Rc<dyn Fn(&T) -> Y>>,
35 close: Option<Rc<dyn Fn(&T) -> Y>>,
36 tick_margin: usize,
37 body_width_ratio: f32,
38 max_band_width: Pixels,
39 x_axis: bool,
40 grid: bool,
41 bullish: Option<Hsla>,
42 bearish: Option<Hsla>,
43 id: ElementId,
44 interactive: bool,
45 tooltip_content: TooltipContent<T>,
46}
47
48impl<T, X, Y> CandlestickChart<T, X, Y>
49where
50 X: Eq + Hash + Into<SharedString> + 'static,
51 Y: PlotValue,
52{
53 #[track_caller]
54 pub fn new<I>(data: I) -> Self
55 where
56 I: IntoIterator<Item = T>,
57 {
58 Self {
59 data: data.into_iter().collect(),
60 x: None,
61 open: None,
62 high: None,
63 low: None,
64 close: None,
65 tick_margin: 1,
66 body_width_ratio: 0.8,
67 max_band_width: px(MAX_BAND_WIDTH),
68 x_axis: true,
69 grid: true,
70 bullish: None,
71 bearish: None,
72 id: caller_id(),
73 interactive: true,
74 tooltip_content: TooltipContent::default(),
75 }
76 }
77
78 pub fn id(mut self, id: impl Into<ElementId>) -> Self {
85 self.id = id.into();
86 self
87 }
88
89 pub fn interactive(mut self, interactive: bool) -> Self {
98 self.interactive = interactive;
99 self
100 }
101
102 pub fn tooltip_title(mut self, title: impl Fn(&T) -> SharedString + 'static) -> Self {
104 self.tooltip_content.set_title(title);
105 self
106 }
107
108 pub fn tooltip_value(
113 mut self,
114 value: impl Fn(&T, usize, f64) -> SharedString + 'static,
115 ) -> Self {
116 self.tooltip_content.set_value(value);
117 self
118 }
119
120 pub fn tooltip_value_color<H>(mut self, color: impl Fn(&T, usize, f64) -> H + 'static) -> Self
126 where
127 H: Into<Hsla>,
128 {
129 self.tooltip_content.set_value_color(color);
130 self
131 }
132
133 pub fn tooltip_content<E>(
140 mut self,
141 content: impl Fn(&T, &mut Window, &mut App) -> E + 'static,
142 ) -> Self
143 where
144 E: IntoElement,
145 {
146 self.tooltip_content.set_content(content);
147 self
148 }
149
150 pub fn x(mut self, x: impl Fn(&T) -> X + 'static) -> Self {
151 self.x = Some(Rc::new(x));
152 self
153 }
154
155 pub fn open(mut self, open: impl Fn(&T) -> Y + 'static) -> Self {
156 self.open = Some(Rc::new(open));
157 self
158 }
159
160 pub fn high(mut self, high: impl Fn(&T) -> Y + 'static) -> Self {
161 self.high = Some(Rc::new(high));
162 self
163 }
164
165 pub fn low(mut self, low: impl Fn(&T) -> Y + 'static) -> Self {
166 self.low = Some(Rc::new(low));
167 self
168 }
169
170 pub fn close(mut self, close: impl Fn(&T) -> Y + 'static) -> Self {
171 self.close = Some(Rc::new(close));
172 self
173 }
174
175 pub fn tick_margin(mut self, tick_margin: usize) -> Self {
176 self.tick_margin = tick_margin;
177 self
178 }
179
180 pub fn body_width_ratio(mut self, ratio: f32) -> Self {
181 self.body_width_ratio = ratio;
182 self
183 }
184
185 pub fn max_band_width(mut self, width: impl Into<Pixels>) -> Self {
190 self.max_band_width = width.into();
191 self
192 }
193
194 pub fn x_axis(mut self, x_axis: bool) -> Self {
198 self.x_axis = x_axis;
199 self
200 }
201
202 pub fn grid(mut self, grid: bool) -> Self {
203 self.grid = grid;
204 self
205 }
206
207 pub fn bullish(mut self, color: impl Into<Hsla>) -> Self {
212 self.bullish = Some(color.into());
213 self
214 }
215
216 pub fn bearish(mut self, color: impl Into<Hsla>) -> Self {
220 self.bearish = Some(color.into());
221 self
222 }
223
224 fn candle_colors(&self, cx: &App) -> (Hsla, Hsla) {
226 (
227 self.bullish.unwrap_or(cx.theme().chart_bullish),
228 self.bearish.unwrap_or(cx.theme().chart_bearish),
229 )
230 }
231
232 fn x_scale(&self, bounds: Bounds<Pixels>) -> Option<ScaleBand<X>> {
235 let x_fn = self.x.as_ref()?;
236 Some(
237 ScaleBand::new(
238 self.data.iter().map(|v| x_fn(v)),
239 [0., bounds.size.width.as_f32()],
240 )
241 .max_band_width(self.max_band_width.as_f32())
242 .padding_inner(0.4)
243 .padding_outer(0.2),
244 )
245 }
246
247 fn plot_height(&self, bounds: Bounds<Pixels>) -> f32 {
249 bounds.size.height.as_f32() - if self.x_axis { AXIS_GAP } else { 0. }
250 }
251}
252
253impl<T, X, Y> Plot for CandlestickChart<T, X, Y>
254where
255 X: Eq + Hash + Into<SharedString> + 'static,
256 Y: PlotValue,
257{
258 fn paint(&mut self, bounds: Bounds<Pixels>, window: &mut Window, cx: &mut App) {
259 let (Some(x_fn), Some(open_fn), Some(high_fn), Some(low_fn), Some(close_fn)) = (
260 self.x.as_ref(),
261 self.open.as_ref(),
262 self.high.as_ref(),
263 self.low.as_ref(),
264 self.close.as_ref(),
265 ) else {
266 return;
267 };
268
269 let height = self.plot_height(bounds);
270
271 let Some(x) = self.x_scale(bounds) else {
273 return;
274 };
275 let band_width = x.band_width();
276
277 let all_values: Vec<Y> = self
279 .data
280 .iter()
281 .flat_map(|d| vec![high_fn(d), low_fn(d), open_fn(d), close_fn(d)])
282 .collect();
283 let y = ScaleLinear::new(all_values, [height, 10.]);
284
285 let mut axis = PlotAxis::new().stroke(cx.theme().border);
287 if self.x_axis {
288 let labels = build_band_labels(
289 &self.data,
290 x_fn.as_ref(),
291 &x,
292 band_width,
293 &labeled_items(self.data.len(), None, self.tick_margin),
294 cx.theme().muted_foreground,
295 );
296 axis = axis.x(height).x_label(labels);
297 }
298 axis.paint(&bounds, window, cx);
299
300 if self.grid {
302 Grid::new()
303 .y((0..=3).map(|i| height * i as f32 / 4.0))
304 .stroke(cx.theme().border)
305 .dash_array(&[px(4.), px(2.)])
306 .paint(&bounds, window);
307 }
308
309 let (bullish, bearish) = self.candle_colors(cx);
311 let origin = bounds.origin;
312 let x_fn = x_fn.clone();
313 let open_fn = open_fn.clone();
314 let high_fn = high_fn.clone();
315 let low_fn = low_fn.clone();
316 let close_fn = close_fn.clone();
317
318 for d in &self.data {
319 let x_tick = x.tick(&x_fn(d));
320 let Some(x_tick) = x_tick else {
321 continue;
322 };
323
324 let open = open_fn(d);
326 let high = high_fn(d);
327 let low = low_fn(d);
328 let close = close_fn(d);
329
330 let open_y = y.tick(&open);
332 let high_y = y.tick(&high);
333 let low_y = y.tick(&low);
334 let close_y = y.tick(&close);
335
336 let (Some(open_y), Some(high_y), Some(low_y), Some(close_y)) =
337 (open_y, high_y, low_y, close_y)
338 else {
339 continue;
340 };
341
342 let is_bullish = close > open;
344 let color = if is_bullish { bullish } else { bearish };
345
346 let center_x = x_tick + band_width / 2.;
348 let body_width = band_width * self.body_width_ratio;
349 let body_left = center_x - body_width / 2.;
350 let body_right = center_x + body_width / 2.;
351
352 let (wick_top, wick_bottom) = (high_y.min(low_y), high_y.max(low_y));
354 let wick_bounds = Bounds::from_corners(
355 origin_point(px(center_x - 0.5), px(wick_top), origin),
356 origin_point(px(center_x + 0.5), px(wick_bottom), origin),
357 );
358 window.paint_quad(fill(wick_bounds, color));
359
360 let (top, bottom) = if is_bullish {
364 (close_y, open_y)
365 } else {
366 (open_y, close_y)
367 };
368
369 let body_bounds = Bounds::from_corners(
370 origin_point(px(body_left), px(top), origin),
371 origin_point(px(body_right), px(bottom), origin),
372 );
373
374 window.paint_quad(fill(body_bounds, color));
375 }
376 }
377
378 fn id(&self) -> Option<ElementId> {
379 self.interactive.then(|| self.id.clone())
380 }
381
382 fn tooltip_state(
383 &self,
384 position: Point<Pixels>,
385 bounds: Bounds<Pixels>,
386 _cx: &App,
387 ) -> Option<TooltipState> {
388 let x_fn = self.x.as_ref()?;
389 let x = self.x_scale(bounds)?;
390
391 if position.y.as_f32() > self.plot_height(bounds) {
393 return None;
394 }
395
396 let index = x.nearest_index(position.x.as_f32());
397 let d = self.data.get(index)?;
398 let center = x.tick(&x_fn(d))? + x.band_width() / 2.;
399
400 Some(TooltipState::new(
401 index,
402 point(px(center), position.y),
403 vec![],
404 ))
405 }
406
407 fn tooltip(
408 &self,
409 state: &TooltipState,
410 cursor: Point<Pixels>,
411 bounds: Bounds<Pixels>,
412 window: &mut Window,
413 cx: &mut App,
414 ) -> Option<AnyElement> {
415 let (x_fn, open_fn, high_fn, low_fn, close_fn) = (
416 self.x.as_ref()?,
417 self.open.as_ref()?,
418 self.high.as_ref()?,
419 self.low.as_ref()?,
420 self.close.as_ref()?,
421 );
422 let d = self.data.get(state.index)?;
423 let (open, close) = (open_fn(d), close_fn(d));
424 let (bullish, bearish) = self.candle_colors(cx);
425 let color = if close > open { bullish } else { bearish };
426
427 let band_width = self.x_scale(bounds)?.band_width();
431 let cross_line = CrossLine::new(state.cross_line)
432 .span(0., self.plot_height(bounds))
433 .band(px(band_width));
434
435 let tooltip = Tooltip::new(cursor, bounds.size)
436 .gap(px(8.))
437 .cross_line(cross_line);
438 let tooltip = self.tooltip_content.apply(
439 tooltip,
440 d,
441 || Some(x_fn(d).into()),
442 || {
443 [
444 (t!("Chart.open"), open),
445 (t!("Chart.high"), high_fn(d)),
446 (t!("Chart.low"), low_fn(d)),
447 (t!("Chart.close"), close),
448 ]
449 .into_iter()
450 .map(|(label, value)| Some((color, label.to_string().into(), value.to_f64()?)))
451 .collect::<Option<Vec<_>>>()
452 },
453 window,
454 cx,
455 )?;
456
457 Some(tooltip.into_any_element())
458 }
459}