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, PlotAppear, PlotAxis, origin_point,
14 scale::{PlotValue, Scale, ScaleBand, ScaleLinear},
15 tooltip::{CrossLine, Tooltip, TooltipState},
16 },
17};
18
19use super::{
20 AXIS_GAP, ChartAppear, MAX_BAND_WIDTH, TooltipContent, build_band_labels, caller_id,
21 labeled_items, reveal_mask,
22};
23
24#[derive(IntoPlot)]
25pub struct CandlestickChart<T, X, Y>
26where
27 T: 'static,
28 X: Eq + Hash + Into<SharedString> + 'static,
29 Y: PlotValue,
30{
31 data: Vec<T>,
32 x: Option<Rc<dyn Fn(&T) -> X>>,
33 open: Option<Rc<dyn Fn(&T) -> Y>>,
34 high: Option<Rc<dyn Fn(&T) -> Y>>,
35 low: Option<Rc<dyn Fn(&T) -> Y>>,
36 close: Option<Rc<dyn Fn(&T) -> Y>>,
37 tick_margin: usize,
38 body_width_ratio: f32,
39 max_band_width: Pixels,
40 x_axis: bool,
41 grid: bool,
42 bullish: Option<Hsla>,
43 bearish: Option<Hsla>,
44 id: ElementId,
45 interactive: bool,
46 appear: ChartAppear,
47 tooltip_content: TooltipContent<T>,
48}
49
50impl<T, X, Y> CandlestickChart<T, X, Y>
51where
52 X: Eq + Hash + Into<SharedString> + 'static,
53 Y: PlotValue,
54{
55 #[track_caller]
56 pub fn new<I>(data: I) -> Self
57 where
58 I: IntoIterator<Item = T>,
59 {
60 Self {
61 data: data.into_iter().collect(),
62 x: None,
63 open: None,
64 high: None,
65 low: None,
66 close: None,
67 tick_margin: 1,
68 body_width_ratio: 0.8,
69 max_band_width: px(MAX_BAND_WIDTH),
70 x_axis: true,
71 grid: true,
72 bullish: None,
73 bearish: None,
74 id: caller_id(),
75 interactive: true,
76 appear: ChartAppear::default(),
77 tooltip_content: TooltipContent::default(),
78 }
79 }
80
81 pub fn id(mut self, id: impl Into<ElementId>) -> Self {
88 self.id = id.into();
89 self
90 }
91
92 pub fn interactive(mut self, interactive: bool) -> Self {
100 self.interactive = interactive;
101 self
102 }
103
104 pub fn appear(mut self, appear: bool) -> Self {
111 self.appear.set_enabled(appear);
112 self
113 }
114
115 pub fn appear_key(mut self, key: impl Hash) -> Self {
120 self.appear.set_key(key);
121 self
122 }
123
124 pub fn tooltip_title(mut self, title: impl Fn(&T) -> SharedString + 'static) -> Self {
126 self.tooltip_content.set_title(title);
127 self
128 }
129
130 pub fn tooltip_value(
135 mut self,
136 value: impl Fn(&T, usize, f64) -> SharedString + 'static,
137 ) -> Self {
138 self.tooltip_content.set_value(value);
139 self
140 }
141
142 pub fn tooltip_value_color<H>(mut self, color: impl Fn(&T, usize, f64) -> H + 'static) -> Self
148 where
149 H: Into<Hsla>,
150 {
151 self.tooltip_content.set_value_color(color);
152 self
153 }
154
155 pub fn tooltip_content<E>(
162 mut self,
163 content: impl Fn(&T, &mut Window, &mut App) -> E + 'static,
164 ) -> Self
165 where
166 E: IntoElement,
167 {
168 self.tooltip_content.set_content(content);
169 self
170 }
171
172 pub fn x(mut self, x: impl Fn(&T) -> X + 'static) -> Self {
173 self.x = Some(Rc::new(x));
174 self
175 }
176
177 pub fn open(mut self, open: impl Fn(&T) -> Y + 'static) -> Self {
178 self.open = Some(Rc::new(open));
179 self
180 }
181
182 pub fn high(mut self, high: impl Fn(&T) -> Y + 'static) -> Self {
183 self.high = Some(Rc::new(high));
184 self
185 }
186
187 pub fn low(mut self, low: impl Fn(&T) -> Y + 'static) -> Self {
188 self.low = Some(Rc::new(low));
189 self
190 }
191
192 pub fn close(mut self, close: impl Fn(&T) -> Y + 'static) -> Self {
193 self.close = Some(Rc::new(close));
194 self
195 }
196
197 pub fn tick_margin(mut self, tick_margin: usize) -> Self {
198 self.tick_margin = tick_margin;
199 self
200 }
201
202 pub fn body_width_ratio(mut self, ratio: f32) -> Self {
203 self.body_width_ratio = ratio;
204 self
205 }
206
207 pub fn max_band_width(mut self, width: impl Into<Pixels>) -> Self {
212 self.max_band_width = width.into();
213 self
214 }
215
216 pub fn x_axis(mut self, x_axis: bool) -> Self {
220 self.x_axis = x_axis;
221 self
222 }
223
224 pub fn grid(mut self, grid: bool) -> Self {
225 self.grid = grid;
226 self
227 }
228
229 pub fn bullish(mut self, color: impl Into<Hsla>) -> Self {
234 self.bullish = Some(color.into());
235 self
236 }
237
238 pub fn bearish(mut self, color: impl Into<Hsla>) -> Self {
242 self.bearish = Some(color.into());
243 self
244 }
245
246 fn candle_colors(&self, cx: &App) -> (Hsla, Hsla) {
248 (
249 self.bullish.unwrap_or(cx.theme().chart_bullish),
250 self.bearish.unwrap_or(cx.theme().chart_bearish),
251 )
252 }
253
254 fn x_scale(&self, bounds: Bounds<Pixels>) -> Option<ScaleBand<X>> {
257 let x_fn = self.x.as_ref()?;
258 Some(
259 ScaleBand::new(
260 self.data.iter().map(|v| x_fn(v)),
261 [0., bounds.size.width.as_f32()],
262 )
263 .max_band_width(self.max_band_width.as_f32())
264 .padding_inner(0.4)
265 .padding_outer(0.2),
266 )
267 }
268
269 fn plot_height(&self, bounds: Bounds<Pixels>) -> f32 {
271 bounds.size.height.as_f32() - if self.x_axis { AXIS_GAP } else { 0. }
272 }
273}
274
275impl<T, X, Y> Plot for CandlestickChart<T, X, Y>
276where
277 X: Eq + Hash + Into<SharedString> + 'static,
278 Y: PlotValue,
279{
280 fn paint(&mut self, bounds: Bounds<Pixels>, window: &mut Window, cx: &mut App) {
281 let (Some(x_fn), Some(open_fn), Some(high_fn), Some(low_fn), Some(close_fn)) = (
282 self.x.as_ref(),
283 self.open.as_ref(),
284 self.high.as_ref(),
285 self.low.as_ref(),
286 self.close.as_ref(),
287 ) else {
288 return;
289 };
290
291 let height = self.plot_height(bounds);
292
293 let Some(x) = self.x_scale(bounds) else {
295 return;
296 };
297 let band_width = x.band_width();
298
299 let all_values: Vec<Y> = self
301 .data
302 .iter()
303 .flat_map(|d| vec![high_fn(d), low_fn(d), open_fn(d), close_fn(d)])
304 .collect();
305 let y = ScaleLinear::new(all_values, [height, 10.]);
306
307 let mut axis = PlotAxis::new().stroke(cx.theme().border);
309 if self.x_axis {
310 let labels = build_band_labels(
311 &self.data,
312 x_fn.as_ref(),
313 &x,
314 band_width,
315 &labeled_items(self.data.len(), None, self.tick_margin),
316 cx.theme().muted_foreground,
317 );
318 axis = axis.x(height).x_label(labels);
319 }
320 axis.paint(&bounds, window, cx);
321
322 if self.grid {
324 Grid::new()
325 .y((0..=3).map(|i| height * i as f32 / 4.0))
326 .stroke(cx.theme().border)
327 .dash_array(&[px(4.), px(2.)])
328 .paint(&bounds, window);
329 }
330
331 let (bullish, bearish) = self.candle_colors(cx);
333 let origin = bounds.origin;
334 let x_fn = x_fn.clone();
335 let open_fn = open_fn.clone();
336 let high_fn = high_fn.clone();
337 let low_fn = low_fn.clone();
338 let close_fn = close_fn.clone();
339
340 let reveal = reveal_mask(bounds, 0., self.appear.get().progress());
342 window.with_content_mask(reveal, |window| {
343 for d in &self.data {
344 let x_tick = x.tick(&x_fn(d));
345 let Some(x_tick) = x_tick else {
346 continue;
347 };
348
349 let open = open_fn(d);
351 let high = high_fn(d);
352 let low = low_fn(d);
353 let close = close_fn(d);
354
355 let open_y = y.tick(&open);
357 let high_y = y.tick(&high);
358 let low_y = y.tick(&low);
359 let close_y = y.tick(&close);
360
361 let (Some(open_y), Some(high_y), Some(low_y), Some(close_y)) =
362 (open_y, high_y, low_y, close_y)
363 else {
364 continue;
365 };
366
367 let is_bullish = close > open;
369 let color = if is_bullish { bullish } else { bearish };
370
371 let center_x = x_tick + band_width / 2.;
373 let body_width = band_width * self.body_width_ratio;
374 let body_left = center_x - body_width / 2.;
375 let body_right = center_x + body_width / 2.;
376
377 let (wick_top, wick_bottom) = (high_y.min(low_y), high_y.max(low_y));
379 let wick_bounds = Bounds::from_corners(
380 origin_point(px(center_x - 0.5), px(wick_top), origin),
381 origin_point(px(center_x + 0.5), px(wick_bottom), origin),
382 );
383 window.paint_quad(fill(wick_bounds, color));
384
385 let (top, bottom) = if is_bullish {
389 (close_y, open_y)
390 } else {
391 (open_y, close_y)
392 };
393
394 let body_bounds = Bounds::from_corners(
395 origin_point(px(body_left), px(top), origin),
396 origin_point(px(body_right), px(bottom), origin),
397 );
398
399 window.paint_quad(fill(body_bounds, color));
400 }
401 });
402 }
403
404 fn id(&self) -> Option<ElementId> {
405 Some(self.id.clone())
406 }
407
408 fn interactive(&self) -> bool {
409 self.interactive
410 }
411
412 fn appear(&mut self, appear: PlotAppear, _window: &mut Window, _cx: &mut App) {
413 self.appear.update(appear);
414 }
415
416 fn appear_generation(&self) -> Option<u64> {
417 self.appear.generation()
418 }
419
420 fn tooltip_state(
421 &self,
422 position: Point<Pixels>,
423 bounds: Bounds<Pixels>,
424 _cx: &App,
425 ) -> Option<TooltipState> {
426 let x_fn = self.x.as_ref()?;
427 let x = self.x_scale(bounds)?;
428
429 if position.y.as_f32() > self.plot_height(bounds) {
431 return None;
432 }
433
434 let index = x.nearest_index(position.x.as_f32());
435 let d = self.data.get(index)?;
436 let center = x.tick(&x_fn(d))? + x.band_width() / 2.;
437
438 Some(TooltipState::new(
439 index,
440 point(px(center), position.y),
441 vec![],
442 ))
443 }
444
445 fn tooltip(
446 &self,
447 state: &TooltipState,
448 cursor: Point<Pixels>,
449 bounds: Bounds<Pixels>,
450 window: &mut Window,
451 cx: &mut App,
452 ) -> Option<AnyElement> {
453 let (x_fn, open_fn, high_fn, low_fn, close_fn) = (
454 self.x.as_ref()?,
455 self.open.as_ref()?,
456 self.high.as_ref()?,
457 self.low.as_ref()?,
458 self.close.as_ref()?,
459 );
460 let d = self.data.get(state.index)?;
461 let (open, close) = (open_fn(d), close_fn(d));
462 let (bullish, bearish) = self.candle_colors(cx);
463 let color = if close > open { bullish } else { bearish };
464
465 let band_width = self.x_scale(bounds)?.band_width();
469 let cross_line = CrossLine::new(state.cross_line)
470 .span(0., self.plot_height(bounds))
471 .band(px(band_width));
472
473 let tooltip = Tooltip::new(cursor, bounds.size)
474 .gap(px(8.))
475 .cross_line(cross_line);
476 let tooltip = self.tooltip_content.apply(
477 tooltip,
478 d,
479 || Some(x_fn(d).into()),
480 || {
481 [
482 (t!("Chart.open"), open),
483 (t!("Chart.high"), high_fn(d)),
484 (t!("Chart.low"), low_fn(d)),
485 (t!("Chart.close"), close),
486 ]
487 .into_iter()
488 .map(|(label, value)| Some((color, label.to_string().into(), value.to_f64()?)))
489 .collect::<Option<Vec<_>>>()
490 },
491 window,
492 cx,
493 )?;
494
495 Some(tooltip.into_any_element())
496 }
497}