Просмотр исходного кода

列车行驶区段计算优化

zhanghaijun лет назад: 2
Родитель
Сommit
0ab12a4bbf

+ 4 - 0
train-server/pom.xml

@@ -54,6 +54,10 @@
             <groupId>cn.hutool</groupId>
             <artifactId>hutool-json</artifactId>
         </dependency>
+        <dependency>
+            <groupId>org.springframework.boot</groupId>
+            <artifactId>spring-boot-starter-cache</artifactId>
+        </dependency>
     </dependencies>
     <build>
         <plugins>

+ 5 - 2
train-server/src/main/java/top/haijunit/train/TrainServerMain.java

@@ -1,8 +1,10 @@
 package top.haijunit.train;
 
+import org.springframework.boot.Banner;
 import org.springframework.boot.SpringApplication;
 import org.springframework.boot.autoconfigure.SpringBootApplication;
 import org.springframework.boot.context.metrics.buffering.BufferingApplicationStartup;
+import org.springframework.cache.annotation.EnableCaching;
 import org.springframework.scheduling.annotation.EnableScheduling;
 
 /**
@@ -10,13 +12,14 @@ import org.springframework.scheduling.annotation.EnableScheduling;
  * @date 2023/11/22 22:50
  * @description 启动类
  */
+@EnableCaching
 @EnableScheduling
 @SpringBootApplication
 public class TrainServerMain {
     public static void main(String[] args) {
         SpringApplication application = new SpringApplication(TrainServerMain.class);
-        // application.setApplicationStartup(new BufferingApplicationStartup(2048));
-        // application.setBannerMode(Banner.Mode.OFF);
+        application.setApplicationStartup(new BufferingApplicationStartup(2048));
+        application.setBannerMode(Banner.Mode.OFF);
         application.run(args);
     }
 }

+ 20 - 0
train-server/src/main/java/top/haijunit/train/domain/constant/CacheNames.java

@@ -0,0 +1,20 @@
+package top.haijunit.train.domain.constant;
+
+/**
+ * @author zhanghaijun
+ * @date 2023/12/18 13:20
+ * @description 缓存组名称常量
+ * <p>
+ * key 格式为 cacheNames#ttl#maxIdleTime#maxSize
+ * <p>
+ * ttl 过期时间 如果设置为0则不过期 默认为0
+ * maxIdleTime 最大空闲时间 根据LRU算法清理空闲数据 如果设置为0则不检测 默认为0
+ * maxSize 组最大长度 根据LRU算法清理溢出数据 如果设置为0则无限长 默认为0
+ * <p>
+ * 例子: test#60s、test#0#60s、test#0#1m#1000、test#1h#0#500
+ */
+public interface CacheNames {
+
+    String SECTION = "section";
+
+}

+ 0 - 4
train-server/src/main/java/top/haijunit/train/listener/ReportTimeHandle.java

@@ -52,10 +52,6 @@ public class ReportTimeHandle {
     @Async
     @Scheduled(fixedDelay = 1000)
     public void runRouteCompute() {
-        if (TrainSimulationHolder.getRouteMap().isEmpty()) {
-            routeComputerService.init();
-            return;
-        }
         for (TrainItem trainItem : TrainSimulationHolder.getTrainItem()) {
             routeComputerService.run(trainItem);
         }

+ 17 - 4
train-server/src/main/java/top/haijunit/train/listener/TrainApplication.java

@@ -10,8 +10,13 @@ import top.haijunit.train.domain.repository.TrainRepository;
 import top.haijunit.train.service.SectionService;
 import top.haijunit.train.simulation.domain.SectionItem;
 import top.haijunit.train.simulation.domain.TrainItem;
+import top.haijunit.train.simulation.train.ComputerService;
+import top.haijunit.train.simulation.train.RouteComputerService;
 import top.haijunit.train.simulation.train.TrainSimulationHolder;
 
+import java.util.ArrayList;
+import java.util.Collection;
+import java.util.LinkedList;
 import java.util.List;
 
 /**
@@ -25,7 +30,8 @@ import java.util.List;
 public class TrainApplication implements ApplicationRunner {
 
     private final TrainRepository trainRepository;
-    private final SectionService sectionService;
+    private final ComputerService computerService;
+    private final RouteComputerService routeComputerService;
 
     @Override
     public void run(ApplicationArguments args) {
@@ -33,15 +39,22 @@ public class TrainApplication implements ApplicationRunner {
         // List<TrainEntity> list = trainRepository.findAllById(new ArrayList<>() {{
         //     add(110001L);
         // }});
+        // 列车初始化
+        ArrayList<TrainItem> trainItems = new ArrayList<>();
         for (TrainEntity train : list) {
             try {
-                SectionItem sectionItem = sectionService.getSection(train.getSectionId());
+                SectionItem sectionItem = computerService.getSectionItem(train.getSectionId());
                 if (null != sectionItem) {
-                    TrainSimulationHolder.addTrainItem(new TrainItem(train, sectionItem.getSectionName(), sectionItem.getDirection()));
+                    trainItems.add(new TrainItem(train, sectionItem.getSectionName(), sectionItem.getDirection()));
                 }
             } catch (Exception exception) {
-                log.error("列车初始化错误,立车:{},区段:{},便批量:{}", train.getId(), train.getSectionId(), train.getSectionOffset());
+                log.error("列车初始化错误, 列车:{}, 区段: {}, 便批量: {}", train.getId(), train.getSectionId(), train.getSectionOffset());
             }
         }
+        for (TrainItem trainItem : trainItems) {
+            List<SectionItem> trainSectionList = routeComputerService.getTrainSectionList(trainItem);
+            TrainSimulationHolder.addRouteList(trainItem.getTrainId(), routeComputerService.getSectionList(new LinkedList<>(trainSectionList), trainItem));
+            TrainSimulationHolder.addTrainItem(trainItem);
+        }
     }
 }

+ 47 - 51
train-server/src/main/java/top/haijunit/train/service/SectionService.java

@@ -1,12 +1,15 @@
 package top.haijunit.train.service;
 
 import cn.hutool.core.collection.CollUtil;
+import cn.hutool.core.util.StrUtil;
 import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
+import org.springframework.cache.annotation.Cacheable;
 import org.springframework.stereotype.Service;
 import top.haijunit.common.utils.StreamUtils;
+import top.haijunit.train.domain.constant.CacheNames;
+import top.haijunit.train.domain.constant.DirectionEnum;
 import top.haijunit.train.domain.constant.SectionStatusEnum;
-import top.haijunit.train.domain.constant.TrainConstant;
 import top.haijunit.train.domain.entity.RouteEntity;
 import top.haijunit.train.domain.entity.SelectionEntity;
 import top.haijunit.train.domain.entity.SignStopEntity;
@@ -14,9 +17,7 @@ import top.haijunit.train.domain.repository.RouteRepository;
 import top.haijunit.train.domain.repository.SelectionRepository;
 import top.haijunit.train.domain.repository.SignStopRepository;
 import top.haijunit.train.simulation.domain.SectionItem;
-import top.haijunit.train.utils.NumberUtil;
 
-import java.math.BigDecimal;
 import java.util.*;
 import java.util.function.Function;
 import java.util.stream.Collectors;
@@ -44,35 +45,26 @@ public class SectionService {
         return list.stream().collect(HashMap::new, (map, entity) -> map.put(entity.getSelectionId(), entity), HashMap::putAll);
     }
 
-    /**
-     * 根据sectionId获取section
-     *
-     * @param sectionId 当前区段
-     * @param offset    偏移量
-     * @param distance  距离 可能是负数,计算车尾的位置
-     * @return 前进指定距离后的所在区段
-     */
-    public SectionItem getSection(Long sectionId, BigDecimal offset, BigDecimal distance) {
-        SectionItem sectionItem = this.getSection(sectionId);
-        BigDecimal decimal = NumberUtil.add(offset, NumberUtil.mul(distance, sectionItem.getDirection().getFactor()));
-        if (NumberUtil.isGreater(decimal, sectionItem.getSectionLength())) {
-            Long sId = switch (sectionItem.getDirection()) {
-                case UP -> sectionItem.getSectionPreId();
-                case DOWN -> sectionItem.getSectionNextId();
-            };
-            return getSection(sId, new BigDecimal("0"), NumberUtil.sub(decimal, sectionItem.getSectionLength()));
-        }
-        return sectionItem;
+    @Cacheable(cacheNames = CacheNames.SECTION, key = "'id-'+#id")
+    public SectionItem getSectionItem(String id) {
+        return StreamUtils.filter(getAllSection(), item -> {
+            return item.getId().equals(id);
+        }).stream().findFirst().orElseThrow(() -> new IllegalArgumentException(String.format("不存在该区段,uid: %s", id)));
     }
 
-    public SectionItem getSection(Long sectionId) {
-        return StreamUtils.filter(getAllSection(), item -> item.getSectionId().equals(sectionId)).stream().findFirst().orElseThrow(() -> new IllegalArgumentException(String.format("不存在该区段,id: %s", sectionId)));
+    @Cacheable(cacheNames = CacheNames.SECTION, key = "'sectionId-'+#sectionId+'-direction-'+#direction.name()")
+    public SectionItem getSectionItem(Long sectionId, DirectionEnum direction) {
+        return StreamUtils.filter(getAllSection(), item -> {
+            return item.getSectionId().equals(sectionId) && item.getDirection().equals(direction);
+        }).stream().findFirst().orElseThrow(() -> new IllegalArgumentException(String.format("不存在该区段,sectionId: %s", sectionId)));
     }
 
+    @Cacheable(cacheNames = CacheNames.SECTION, key = "'routeId-'+#routeId")
     public List<SectionItem> getSectionAll(Long routeId) {
         return StreamUtils.filter(getAllSection(), item -> item.getRouteId().equals(routeId));
     }
 
+    @Cacheable(cacheNames = CacheNames.SECTION, key = "'routeIds-' + T(java.util.Objects).hash(#routeIds)")
     public List<SectionItem> getSectionAll(Collection<Long> routeIds) {
         if (CollUtil.isEmpty(routeIds)) {
             return new ArrayList<>();
@@ -80,32 +72,23 @@ public class SectionService {
         return StreamUtils.filter(getAllSection(), item -> CollUtil.contains(routeIds, item.getRouteId()));
     }
 
-    public Long getRouteNextId(Long sectionId) {
-        SectionItem sectionItem = this.getSection(sectionId);
-        RouteEntity route = routeRepository.findById(sectionItem.getRouteId()).orElseThrow(() -> new IllegalArgumentException(String.format("未找到进路数据,进路Id: %s", sectionItem.getRouteId())));
-        RouteEntity entity = routeRepository.findById(route.getNextId()).orElseThrow(() -> new IllegalArgumentException(String.format("未找到进路数据,进路Id: %s", route.getNextId())));
+    @Cacheable(cacheNames = CacheNames.SECTION + "-RouteNextId", key = "'sectionId-'+#sectionId+'-direction-'+#direction.name()")
+    public Long getRouteNextId(Long sectionId, DirectionEnum direction) {
+        SectionItem sectionItem = this.getSectionItem(sectionId, direction);
+        RouteEntity route = routeRepository.findById(sectionItem.getRouteId()).orElseThrow(() -> new IllegalArgumentException(String.format("未找到进路数据,routeId: %s", sectionItem.getRouteId())));
+        RouteEntity entity = routeRepository.findById(route.getNextId()).orElseThrow(() -> new IllegalArgumentException(String.format("未找到进路数据, routeId: %s", route.getNextId())));
         if (null == entity) {
             throw new IllegalArgumentException(String.format("没有找到下一个进路,区段Id:%s", sectionId));
         }
         return entity.getId();
     }
 
-    public boolean isTrigger(Long sectionId, BigDecimal offset) {
-        SectionItem section = this.getSection(sectionId);
-        RouteEntity route = routeRepository.findById(section.getRouteId()).orElseThrow(() -> new IllegalArgumentException(String.format("不存在该进路,id: %s", section.getRouteId())));
-        if (sectionId.equals(CollUtil.getLast(route.getSelectionList()))) {
-            // 进入最后一个区段才触发下一个进路
-            BigDecimal distance = NumberUtil.add(offset, TrainConstant.ROUTE_TRIGGER_OFFSET);
-            return NumberUtil.isLessOrEqual(section.getSectionLength(), distance);
-        }
-        return false;
-    }
-
     /**
      * 获取所有进路的区段
      *
      * @return 全部进路的区段
      */
+    @Cacheable(cacheNames = CacheNames.SECTION, key = "'section_all'")
     public List<SectionItem> getAllSection() {
         List<SectionItem> sectionList = new ArrayList<>();
         List<RouteEntity> routeList = routeRepository.findAll();
@@ -125,23 +108,25 @@ public class SectionService {
                 sectionItem.setSectionLength(selection.getSelectionLength());
                 sectionItem.setStartKilometer(selection.getStartKilometer());
                 sectionItem.setSwitch(selection.isSwitch());
-                sectionItem.setDirection(selection.getDirection());
+                sectionItem.setTurnBack(selection.getZfState());
                 // 进路数据
+                sectionItem.setId(getSectionUId(route.getId(), selection.getId()));
                 sectionItem.setRouteId(route.getId());
                 sectionItem.setRouteName(route.getName());
+                sectionItem.setDirection(route.getDirection());
                 sectionItem.setSectionStatus(SectionStatusEnum.IDLE);
                 // 区段连接数据
                 if (i == 0) {
                     // 进路的第一个区段
-                    sectionItem.setSectionPreId(this.getRouteSectionLast(route.getId()));
-                    sectionItem.setSectionNextId(selectionIds.get(i + 1));
+                    sectionItem.setPreId(this.getRouteSectionLast(route.getId()));
+                    sectionItem.setNextId(getSectionUId(route.getId(), selectionIds.get(i + 1)));
                 } else if (i + 1 >= selectionIds.size()) {
                     // 进路的最后一个区段
-                    sectionItem.setSectionPreId(selectionIds.get(i - 1));
-                    sectionItem.setSectionNextId(this.getRouteSectionFirst(route.getNextId()));
+                    sectionItem.setPreId(getSectionUId(route.getId(), selectionIds.get(i - 1)));
+                    sectionItem.setNextId(this.getRouteSectionFirst(route.getNextId()));
                 } else {
-                    sectionItem.setSectionPreId(selectionIds.get(i - 1));
-                    sectionItem.setSectionNextId(selectionIds.get(i + 1));
+                    sectionItem.setPreId(getSectionUId(route.getId(), selectionIds.get(i - 1)));
+                    sectionItem.setNextId(getSectionUId(route.getId(), selectionIds.get(i + 1)));
                 }
                 sectionList.add(sectionItem);
             }
@@ -149,13 +134,24 @@ public class SectionService {
         return sectionList;
     }
 
-    private Long getRouteSectionFirst(Long routeId) {
-        Optional<RouteEntity> optional = routeRepository.findById(routeId);
-        return optional.map(item -> CollUtil.getFirst(item.getSelectionList())).orElse(-1L);
+    private String getRouteSectionFirst(Long routeId) {
+        RouteEntity route = routeRepository.findById(routeId).orElseThrow(() -> new IllegalArgumentException(String.format("未找到进路数据,进路Id: %s", routeId)));
+        return getSectionUId(routeId, CollUtil.getFirst(route.getSelectionList()));
     }
 
-    private Long getRouteSectionLast(Long routeNextId) {
+    private String getRouteSectionLast(Long routeNextId) {
         RouteEntity entity = routeRepository.findByNextId(routeNextId);
-        return Optional.of(entity).map(item -> CollUtil.getLast(item.getSelectionList())).orElse(-1L);
+        RouteEntity route = Optional.of(entity).orElseThrow(() -> new IllegalArgumentException(String.format("未找到进路数据,进路Id: %s", routeNextId)));
+        return getSectionUId(route.getId(), CollUtil.getLast(route.getSelectionList()));
+    }
+
+    /**
+     * 进路区段的唯一标识 用于查询唯一的下一个、下一个
+     * @param routeI 所在的进路Id
+     * @param selectionId 区段Id
+     * @return 进路区段Id
+     */
+    private String getSectionUId(long routeI, long selectionId) {
+        return StrUtil.join("-", routeI, selectionId);
     }
 }

+ 9 - 5
train-server/src/main/java/top/haijunit/train/simulation/domain/SectionItem.java

@@ -9,11 +9,13 @@ import java.math.BigDecimal;
 /**
  * @author zhanghaijun
  * @date 2023/12/9 11:08
- * @description [一句话描述该类的功能]
+ * @description 进路区段信息
  */
 @Data
 public class SectionItem {
 
+    // 唯一标识 进路Id-区段Id
+    private String id;
     // 区段Id
     private Long sectionId;
     // 区段名称
@@ -26,6 +28,8 @@ public class SectionItem {
     private BigDecimal startKilometer;
     // 是否是道岔
     private boolean isSwitch;
+    // 是否折返
+    private boolean isTurnBack;
     // 方向
     private DirectionEnum direction;
 
@@ -34,8 +38,8 @@ public class SectionItem {
     // 进路名称
     private String routeName;
 
-    // 下一个区段Id
-    private Long sectionNextId;
-    // 上一个区段Id
-    private Long sectionPreId;
+    // 上一个Id->uid
+    private String preId;
+    // 下一个Id->uid
+    private String nextId;
 }

+ 15 - 5
train-server/src/main/java/top/haijunit/train/simulation/domain/TrainTravel.java

@@ -1,8 +1,11 @@
 package top.haijunit.train.simulation.domain;
 
+import cn.hutool.core.util.StrUtil;
+import liquibase.util.StringUtil;
 import lombok.AllArgsConstructor;
 import lombok.Data;
 import lombok.RequiredArgsConstructor;
+import top.haijunit.train.domain.constant.DirectionEnum;
 import top.haijunit.train.domain.entity.SignStopEntity;
 import top.haijunit.train.utils.NumberUtil;
 
@@ -35,6 +38,13 @@ public class TrainTravel {
         this.signStopMap = signStopMap;
     }
 
+    public TravelItem travel(TrainItem train, BigDecimal distance) {
+        Optional<SectionItem> sectionOptional = sectionList.stream().filter(item -> {
+            return item.getSectionId().equals(train.getBlockId()) && item.getDirection().equals(train.getDirection());
+        }).findFirst();
+        return sectionOptional.map(sectionItem -> travel(sectionItem.getId(), train.getOffset(), distance)).orElse(null);
+    }
+
     /**
      * 列车行驶
      * @param section 当前区段
@@ -42,8 +52,8 @@ public class TrainTravel {
      * @param distance 行驶距离
      * @return 行驶后的位置
      */
-    public TravelItem travel(long sectionId, BigDecimal offset, BigDecimal distance) {
-        Optional<SectionItem> sectionOptional = sectionList.stream().filter(item -> item.getSectionId().equals(sectionId)).findFirst();
+    public TravelItem travel(String id, BigDecimal offset, BigDecimal distance) {
+        Optional<SectionItem> sectionOptional = sectionList.stream().filter(item -> StrUtil.equals(item.getId(), id)).findFirst();
         if (sectionOptional.isEmpty()) {
             return null;
         }
@@ -55,7 +65,7 @@ public class TrainTravel {
         BigDecimal travelDistance = NumberUtil.mul(sectionItem.getDirection().getFactor(), distance.abs());
         BigDecimal travelOffset = NumberUtil.add(offset, travelDistance);
         // 计算停车标
-        SignStopEntity stopEntity = signStopMap.get(sectionId);
+        SignStopEntity stopEntity = signStopMap.get(sectionItem.getSectionId());
         if (null != stopEntity) {
             BigDecimal minInclude = NumberUtil.min(offset.abs(), travelOffset);
             BigDecimal maxInclude = NumberUtil.max(offset.abs(), travelOffset);
@@ -69,10 +79,10 @@ public class TrainTravel {
         }
         if (NumberUtil.isLess(travelOffset, BigDecimal.ZERO)) {
             // <0 上行 行驶到下一个区段
-            return travel(sectionItem.getSectionNextId(), null, travelOffset.abs());
+            return travel(sectionItem.getNextId(), null, travelOffset.abs());
         } else if (NumberUtil.isGreaterOrEqual(travelOffset, sectionItem.getSectionLength())) {
             // >= 区段的长度 下行 实行到下一个区段
-            return travel(sectionItem.getSectionNextId(), BigDecimal.ZERO, travelOffset.subtract(sectionItem.getSectionLength()));
+            return travel(sectionItem.getNextId(), BigDecimal.ZERO, travelOffset.subtract(sectionItem.getSectionLength()));
         } else {
             // 行驶后的位置在当前区段
             return new TravelItem(sectionItem, travelOffset, false);

+ 61 - 0
train-server/src/main/java/top/haijunit/train/simulation/train/ComputerService.java

@@ -0,0 +1,61 @@
+package top.haijunit.train.simulation.train;
+
+import lombok.RequiredArgsConstructor;
+import lombok.extern.slf4j.Slf4j;
+import org.springframework.cache.annotation.Cacheable;
+import org.springframework.stereotype.Service;
+import top.haijunit.common.utils.StreamUtils;
+import top.haijunit.train.domain.constant.CacheNames;
+import top.haijunit.train.domain.constant.DirectionEnum;
+import top.haijunit.train.service.SectionService;
+import top.haijunit.train.simulation.domain.SectionItem;
+import top.haijunit.train.simulation.domain.TrainItem;
+import top.haijunit.train.utils.NumberUtil;
+
+import java.math.BigDecimal;
+
+/**
+ * @author zhanghaijun
+ * @date 2023/12/18 13:43
+ * @description 仿真计算服务
+ */
+@Slf4j
+@Service
+@RequiredArgsConstructor
+public class ComputerService {
+
+    private final SectionService sectionService;
+
+    public SectionItem getSectionItem(TrainItem train, BigDecimal distance) {
+        return getSectionItem(train.getBlockId(), train.getDirection(), train.getOffset(), distance);
+    }
+
+    /**
+     * 根据sectionId获取section
+     *
+     * @param sectionId 当前区段
+     * @param offset    偏移量
+     * @param distance  距离 可能是负数,计算车尾的位置
+     * @return 前进指定距离后的所在区段
+     */
+    public SectionItem getSectionItem(Long sectionId, DirectionEnum direction, BigDecimal offset, BigDecimal distance) {
+        SectionItem sectionItem = sectionService.getSectionItem(sectionId, direction);
+        BigDecimal decimal = NumberUtil.add(offset, NumberUtil.mul(distance, sectionItem.getDirection().getFactor()));
+        if (NumberUtil.isGreater(decimal, sectionItem.getSectionLength())) {
+            String uId = switch (sectionItem.getDirection()) {
+                case UP -> sectionItem.getPreId();
+                case DOWN -> sectionItem.getNextId();
+            };
+            SectionItem item = sectionService.getSectionItem(uId);
+            return getSectionItem(item.getSectionId(), item.getDirection(), new BigDecimal("0"), NumberUtil.sub(decimal, sectionItem.getSectionLength()));
+        }
+        return sectionItem;
+    }
+
+    @Deprecated
+    public SectionItem getSectionItem(Long sectionId) {
+        return StreamUtils.filter(sectionService.getAllSection(), item -> {
+            return item.getSectionId().equals(sectionId);
+        }).stream().findFirst().orElseThrow(() -> new IllegalArgumentException(String.format("不存在该区段,sectionId: %s", sectionId)));
+    }
+}

+ 25 - 9
train-server/src/main/java/top/haijunit/train/simulation/train/RouteComputerService.java

@@ -3,19 +3,22 @@ package top.haijunit.train.simulation.train;
 import cn.hutool.core.collection.CollUtil;
 import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
-import org.springframework.scheduling.annotation.Async;
 import org.springframework.stereotype.Service;
 import top.haijunit.common.utils.ErrorUtil;
 import top.haijunit.common.utils.StreamUtils;
+import top.haijunit.train.domain.constant.DirectionEnum;
 import top.haijunit.train.domain.constant.SectionStatusEnum;
+import top.haijunit.train.domain.constant.TrainConstant;
+import top.haijunit.train.domain.entity.RouteEntity;
+import top.haijunit.train.domain.repository.RouteRepository;
 import top.haijunit.train.service.SectionService;
 import top.haijunit.train.service.RouteService;
 import top.haijunit.train.simulation.domain.SectionItem;
 import top.haijunit.train.simulation.domain.TrainItem;
 import top.haijunit.train.utils.NumberUtil;
 
+import java.math.BigDecimal;
 import java.util.*;
-import java.util.stream.Collectors;
 
 /**
  * @author zhanghaijun
@@ -29,6 +32,8 @@ public class RouteComputerService {
 
     private final RouteService routeService;
     private final SectionService sectionService;
+    private final RouteRepository routeRepository;
+    private final ComputerService computerService;
 
     public synchronized void init() {
         Collection<TrainItem> trainItems = TrainSimulationHolder.getTrainItem();
@@ -65,11 +70,11 @@ public class RouteComputerService {
             routeListAll.addAll(getTrainSectionList(train));
         }
         LinkedList<SectionItem> routeList = this.getSectionList(routeListAll, train);
-        if (routeList.size() > 1 && !sectionService.isTrigger(train.getBlockId(), train.getOffset())) {
+        if (routeList.size() > 1 && !isTrigger(train.getBlockId(), train.getDirection(), train.getOffset())) {
             // 没有触发进路
             return routeList;
         }
-        Long routeId = sectionService.getRouteNextId(train.getBlockId());
+        Long routeId = sectionService.getRouteNextId(train.getBlockId(), train.getDirection());
         if (CollUtil.isNotEmpty(StreamUtils.filter(routeList, item -> item.getRouteId().equals(routeId)))) {
             // 已经添加过了
             return routeList;
@@ -92,12 +97,23 @@ public class RouteComputerService {
         return routeList;
     }
 
-    private List<SectionItem> getTrainSectionList(TrainItem train) {
+    public boolean isTrigger(Long sectionId, DirectionEnum direction, BigDecimal offset) {
+        SectionItem section = sectionService.getSectionItem(sectionId, direction);
+        RouteEntity route = routeRepository.findById(section.getRouteId()).orElseThrow(() -> new IllegalArgumentException(String.format("不存在该进路, routeId: %s", section.getRouteId())));
+        if (sectionId.equals(CollUtil.getLast(route.getSelectionList()))) {
+            // 进入最后一个区段才触发下一个进路
+            BigDecimal distance = NumberUtil.add(offset, TrainConstant.ROUTE_TRIGGER_OFFSET);
+            return NumberUtil.isLessOrEqual(section.getSectionLength(), distance);
+        }
+        return false;
+    }
+
+    public List<SectionItem> getTrainSectionList(TrainItem train) {
         log.info("初始化进路计算,列车: {}", train.getTrainId());
         // 车头所在的区段
-        SectionItem sectionHead = sectionService.getSection(train.getBlockId());
+        SectionItem sectionHead = sectionService.getSectionItem(train.getBlockId(), train.getDirection());
         // 车尾所在的区段
-        SectionItem sectionTail = sectionService.getSection(train.getBlockId(), train.getOffset(), NumberUtil.mul(train.getTrainLength(), -1));
+        SectionItem sectionTail = computerService.getSectionItem(train, NumberUtil.mul(train.getTrainLength(), -1));
         List<SectionItem> sectionAll = sectionService.getSectionAll(new ArrayList<>() {{
             add(sectionHead.getRouteId());
             add(sectionTail.getRouteId());
@@ -112,14 +128,14 @@ public class RouteComputerService {
         return sectionAll;
     }
 
-    private LinkedList<SectionItem> getSectionList(LinkedList<SectionItem> routeList, TrainItem train) {
+    public LinkedList<SectionItem> getSectionList(LinkedList<SectionItem> routeList, TrainItem train) {
         if (CollUtil.isEmpty(routeList)) {
             return new LinkedList<>();
         }
         // 车头所在的区段
         Long sectionIdByHead = train.getBlockId();
         // 车尾的所在的区段
-        Long sectionIdByTail = sectionService.getSection(train.getBlockId(), train.getOffset(), NumberUtil.mul(train.getTrainLength(), -1)).getSectionId();
+        Long sectionIdByTail = computerService.getSectionItem(train, NumberUtil.mul(train.getTrainLength(), -1)).getSectionId();
         Iterator<SectionItem> iterator = routeList.iterator();
         boolean isPass = true;
         while (iterator.hasNext()) {

+ 21 - 49
train-server/src/main/java/top/haijunit/train/simulation/train/TrainComputerService.java

@@ -6,6 +6,7 @@ import lombok.extern.slf4j.Slf4j;
 import org.springframework.stereotype.Service;
 import top.haijunit.common.utils.ErrorUtil;
 import top.haijunit.common.utils.StreamUtils;
+import top.haijunit.train.domain.constant.DirectionEnum;
 import top.haijunit.train.domain.constant.TrainConstant;
 import top.haijunit.train.domain.entity.SignStopEntity;
 import top.haijunit.train.service.SectionService;
@@ -64,7 +65,7 @@ public class TrainComputerService {
             }
             HashMap<Long, SignStopEntity> stopSignalMap = sectionService.getStopSignalMap();
             TrainTravel trainTravel = new TrainTravel(sectionList, stopSignalMap);
-            TrainTravel.TravelItem travel = trainTravel.travel(train.getBlockId(), train.getOffset(), distance);
+            TrainTravel.TravelItem travel = trainTravel.travel(train, distance);
             if (null != travel) {
                 train.setSpeed(CalculatorUtil.velocity(train.getSpeed(), train.getAcceleration(), train.getUpdateTimeMillis(), nowMillis).abs());
                 train.setOffset(travel.getOffset());
@@ -83,8 +84,7 @@ public class TrainComputerService {
             train.setUpdateTimeMillis(nowMillis);
         }
         // 计算MA
-        List<BigDecimal> list = this.getDistanceMa(sectionList, train.getBlockId(), train.getOffset(), 0);
-        BigDecimal distanceMa = list.stream().reduce(BigDecimal.ZERO, BigDecimal::add);
+        BigDecimal distanceMa = this.getDistanceMa(sectionList, train);
         BigDecimal distanceLimit = CalculatorUtil.distanceStopLimit(train.getSpeed());
         // log.info("列车信息:列车:{},速度:{},加速度:{}, 滑行距离:{}", trainItem.getTrainId(), trainItem.getSpeed(), trainItem.getAcceleration(), distance);
         if (NumberUtil.isLessOrEqual(distanceMa, distanceLimit)) {
@@ -103,55 +103,27 @@ public class TrainComputerService {
         return train;
     }
 
-    private List<BigDecimal> getDistanceMa(LinkedList<SectionItem> list, Long sectionId, BigDecimal offset, int count) {
-        List<BigDecimal> maList = new ArrayList<>();
-        if (count > 10) {
-            log.error("循环太多了-- 计算MA,当前count:{}", count);
-        }
-        Optional<SectionItem> sectionItem = list.stream().filter(item -> item.getSectionId().equals(sectionId)).findFirst();
-        if (sectionItem.isPresent()) {
-            SectionItem item = sectionItem.get();
-            maList.add(switch (item.getDirection()) {
-                case UP -> NumberUtil.equals(offset, BigDecimal.ZERO) ? item.getSectionLength().abs() : offset.abs();
-                case DOWN -> NumberUtil.sub(sectionItem.get().getSectionLength(), offset).abs();
-            });
-            if (!list.getLast().getSectionId().equals(sectionId)) {
-                // list 是有序的,最后一个不需要再去找了,避免无限循环
-                List<BigDecimal> distanceMa = getDistanceMa(list, sectionItem.get().getSectionNextId(), BigDecimal.ZERO, count + 1);
-                maList.addAll(distanceMa);
-            }
-        }
-        return maList;
-    }
 
-    /**
-     * 获取列车心行驶后的区段
-     *
-     * @param list      区段数据
-     * @param sectionId 当前的区段
-     * @param offset    当前偏移量
-     * @param distance  行驶的距离
-     * @return 行驶后所在的区段
-     */
-    private List<SectionItem> getSectionByDistance(LinkedList<SectionItem> list, Long sectionId, BigDecimal offset, BigDecimal distance) {
-        Optional<SectionItem> sectionOptional = StreamUtils.filter(list, item -> item.getSectionId().equals(sectionId)).stream().findFirst();
+    private BigDecimal getDistanceMa(LinkedList<SectionItem> list, TrainItem train) {
+        Optional<SectionItem> sectionOptional = list.stream().filter(item -> {
+            return item.getSectionId().equals(train.getBlockId()) && item.getDirection().equals(train.getDirection());
+        }).findFirst();
         if (sectionOptional.isEmpty()) {
-            log.error("没有找到当前列车的所在的区段, 区段Id: {}, offset: {}", sectionId, offset);
-            return new ArrayList<>();
-            // throw new IllegalArgumentException(String.format("没有找到当前列车的所在的区段, 区段Id: %s, offset: %s", sectionId, offset));
+            return BigDecimal.ZERO;
         }
-        List<SectionItem> resultList = new ArrayList<>();
-        SectionItem sectionItem = sectionOptional.get();
-        BigDecimal offsetDistance = NumberUtil.add(offset, NumberUtil.mul(distance, sectionItem.getDirection().getFactor()));
-        if (!NumberUtil.isContain(offsetDistance, BigDecimal.ZERO, sectionItem.getSectionLength())) {
-            // 不在当前的区段
-            BigDecimal distanceNext = switch (sectionItem.getDirection()) {
-                case UP -> offsetDistance.abs();
-                case DOWN -> NumberUtil.sub(offsetDistance, sectionItem.getSectionLength());
-            };
-            resultList.addAll(getSectionByDistance(list, sectionItem.getSectionNextId(), new BigDecimal("0.00"), distanceNext));
+        return getDistanceMa(list, sectionOptional.get(), train.getOffset(), BigDecimal.ZERO);
+    }
+
+    private BigDecimal getDistanceMa(LinkedList<SectionItem> list, SectionItem sectionItem, BigDecimal offset, BigDecimal distance) {
+        BigDecimal sectionDistance = null == offset ? sectionItem.getSectionLength() : switch (sectionItem.getDirection()) {
+            case UP -> offset.abs();
+            case DOWN -> NumberUtil.sub(sectionItem.getSectionLength().abs(), offset.abs());
+        };
+        BigDecimal totalDistance = NumberUtil.add(distance.abs(), sectionDistance.abs());
+        Optional<SectionItem> nextSectionOptional = list.stream().filter(item -> item.getId().equals(sectionItem.getNextId())).findFirst();
+        if (nextSectionOptional.isEmpty()) {
+            return totalDistance;
         }
-        resultList.add(sectionItem);
-        return resultList;
+        return getDistanceMa(list, nextSectionOptional.get(), null, totalDistance);
     }
 }

+ 40 - 14
train-server/src/test/java/top/haijunit/train/test/ComputerRouteTest.java

@@ -1,10 +1,17 @@
 package top.haijunit.train.test;
 
 import lombok.extern.slf4j.Slf4j;
+import org.aspectj.lang.annotation.Before;
+import org.junit.jupiter.api.BeforeAll;
+import org.junit.jupiter.api.BeforeEach;
 import org.junit.jupiter.api.DisplayName;
 import org.junit.jupiter.api.Test;
+import org.junit.jupiter.params.ParameterizedTest;
+import org.junit.jupiter.params.provider.CsvSource;
+import org.junit.jupiter.params.provider.ValueSource;
 import org.springframework.beans.factory.annotation.Autowired;
 import org.springframework.boot.test.context.SpringBootTest;
+import top.haijunit.train.TrainServerMain;
 import top.haijunit.train.domain.entity.SignStopEntity;
 import top.haijunit.train.service.SectionService;
 import top.haijunit.train.simulation.domain.SectionItem;
@@ -20,29 +27,48 @@ import java.util.List;
  * @description 进路触发 单元测试
  */
 @Slf4j
-@SpringBootTest
+@DisplayName("进路测试")
+@SpringBootTest(classes = TrainServerMain.class, webEnvironment = SpringBootTest.WebEnvironment.NONE)
 public class ComputerRouteTest {
 
     @Autowired
     private SectionService sectionService;
+    // 进路区段集合
+    private List<SectionItem> sectionItems;
 
     @DisplayName("进路集合")
-    @Test
+    @BeforeEach
     void localAllRoute() {
-        List<SectionItem> allSection = sectionService.getAllSection();
-        for (SectionItem item : allSection) {
-            log.info("进路:{},区段名称:{},区段Id:{},上一个区段:{},下一个区段:{}", item.getRouteName(), item.getSectionName(), item.getSectionId(), item.getSectionPreId(), item.getSectionNextId());
-        }
+        this.sectionItems = sectionService.getAllSection();
+        this.sectionItems.forEach(item -> log.info("id:{}, 进路:{}, 区段名称:{}, 区段Id:{}, 上一个:{}, 下一个:{}, 道岔:{}, 折返:{}", item.getId(), item.getRouteName(), item.getSectionName(), item.getSectionId(), item.getPreId(), item.getNextId(), item.isSwitch(), item.isTurnBack()));
+        log.info("------ 进路区段集合 ------");
     }
 
-    @DisplayName("列车行驶")
-    @Test
-    void trainTravelTest() {
-        List<SectionItem> sectionItems = sectionService.getAllSection();
-        sectionItems.forEach(item -> log.info("进路:{},区段名称:{},区段Id:{},上一个区段:{},下一个区段:{}", item.getRouteName(), item.getSectionName(), item.getSectionId(), item.getSectionPreId(), item.getSectionNextId()));
+    @DisplayName("列车行驶-上行")
+    @ParameterizedTest
+    @CsvSource({"108001-103009,50,500", "108001-103009,10,500", "108001-103009,10,3500"})
+    void trainUpTravelTest(String id, String offset, Integer distance) {
+        this.trainTravelTest(id, offset, distance);
+    }
+
+    @DisplayName("列车行驶-下行")
+    @ParameterizedTest
+    @CsvSource({"108005-103018,20,500", "108005-103018,145.63,500", "108005-103018,145.63,3500"})
+    void trainDownTravelTest(String id, String offset, Integer distance) {
+        this.trainTravelTest(id, offset, distance);
+    }
+
+    @DisplayName("列车行驶-道岔")
+    @ParameterizedTest
+    @CsvSource({"108003-103005,100,500", "108003-103005,10,200", "108008-103026,0,200", "108008-103026,149,500"})
+    void trainSwitchTravelTest(String id, String offset, Integer distance) {
+        this.trainTravelTest(id, offset, distance);
+    }
+
+    private void trainTravelTest(String id, String offset, Integer distance) {
         HashMap<Long, SignStopEntity> stopSignalMap = sectionService.getStopSignalMap();
-        TrainTravel trainTravel = new TrainTravel(sectionItems, stopSignalMap);
-        TrainTravel.TravelItem travel = trainTravel.travel(103009L, new BigDecimal("10"), new BigDecimal("200"));
-        log.info("当前的位置,是否停车:{},区段:{},偏移量:{}", travel.isStop(), travel.getSection().getSectionName(), travel.getOffset());
+        TrainTravel trainTravel = new TrainTravel(this.sectionItems, stopSignalMap);
+        TrainTravel.TravelItem travel = trainTravel.travel(id, new BigDecimal(offset), new BigDecimal(distance));
+        log.info("当前的位置, 方向:{}, 是否停车:{}, 区段:{}, 偏移量:{}", travel.getSection().getDirection().getDescribe(), travel.isStop(), travel.getSection().getSectionName(), travel.getOffset());
     }
 }