Skip to content

Commit d3a01c5

Browse files
authored
Merge pull request #209 from smly/fix/prefer-non-red-five
fix: prefer non-red 5 when resolving discard actions
2 parents 0c1e575 + a8daeef commit d3a01c5

2 files changed

Lines changed: 179 additions & 22 deletions

File tree

riichienv-core/src/observation/mod.rs

Lines changed: 97 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@ pub(crate) mod sequence_features;
1212
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
1313
use serde::{Deserialize, Serialize};
1414

15-
use crate::action::{Action, ActionEncoder};
15+
use crate::action::{Action, ActionEncoder, ActionType};
1616
use crate::errors::{RiichiError, RiichiResult};
1717
use crate::types::Meld;
1818

@@ -116,16 +116,23 @@ impl Observation {
116116

117117
pub fn find_action(&self, action_id: usize) -> Option<Action> {
118118
let encoder = ActionEncoder::FourPlayer;
119-
self._legal_actions
120-
.iter()
121-
.find(|a| {
122-
if let Ok(idx) = encoder.encode(a) {
123-
(idx as usize) == action_id
124-
} else {
125-
false
126-
}
127-
})
128-
.cloned()
119+
// Prefer non-red-five candidates so that 5m/5p/5s discards do not
120+
// accidentally drop the akadora when a normal 5 is also legal.
121+
let mut fallback: Option<&Action> = None;
122+
for action in &self._legal_actions {
123+
let Ok(idx) = encoder.encode(action) else {
124+
continue;
125+
};
126+
if (idx as usize) != action_id {
127+
continue;
128+
}
129+
if is_red_five_discard(action) {
130+
fallback.get_or_insert(action);
131+
} else {
132+
return Some(action.clone());
133+
}
134+
}
135+
fallback.cloned()
129136
}
130137

131138
/// Return absolute player indices in relative order: [self, shimocha, toimen, kamicha].
@@ -159,3 +166,82 @@ impl Observation {
159166
Ok(obs)
160167
}
161168
}
169+
170+
fn is_red_five_discard(action: &Action) -> bool {
171+
matches!(action.action_type, ActionType::Discard)
172+
&& matches!(action.tile, Some(16) | Some(52) | Some(88))
173+
}
174+
175+
#[cfg(test)]
176+
mod tests {
177+
use super::*;
178+
179+
fn obs_with_actions(actions: Vec<Action>) -> Observation {
180+
Observation::new(
181+
0,
182+
[vec![], vec![], vec![], vec![]],
183+
[vec![], vec![], vec![], vec![]],
184+
[vec![], vec![], vec![], vec![]],
185+
vec![],
186+
[25000, 25000, 25000, 25000],
187+
[false, false, false, false],
188+
actions,
189+
vec![],
190+
0,
191+
0,
192+
0,
193+
0,
194+
0,
195+
vec![],
196+
false,
197+
[None, None, None, None],
198+
[None, None, None, None],
199+
None,
200+
None,
201+
)
202+
}
203+
204+
fn discard(tile: u8) -> Action {
205+
Action::new(ActionType::Discard, Some(tile), vec![], Some(0))
206+
}
207+
208+
#[test]
209+
fn find_action_prefers_non_red_5m() {
210+
// 5m action_id = 16 / 4 = 4; legal actions list red 5m first.
211+
let obs = obs_with_actions(vec![discard(16), discard(17), discard(18), discard(19)]);
212+
let chosen = obs.find_action(4).expect("discard 5m should resolve");
213+
assert_eq!(chosen.tile, Some(17), "non-red 5m must win over red 5m");
214+
}
215+
216+
#[test]
217+
fn find_action_prefers_non_red_5p() {
218+
// 5p action_id = 52 / 4 = 13.
219+
let obs = obs_with_actions(vec![discard(52), discard(54)]);
220+
let chosen = obs.find_action(13).expect("discard 5p should resolve");
221+
assert_eq!(chosen.tile, Some(54));
222+
}
223+
224+
#[test]
225+
fn find_action_prefers_non_red_5s() {
226+
// 5s action_id = 88 / 4 = 22.
227+
let obs = obs_with_actions(vec![discard(88), discard(90)]);
228+
let chosen = obs.find_action(22).expect("discard 5s should resolve");
229+
assert_eq!(chosen.tile, Some(90));
230+
}
231+
232+
#[test]
233+
fn find_action_falls_back_to_red_when_only_red_legal() {
234+
// Only the red 5m is in the hand; we still need a working discard.
235+
let obs = obs_with_actions(vec![discard(16)]);
236+
let chosen = obs.find_action(4).expect("red-only 5m must still resolve");
237+
assert_eq!(chosen.tile, Some(16));
238+
}
239+
240+
#[test]
241+
fn find_action_unaffected_for_non_five_tiles() {
242+
// 4m action_id = 12 / 4 = 3.
243+
let obs = obs_with_actions(vec![discard(12), discard(13)]);
244+
let chosen = obs.find_action(3).expect("discard 4m should resolve");
245+
assert_eq!(chosen.tile, Some(12), "first match wins for non-red ties");
246+
}
247+
}

riichienv-core/src/observation_3p/mod.rs

Lines changed: 82 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@ mod python;
88
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
99
use serde::{Deserialize, Serialize};
1010

11-
use crate::action::{Action, Action3P};
11+
use crate::action::{Action, Action3P, ActionType};
1212
use crate::errors::{RiichiError, RiichiResult};
1313
use crate::types::Meld;
1414

@@ -107,16 +107,23 @@ impl Observation3P {
107107
}
108108

109109
pub fn find_action(&self, action_id: usize) -> Option<Action3P> {
110-
self._legal_actions
111-
.iter()
112-
.find(|a| {
113-
if let Ok(idx) = a.encode() {
114-
(idx as usize) == action_id
115-
} else {
116-
false
117-
}
118-
})
119-
.cloned()
110+
// Prefer non-red-five candidates so that 5m/5p/5s discards do not
111+
// accidentally drop the akadora when a normal 5 is also legal.
112+
let mut fallback: Option<&Action3P> = None;
113+
for action in &self._legal_actions {
114+
let Ok(idx) = action.encode() else {
115+
continue;
116+
};
117+
if (idx as usize) != action_id {
118+
continue;
119+
}
120+
if is_red_five_discard(&action.0) {
121+
fallback.get_or_insert(action);
122+
} else {
123+
return Some(action.clone());
124+
}
125+
}
126+
fallback.cloned()
120127
}
121128

122129
/// Return absolute player indices in relative order: [self, next, prev].
@@ -150,3 +157,67 @@ impl Observation3P {
150157
Ok(obs)
151158
}
152159
}
160+
161+
fn is_red_five_discard(action: &Action) -> bool {
162+
matches!(action.action_type, ActionType::Discard)
163+
&& matches!(action.tile, Some(16) | Some(52) | Some(88))
164+
}
165+
166+
#[cfg(test)]
167+
mod tests {
168+
use super::*;
169+
170+
fn obs_with_actions(actions: Vec<Action>) -> Observation3P {
171+
Observation3P::new(
172+
0,
173+
[vec![], vec![], vec![]],
174+
[vec![], vec![], vec![]],
175+
[vec![], vec![], vec![]],
176+
vec![],
177+
[35000, 35000, 35000],
178+
[false, false, false],
179+
actions,
180+
vec![],
181+
0,
182+
0,
183+
0,
184+
0,
185+
0,
186+
vec![],
187+
false,
188+
[None, None, None],
189+
[None, None, None],
190+
None,
191+
None,
192+
)
193+
}
194+
195+
fn discard(tile: u8) -> Action {
196+
Action::new(ActionType::Discard, Some(tile), vec![], Some(0))
197+
}
198+
199+
#[test]
200+
fn find_action_3p_prefers_non_red_5p() {
201+
let obs = obs_with_actions(vec![discard(52), discard(54)]);
202+
// 5p compact id matches the encoded discard id of either action.
203+
let id = obs._legal_actions[0].encode().unwrap() as usize;
204+
let chosen = obs.find_action(id).expect("discard 5p should resolve");
205+
assert_eq!(chosen.0.tile, Some(54), "non-red 5p must win over red 5p");
206+
}
207+
208+
#[test]
209+
fn find_action_3p_prefers_non_red_5s() {
210+
let obs = obs_with_actions(vec![discard(88), discard(90)]);
211+
let id = obs._legal_actions[0].encode().unwrap() as usize;
212+
let chosen = obs.find_action(id).expect("discard 5s should resolve");
213+
assert_eq!(chosen.0.tile, Some(90));
214+
}
215+
216+
#[test]
217+
fn find_action_3p_falls_back_to_red_when_only_red_legal() {
218+
let obs = obs_with_actions(vec![discard(52)]);
219+
let id = obs._legal_actions[0].encode().unwrap() as usize;
220+
let chosen = obs.find_action(id).expect("red-only 5p must still resolve");
221+
assert_eq!(chosen.0.tile, Some(52));
222+
}
223+
}

0 commit comments

Comments
 (0)